spark_model/kimi_k3/
mla_cuda.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Host launch for K3 CUDA gated-NoPE MLA decode (`mla_decode` PTX module).
4//!
5//! Two kernels, one token: maybe_rope (NoPE skips rotate), then SDPA + gate.
6//! CPU oracle: [`atlas_core::kimi_k3::mla_decode_token`]. BoundLayer serve
7//! FullAttention default is this launch (`K3_CUDA_MLA=0` keeps CPU).
8
9use anyhow::{Context, Result, bail};
10use atlas_core::kimi_k3::{MlaConfig, MlaKv};
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/mla_decode.cu`.
15pub const MODULE: &str = "mla_decode";
16pub const ROPE_ENTRY: &str = "k3_mla_maybe_rope_f32";
17pub const SDPA_ENTRY: &str = "k3_mla_sdpa_gate_f32";
18const ROPE_BLOCK: u32 = 128;
19const SDPA_BLOCK: u32 = 32;
20
21#[derive(Clone, Copy, Debug)]
22pub struct K3MlaDecodeKernels {
23    pub rope: KernelHandle,
24    pub sdpa: KernelHandle,
25}
26
27impl K3MlaDecodeKernels {
28    pub fn resolve(gpu: &dyn GpuBackend) -> Result<Self> {
29        Ok(Self {
30            rope: gpu.kernel(MODULE, ROPE_ENTRY)?,
31            sdpa: gpu.kernel(MODULE, SDPA_ENTRY)?,
32        })
33    }
34}
35
36fn f32_bytes(v: &[f32]) -> Vec<u8> {
37    v.iter().flat_map(|x| x.to_le_bytes()).collect()
38}
39
40fn bytes_f32(b: &[u8]) -> Vec<f32> {
41    b.chunks_exact(4)
42        .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
43        .collect()
44}
45
46fn up(gpu: &dyn GpuBackend, v: &[f32], hold: &mut Vec<DevicePtr>) -> Result<DevicePtr> {
47    let b = f32_bytes(v);
48    let p = gpu.alloc(b.len().max(1))?;
49    hold.push(p);
50    gpu.copy_h2d(&b, p)?;
51    Ok(p)
52}
53
54/// Device-resident MLA KV. Append is one-row D2D/H2D; SDPA reads the buffer.
55/// Host `MlaKv` remains the CPU oracle. Do not re-upload `[0..T]` each token.
56pub struct MlaDeviceKv {
57    pub k: DevicePtr,
58    pub v: DevicePtr,
59    pub seq_len: usize,
60    pub cap: usize,
61    k_row: usize,
62    v_row: usize,
63}
64
65impl MlaDeviceKv {
66    pub fn alloc(gpu: &dyn GpuBackend, cap: usize, k_row: usize, v_row: usize) -> Result<Self> {
67        let cap = cap.max(1);
68        Ok(Self {
69            k: gpu.alloc((cap * k_row * 4).max(1))?,
70            v: gpu.alloc((cap * v_row * 4).max(1))?,
71            seq_len: 0,
72            cap,
73            k_row,
74            v_row,
75        })
76    }
77
78    fn append_from_device(
79        &mut self,
80        gpu: &dyn GpuBackend,
81        k_new: DevicePtr,
82        v_host: &[f32],
83    ) -> Result<()> {
84        if self.seq_len >= self.cap {
85            bail!("k3 mla: device KV cap {} full", self.cap);
86        }
87        let k_off = DevicePtr(self.k.0 + (self.seq_len * self.k_row * 4) as u64);
88        gpu.copy_d2d(k_new, k_off, self.k_row * 4)?;
89        let v_off = DevicePtr(self.v.0 + (self.seq_len * self.v_row * 4) as u64);
90        gpu.copy_h2d(&f32_bytes(v_host), v_off)?;
91        self.seq_len += 1;
92        Ok(())
93    }
94}
95
96/// One-token CUDA gated-NoPE MLA. Updates `q`/`k` (rope) and `kv` (append).
97#[allow(clippy::too_many_arguments)]
98pub fn launch_k3_mla_decode_token(
99    gpu: &dyn GpuBackend,
100    kernels: &K3MlaDecodeKernels,
101    q: &mut [f32],
102    k: &mut [f32],
103    v: &[f32],
104    g: &[f32],
105    kv: &mut MlaKv,
106    cfg: &MlaConfig,
107    pos: usize,
108    theta: f32,
109    stream: u64,
110) -> Result<Vec<f32>> {
111    let (h, dq, dv) = (cfg.heads, cfg.qk_head_dim(), cfg.v_head_dim);
112    if q.len() != h * dq || k.len() != h * dq {
113        bail!("k3 mla: q/k rank");
114    }
115    if v.len() != h * dv || g.len() != h * dv {
116        bail!("k3 mla: v/g rank");
117    }
118    if dq == 0 {
119        bail!("k3 mla: dq must be > 0");
120    }
121
122    let mut hold = Vec::new();
123    let run = (|| {
124        let dq_ptr = up(gpu, q, &mut hold)?;
125        let dk_new = up(gpu, k, &mut hold)?;
126        KernelLaunch::new(gpu, kernels.rope)
127            .grid([div_ceil(h as u32, ROPE_BLOCK), 1, 1])
128            .block([ROPE_BLOCK, 1, 1])
129            .arg_ptr(dq_ptr)
130            .arg_ptr(dk_new)
131            .arg_u32(h as u32)
132            .arg_u32(cfg.qk_nope_head_dim as u32)
133            .arg_u32(cfg.qk_rope_head_dim as u32)
134            .arg_u32(pos as u32)
135            .arg_f32(theta)
136            .arg_u32(u32::from(cfg.mla_use_nope))
137            .launch(stream)
138            .context("k3_mla_maybe_rope_f32")?;
139        gpu.synchronize(stream)?;
140
141        let mut qb = vec![0u8; q.len() * 4];
142        let mut kb = vec![0u8; k.len() * 4];
143        gpu.copy_d2h(dq_ptr, &mut qb)?;
144        gpu.copy_d2h(dk_new, &mut kb)?;
145        q.copy_from_slice(&bytes_f32(&qb));
146        k.copy_from_slice(&bytes_f32(&kb));
147        kv.append(k, v);
148
149        let dq2 = up(gpu, q, &mut hold)?;
150        let dk = up(gpu, &kv.k, &mut hold)?;
151        let dv_ptr = up(gpu, &kv.v, &mut hold)?;
152        let dg = up(gpu, g, &mut hold)?;
153        let dout = gpu.alloc((h * dv * 4).max(1))?;
154        hold.push(dout);
155        KernelLaunch::new(gpu, kernels.sdpa)
156            .grid([h as u32, 1, 1])
157            .block([SDPA_BLOCK, 1, 1])
158            .arg_ptr(dq2)
159            .arg_ptr(dk)
160            .arg_ptr(dv_ptr)
161            .arg_ptr(dg)
162            .arg_ptr(dout)
163            .arg_u32(kv.seq_len as u32)
164            .arg_u32(h as u32)
165            .arg_u32(dq as u32)
166            .arg_u32(dv as u32)
167            .arg_u32(u32::from(cfg.mla_use_output_gate))
168            .launch(stream)
169            .context("k3_mla_sdpa_gate_f32")?;
170        gpu.synchronize(stream)?;
171
172        let mut out_b = vec![0u8; h * dv * 4];
173        gpu.copy_d2h(dout, &mut out_b)?;
174        Ok(bytes_f32(&out_b))
175    })();
176    for p in hold {
177        let _ = gpu.free(p);
178    }
179    run
180}
181
182/// Device-resident KV: rope stays on device, append is one row, one D2H (output).
183#[allow(clippy::too_many_arguments)]
184pub fn launch_k3_mla_decode_token_on_device(
185    gpu: &dyn GpuBackend,
186    kernels: &K3MlaDecodeKernels,
187    q: &[f32],
188    k: &[f32],
189    v: &[f32],
190    g: &[f32],
191    kv: &mut MlaDeviceKv,
192    cfg: &MlaConfig,
193    pos: usize,
194    theta: f32,
195    stream: u64,
196) -> Result<Vec<f32>> {
197    let (h, dq, dv) = (cfg.heads, cfg.qk_head_dim(), cfg.v_head_dim);
198    if q.len() != h * dq || k.len() != h * dq || v.len() != h * dv || g.len() != h * dv {
199        bail!("k3 mla: q/k/v/g rank");
200    }
201    let mut hold = Vec::new();
202    let run = (|| {
203        let dq_ptr = up(gpu, q, &mut hold)?;
204        let dk_new = up(gpu, k, &mut hold)?;
205        KernelLaunch::new(gpu, kernels.rope)
206            .grid([div_ceil(h as u32, ROPE_BLOCK), 1, 1])
207            .block([ROPE_BLOCK, 1, 1])
208            .arg_ptr(dq_ptr)
209            .arg_ptr(dk_new)
210            .arg_u32(h as u32)
211            .arg_u32(cfg.qk_nope_head_dim as u32)
212            .arg_u32(cfg.qk_rope_head_dim as u32)
213            .arg_u32(pos as u32)
214            .arg_f32(theta)
215            .arg_u32(u32::from(cfg.mla_use_nope))
216            .launch(stream)
217            .context("k3_mla_maybe_rope_f32")?;
218        kv.append_from_device(gpu, dk_new, v)?;
219        let dg = up(gpu, g, &mut hold)?;
220        let dout = gpu.alloc((h * dv * 4).max(1))?;
221        hold.push(dout);
222        KernelLaunch::new(gpu, kernels.sdpa)
223            .grid([h as u32, 1, 1])
224            .block([SDPA_BLOCK, 1, 1])
225            .arg_ptr(dq_ptr)
226            .arg_ptr(kv.k)
227            .arg_ptr(kv.v)
228            .arg_ptr(dg)
229            .arg_ptr(dout)
230            .arg_u32(kv.seq_len as u32)
231            .arg_u32(h as u32)
232            .arg_u32(dq as u32)
233            .arg_u32(dv as u32)
234            .arg_u32(u32::from(cfg.mla_use_output_gate))
235            .launch(stream)
236            .context("k3_mla_sdpa_gate_f32")?;
237        gpu.synchronize(stream)?;
238        let mut out_b = vec![0u8; h * dv * 4];
239        gpu.copy_d2h(dout, &mut out_b)?;
240        Ok(bytes_f32(&out_b))
241    })();
242    for p in hold {
243        let _ = gpu.free(p);
244    }
245    run
246}
247
248#[cfg(test)]
249mod tests {
250    use super::*;
251    use spark_runtime::gpu::mock::{MockArg, MockGpuBackend};
252
253    #[test]
254    fn resolve_looks_up_k3_entries() {
255        let gpu = MockGpuBackend::new();
256        let _ = K3MlaDecodeKernels::resolve(&gpu).unwrap();
257        assert_eq!(
258            gpu.kernel_lookups_snapshot(),
259            vec![
260                (MODULE.to_string(), ROPE_ENTRY.to_string()),
261                (MODULE.to_string(), SDPA_ENTRY.to_string()),
262            ]
263        );
264    }
265
266    #[test]
267    fn mock_launch_contract_twin_geometry() {
268        let gpu = MockGpuBackend::new();
269        let kernels = K3MlaDecodeKernels::resolve(&gpu).unwrap();
270        let cfg = MlaConfig::twin_0_40b();
271        let mut q = vec![0.1f32; cfg.heads * cfg.qk_head_dim()];
272        let mut k = vec![0.2f32; cfg.heads * cfg.qk_head_dim()];
273        let v = vec![0.3f32; cfg.heads * cfg.v_head_dim];
274        let g = vec![0.0f32; cfg.heads * cfg.v_head_dim];
275        let mut kv = MlaKv::default();
276        let _ = launch_k3_mla_decode_token(
277            &gpu, &kernels, &mut q, &mut k, &v, &g, &mut kv, &cfg, 3, 10000.0, 3,
278        )
279        .unwrap();
280        assert_eq!(kv.seq_len, 1);
281        let launches = gpu.launches_snapshot();
282        assert_eq!(launches.len(), 2, "rope then sdpa_gate");
283        assert_eq!(
284            launches[0].grid,
285            [div_ceil(cfg.heads as u32, ROPE_BLOCK), 1, 1]
286        );
287        assert_eq!(launches[0].block, [ROPE_BLOCK, 1, 1]);
288        assert_eq!(launches[0].args.len(), 8);
289        assert_eq!(
290            launches[0].args[7],
291            MockArg::Bytes(1u32.to_le_bytes().to_vec()),
292            "twin is NoPE"
293        );
294        let sdpa = &launches[1];
295        assert_eq!(sdpa.grid, [cfg.heads as u32, 1, 1]);
296        assert_eq!(sdpa.block, [SDPA_BLOCK, 1, 1]);
297        assert_eq!(sdpa.shared_mem, 0);
298        assert_eq!(sdpa.stream, 3);
299        assert_eq!(sdpa.args.len(), 10);
300        assert_eq!(
301            sdpa.args[9],
302            MockArg::Bytes(1u32.to_le_bytes().to_vec()),
303            "twin uses output gate"
304        );
305    }
306
307    #[test]
308    fn device_kv_second_token_d2h_is_output_only() {
309        let gpu = MockGpuBackend::new();
310        let kernels = K3MlaDecodeKernels::resolve(&gpu).unwrap();
311        let cfg = MlaConfig::twin_0_40b();
312        let q = vec![0.1f32; cfg.heads * cfg.qk_head_dim()];
313        let k = vec![0.2f32; cfg.heads * cfg.qk_head_dim()];
314        let v = vec![0.3f32; cfg.heads * cfg.v_head_dim];
315        let g = vec![0.0f32; cfg.heads * cfg.v_head_dim];
316        let mut kv = MlaDeviceKv::alloc(
317            &gpu,
318            8,
319            cfg.heads * cfg.qk_head_dim(),
320            cfg.heads * cfg.v_head_dim,
321        )
322        .unwrap();
323        let _ = launch_k3_mla_decode_token_on_device(
324            &gpu, &kernels, &q, &k, &v, &g, &mut kv, &cfg, 0, 10000.0, 0,
325        )
326        .unwrap();
327        let before = gpu.d2h_blocking_count();
328        let _ = launch_k3_mla_decode_token_on_device(
329            &gpu, &kernels, &q, &k, &v, &g, &mut kv, &cfg, 1, 10000.0, 0,
330        )
331        .unwrap();
332        let pulled = gpu.d2h_blocking_count() - before;
333        assert_eq!(
334            pulled, 1,
335            "resident MLA must not re-download KV; D2H is output only (got {pulled})"
336        );
337        assert_eq!(kv.seq_len, 2);
338    }
339}