1use anyhow::{Context, Result, bail};
10use atlas_core::kimi_k3::{KDA_L2_EPS, KdaConfig, KdaState};
11use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
12use spark_runtime::kernel_args::{KernelLaunch, div_ceil};
13
14pub const MODULE: &str = "kda_decode";
16pub const CONV_ENTRY: &str = "k3_kda_conv_update_f32";
17pub const RECURRENT_ENTRY: &str = "k3_kda_recurrent_step_f32";
18const CONV_BLOCK: u32 = 128;
19
20#[derive(Clone, Copy, Debug)]
21pub struct K3KdaDecodeKernels {
22 pub conv: KernelHandle,
23 pub recurrent: KernelHandle,
24}
25
26impl K3KdaDecodeKernels {
27 pub fn resolve(gpu: &dyn GpuBackend) -> Result<Self> {
28 Ok(Self {
29 conv: gpu.kernel(MODULE, CONV_ENTRY)?,
30 recurrent: gpu.kernel(MODULE, RECURRENT_ENTRY)?,
31 })
32 }
33}
34
35fn f32_bytes(v: &[f32]) -> Vec<u8> {
36 v.iter().flat_map(|x| x.to_le_bytes()).collect()
37}
38
39fn bytes_f32(b: &[u8]) -> Vec<f32> {
40 b.chunks_exact(4)
41 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
42 .collect()
43}
44
45fn up(gpu: &dyn GpuBackend, v: &[f32], hold: &mut Vec<DevicePtr>) -> Result<DevicePtr> {
46 let b = f32_bytes(v);
47 let p = gpu.alloc(b.len().max(1))?;
48 hold.push(p);
49 gpu.copy_h2d(&b, p)?;
50 Ok(p)
51}
52
53pub struct KdaDeviceState {
59 pub conv: DevicePtr,
60 pub recurrent: DevicePtr,
61}
62
63impl KdaDeviceState {
64 pub fn alloc_and_upload(gpu: &dyn GpuBackend, host: &KdaState) -> Result<Self> {
65 let conv_b = f32_bytes(&host.conv);
66 let rec_b = f32_bytes(&host.recurrent);
67 let conv = gpu.alloc(conv_b.len().max(1))?;
68 let recurrent = gpu.alloc(rec_b.len().max(1))?;
69 gpu.copy_h2d(&conv_b, conv)?;
70 gpu.copy_h2d(&rec_b, recurrent)?;
71 Ok(Self { conv, recurrent })
72 }
73
74 pub fn download(&self, gpu: &dyn GpuBackend, host: &mut KdaState) -> Result<()> {
75 let mut conv_b = vec![0u8; host.conv.len() * 4];
76 let mut rec_b = vec![0u8; host.recurrent.len() * 4];
77 gpu.copy_d2h(self.conv, &mut conv_b)?;
78 gpu.copy_d2h(self.recurrent, &mut rec_b)?;
79 host.conv = bytes_f32(&conv_b);
80 host.recurrent = bytes_f32(&rec_b);
81 Ok(())
82 }
83
84 pub fn free(self, gpu: &dyn GpuBackend) -> Result<()> {
85 gpu.free(self.conv)?;
86 gpu.free(self.recurrent)?;
87 Ok(())
88 }
89}
90
91pub fn launch_k3_kda_decode_token(
93 gpu: &dyn GpuBackend,
94 kernels: &K3KdaDecodeKernels,
95 x_qkv: &[f32],
96 conv_w: &[f32],
97 gate: &[f32],
98 beta: &[f32],
99 cfg: &KdaConfig,
100 state: &mut KdaState,
101 stream: u64,
102) -> Result<Vec<f32>> {
103 launch_k3_kda_decode_inner(
104 gpu,
105 kernels,
106 x_qkv,
107 conv_w,
108 gate,
109 beta,
110 cfg,
111 Some(state),
112 None,
113 stream,
114 )
115}
116
117pub fn launch_k3_kda_decode_token_on_device(
119 gpu: &dyn GpuBackend,
120 kernels: &K3KdaDecodeKernels,
121 x_qkv: &[f32],
122 conv_w: &[f32],
123 gate: &[f32],
124 beta: &[f32],
125 cfg: &KdaConfig,
126 device: &KdaDeviceState,
127 stream: u64,
128) -> Result<Vec<f32>> {
129 launch_k3_kda_decode_inner(
130 gpu,
131 kernels,
132 x_qkv,
133 conv_w,
134 gate,
135 beta,
136 cfg,
137 None,
138 Some(device),
139 stream,
140 )
141}
142
143#[allow(clippy::too_many_arguments)]
144fn launch_k3_kda_decode_inner(
145 gpu: &dyn GpuBackend,
146 kernels: &K3KdaDecodeKernels,
147 x_qkv: &[f32],
148 conv_w: &[f32],
149 gate: &[f32],
150 beta: &[f32],
151 cfg: &KdaConfig,
152 mut host: Option<&mut KdaState>,
153 device: Option<&KdaDeviceState>,
154 stream: u64,
155) -> Result<Vec<f32>> {
156 let (h, d, k) = (cfg.heads, cfg.head_dim, cfg.conv_kernel);
157 let c = cfg.conv_dim();
158 if x_qkv.len() != c {
159 bail!("k3 kda: x_qkv {} != conv_dim {c}", x_qkv.len());
160 }
161 if conv_w.len() != cfg.conv_elems() {
162 bail!("k3 kda: conv_w {} != {}", conv_w.len(), cfg.conv_elems());
163 }
164 if gate.len() != cfg.qkv_dim() || beta.len() != h {
165 bail!("k3 kda: gate/beta rank");
166 }
167 if let Some(state) = host.as_ref()
168 && (state.conv.len() != cfg.conv_elems() || state.recurrent.len() != cfg.recurrent_elems())
169 {
170 bail!("k3 kda: state rank");
171 }
172 if k == 0 || d == 0 {
173 bail!("k3 kda: D and conv_kernel must be > 0");
174 }
175
176 let mut hold = Vec::new();
177 let run = (|| {
178 let dx = up(gpu, x_qkv, &mut hold)?;
179 let dw = up(gpu, conv_w, &mut hold)?;
180 let dgate = up(gpu, gate, &mut hold)?;
181 let dbeta = up(gpu, beta, &mut hold)?;
182 let (dconv, drec, pull_state) = if let Some(d) = device {
183 (d.conv, d.recurrent, false)
184 } else {
185 let s = host.as_ref().context("k3 kda: host or device state")?;
186 (
187 up(gpu, &s.conv, &mut hold)?,
188 up(gpu, &s.recurrent, &mut hold)?,
189 true,
190 )
191 };
192 let dy = gpu.alloc((c * 4).max(1))?;
193 hold.push(dy);
194 let dout = gpu.alloc((cfg.qkv_dim() * 4).max(1))?;
195 hold.push(dout);
196
197 KernelLaunch::new(gpu, kernels.conv)
198 .grid([div_ceil(c as u32, CONV_BLOCK), 1, 1])
199 .block([CONV_BLOCK, 1, 1])
200 .arg_ptr(dx)
201 .arg_ptr(dw)
202 .arg_ptr(dconv)
203 .arg_ptr(dy)
204 .arg_u32(c as u32)
205 .arg_u32(k as u32)
206 .launch(stream)
207 .context("k3_kda_conv_update_f32")?;
208
209 let rec_block = (d as u32).min(128);
210 KernelLaunch::new(gpu, kernels.recurrent)
211 .grid([h as u32, 1, 1])
212 .block([rec_block, 1, 1])
213 .shared_mem((3 * d * 4) as u32)
214 .arg_ptr(dy)
215 .arg_ptr(dgate)
216 .arg_ptr(dbeta)
217 .arg_ptr(drec)
218 .arg_ptr(dout)
219 .arg_u32(h as u32)
220 .arg_u32(d as u32)
221 .arg_f32(KDA_L2_EPS)
222 .launch(stream)
223 .context("k3_kda_recurrent_step_f32")?;
224 gpu.synchronize(stream)?;
225
226 let mut out_b = vec![0u8; cfg.qkv_dim() * 4];
227 gpu.copy_d2h(dout, &mut out_b)?;
228 if pull_state {
229 let state = host.as_mut().context("k3 kda: host state for D2H")?;
230 let mut conv_b = vec![0u8; cfg.conv_elems() * 4];
231 let mut rec_b = vec![0u8; cfg.recurrent_elems() * 4];
232 gpu.copy_d2h(dconv, &mut conv_b)?;
233 gpu.copy_d2h(drec, &mut rec_b)?;
234 state.conv = bytes_f32(&conv_b);
235 state.recurrent = bytes_f32(&rec_b);
236 }
237 Ok(bytes_f32(&out_b))
238 })();
239 for p in hold {
240 let _ = gpu.free(p);
241 }
242 run
243}
244
245#[cfg(test)]
246mod tests {
247 use super::*;
248 use atlas_core::kimi_k3::kda_decode_token;
249 use spark_runtime::gpu::mock::{MockArg, MockGpuBackend};
250
251 #[test]
252 fn resolve_looks_up_k3_entries() {
253 let gpu = MockGpuBackend::new();
254 let _ = K3KdaDecodeKernels::resolve(&gpu).unwrap();
255 assert_eq!(
256 gpu.kernel_lookups_snapshot(),
257 vec![
258 (MODULE.to_string(), CONV_ENTRY.to_string()),
259 (MODULE.to_string(), RECURRENT_ENTRY.to_string()),
260 ]
261 );
262 }
263
264 #[test]
265 fn mock_launch_contract_twin_geometry() {
266 let gpu = MockGpuBackend::new();
267 let k = K3KdaDecodeKernels::resolve(&gpu).unwrap();
268 let cfg = KdaConfig::twin_0_40b();
269 let mut state = KdaState::new(&cfg);
270 let x = vec![0.1f32; cfg.conv_dim()];
271 let w = vec![0.2f32; cfg.conv_elems()];
272 let gate = vec![-0.5f32; cfg.qkv_dim()];
273 let beta = vec![0.25f32; cfg.heads];
274 let _ = launch_k3_kda_decode_token(&gpu, &k, &x, &w, &gate, &beta, &cfg, &mut state, 3)
275 .unwrap();
276 let launches = gpu.launches_snapshot();
277 assert_eq!(launches.len(), 2, "conv then recurrent");
278 let c = cfg.conv_dim() as u32;
279 assert_eq!(launches[0].grid, [div_ceil(c, CONV_BLOCK), 1, 1]);
280 assert_eq!(launches[0].block, [CONV_BLOCK, 1, 1]);
281 assert_eq!(launches[0].shared_mem, 0);
282 assert_eq!(launches[0].stream, 3);
283 assert_eq!(launches[0].args.len(), 6);
284 assert_eq!(
285 launches[0].args[4],
286 MockArg::Bytes(c.to_le_bytes().to_vec())
287 );
288 assert_eq!(
289 launches[0].args[5],
290 MockArg::Bytes((cfg.conv_kernel as u32).to_le_bytes().to_vec())
291 );
292
293 let rec = &launches[1];
294 assert_eq!(rec.grid, [cfg.heads as u32, 1, 1]);
295 assert_eq!(rec.block, [cfg.head_dim as u32, 1, 1]);
296 assert_eq!(rec.shared_mem, (3 * cfg.head_dim * 4) as u32);
297 assert_eq!(rec.args.len(), 8);
298 assert_eq!(
299 rec.args[7],
300 MockArg::Bytes(KDA_L2_EPS.to_le_bytes().to_vec())
301 );
302 }
303
304 #[test]
305 fn cpu_oracle_is_the_compare_target() {
306 let cfg = KdaConfig::twin_0_40b();
307 let mut state = KdaState::new(&cfg);
308 let x = vec![0.1f32; cfg.conv_dim()];
309 let w = vec![0.2f32; cfg.conv_elems()];
310 let gate = vec![-0.5f32; cfg.qkv_dim()];
311 let beta = vec![0.25f32; cfg.heads];
312 let y = kda_decode_token(&x, &w, &gate, &beta, &cfg, &mut state);
313 assert_eq!(y.len(), cfg.qkv_dim());
314 assert!(y.iter().any(|v| v.abs() > 1e-8));
315 }
316
317 #[test]
318 fn device_resident_state_skips_conv_recurrent_d2h() {
319 let gpu = MockGpuBackend::new();
320 let k = K3KdaDecodeKernels::resolve(&gpu).unwrap();
321 let cfg = KdaConfig::twin_0_40b();
322 let host = KdaState::new(&cfg);
323 let x = vec![0.1f32; cfg.conv_dim()];
324 let w = vec![0.2f32; cfg.conv_elems()];
325 let gate = vec![-0.5f32; cfg.qkv_dim()];
326 let beta = vec![0.25f32; cfg.heads];
327 let device = KdaDeviceState::alloc_and_upload(&gpu, &host).unwrap();
328 let before = gpu.d2h_blocking_count();
329 let _ =
330 launch_k3_kda_decode_token_on_device(&gpu, &k, &x, &w, &gate, &beta, &cfg, &device, 0)
331 .unwrap();
332 let pulled = gpu.d2h_blocking_count() - before;
333 assert_eq!(
334 pulled, 1,
335 "resident decode must D2H the output only, not conv/recurrent (got {pulled})"
336 );
337 device.free(&gpu).unwrap();
338 }
339}