1use 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
14pub 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
54pub 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#[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#[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}