spark_model/kimi_k3/
kda_cuda.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Host launch for K3 CUDA KDA decode (`kda_decode` PTX module).
4//!
5//! Two kernels, one token: conv-4 + SiLU, then L2 q/k + delta-rule.
6//! CPU oracle: [`atlas_core::kimi_k3::kda_decode_token`]. BoundLayer serve
7//! LinearAttention default is this launch (`K3_CUDA_KDA=0` keeps CPU).
8
9use 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
14/// PTX module stem = `kernels/gb10/kimi-k3/bf16/kda_decode.cu`.
15pub 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
53/// Device-resident conv + recurrent. Seed once; decode does not D2H them.
54///
55/// Host `KdaState` remains the CPU oracle. Certified tok/s must use this,
56/// not `launch_k3_kda_decode_token` (that path uploads ~6.3 MB/layer/token
57/// at production width).
58pub 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
91/// One-token CUDA KDA decode. Host-state round-trip (oracle / CPU escape).
92pub 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
117/// One-token CUDA KDA on device-resident conv/recurrent. Does not D2H state.
118pub 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}