1use std::collections::HashMap;
9
10use anyhow::{Context, Result, bail, ensure};
11use atlas_core::kimi_k3::{LatentMoeConfig, situ_glu_vec};
12use half::bf16;
13use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
14
15use crate::layers::ops::moe_w4a16_grouped_gemm_ptrtable;
16use crate::weight_map::QuantizedWeight;
17
18pub const MODULE: &str = "moe_w4a16";
20pub const PTRTABLE_E8M0: &str = "moe_w4a16_grouped_gemm_ptrtable_e8m0";
21pub const E8M0_ENTRY: &str = PTRTABLE_E8M0;
22
23#[derive(Clone, Copy, Debug)]
24pub struct K3MoeGemmKernels {
25 pub ptrtable: KernelHandle,
26}
27
28impl K3MoeGemmKernels {
29 pub fn resolve(gpu: &dyn GpuBackend) -> Result<Self> {
30 Ok(Self {
31 ptrtable: gpu.kernel(MODULE, PTRTABLE_E8M0).with_context(|| {
32 format!(
33 "K3 MXFP4: {MODULE}::{PTRTABLE_E8M0} missing; \
34 packed experts cannot silently run host F32"
35 )
36 })?,
37 })
38 }
39}
40
41#[allow(clippy::too_many_arguments)]
45pub fn launch_k3_moe_e8m0_ptrtable(
46 gpu: &dyn GpuBackend,
47 kernel: KernelHandle,
48 a: DevicePtr,
49 packed: &[QuantizedWeight],
50 c: DevicePtr,
51 expert_offsets: DevicePtr,
52 sorted_token_ids: DevicePtr,
53 num_experts: u32,
54 n_out: u32,
55 k: u32,
56 stream: u64,
57) -> Result<()> {
58 ensure!(
59 packed.len() == num_experts as usize && num_experts > 0,
60 "K3 MXFP4: ptrtable length {} != num_experts {num_experts}",
61 packed.len()
62 );
63 let mut hold = Vec::new();
64 let run = (|| {
65 let (packed_ptrs, scale_ptrs, scale2_vals) = upload_ptr_table(gpu, packed, &mut hold)?;
66 moe_w4a16_grouped_gemm_ptrtable(
67 gpu,
68 kernel,
69 a,
70 packed_ptrs,
71 scale_ptrs,
72 scale2_vals,
73 c,
74 expert_offsets,
75 sorted_token_ids,
76 num_experts,
77 n_out,
78 k,
79 1, stream,
81 )
82 .context("moe_w4a16_grouped_gemm_ptrtable_e8m0")
83 })();
84 for p in hold {
85 let _ = gpu.free(p);
86 }
87 run
88}
89
90#[allow(clippy::too_many_arguments)]
92pub fn launch_k3_latent_moe_experts(
93 gpu: &dyn GpuBackend,
94 kernels: &K3MoeGemmKernels,
95 packed: &[(String, QuantizedWeight)],
96 latent: &[f32],
97 ids: &[usize],
98 mix_w: &[f32],
99 cfg: &LatentMoeConfig,
100 stream: u64,
101) -> Result<Vec<f32>> {
102 ensure!(
103 latent.len() == cfg.latent,
104 "K3 MXFP4: latent {} != {}",
105 latent.len(),
106 cfg.latent
107 );
108 ensure!(ids.len() == mix_w.len(), "K3 MXFP4: ids/weights rank");
109 if ids.is_empty() {
110 return Ok(vec![0.0; cfg.latent]);
111 }
112 let table = index_packed(packed)?;
113 let w1 = gather_proj(&table, ids, 0, "w1")?;
114 let w2 = gather_proj(&table, ids, 1, "w2")?;
115 let w3 = gather_proj(&table, ids, 2, "w3")?;
116 let m = ids.len();
117 let a_w1: Vec<f32> = latent
118 .iter()
119 .copied()
120 .cycle()
121 .take(m * cfg.latent)
122 .collect();
123 let gate = gemm_rows(
124 gpu,
125 kernels,
126 &a_w1,
127 m,
128 &w1,
129 cfg.expert_hidden as u32,
130 cfg.latent as u32,
131 stream,
132 )?;
133 let up = gemm_rows(
134 gpu,
135 kernels,
136 &a_w1,
137 m,
138 &w3,
139 cfg.expert_hidden as u32,
140 cfg.latent as u32,
141 stream,
142 )?;
143 let mut mid = Vec::with_capacity(m * cfg.expert_hidden);
144 for e in 0..m {
145 let g = &gate[e * cfg.expert_hidden..(e + 1) * cfg.expert_hidden];
146 let u = &up[e * cfg.expert_hidden..(e + 1) * cfg.expert_hidden];
147 mid.extend(situ_glu_vec(g, u, cfg.situ_beta, cfg.situ_linear_beta));
148 }
149 let down_rows = gemm_rows(
150 gpu,
151 kernels,
152 &mid,
153 m,
154 &w2,
155 cfg.latent as u32,
156 cfg.expert_hidden as u32,
157 stream,
158 )?;
159 let mut mixed = vec![0.0f32; cfg.latent];
160 for (e, &w) in mix_w.iter().enumerate() {
161 let y = &down_rows[e * cfg.latent..(e + 1) * cfg.latent];
162 for (acc, yy) in mixed.iter_mut().zip(y) {
163 *acc += w * *yy;
164 }
165 }
166 Ok(mixed)
167}
168
169fn gemm_rows(
170 gpu: &dyn GpuBackend,
171 kernels: &K3MoeGemmKernels,
172 a_f32: &[f32],
173 m: usize,
174 packed: &[QuantizedWeight],
175 n_out: u32,
176 k: u32,
177 stream: u64,
178) -> Result<Vec<f32>> {
179 ensure!(
180 a_f32.len() == m * k as usize,
181 "K3 MXFP4: A {} vs {m}x{k}",
182 a_f32.len()
183 );
184 let mut hold = Vec::new();
185 let run = (|| {
186 let a = up_bf16(gpu, a_f32, &mut hold)?;
187 let c = gpu.alloc((m * n_out as usize * 2).max(1))?;
188 hold.push(c);
189 let off: Vec<i32> = (0..=m as i32).collect();
190 let ids: Vec<i32> = (0..m as i32).collect();
191 let offsets = up_i32(gpu, &off, &mut hold)?;
192 let sorted = up_i32(gpu, &ids, &mut hold)?;
193 launch_k3_moe_e8m0_ptrtable(
194 gpu,
195 kernels.ptrtable,
196 a,
197 packed,
198 c,
199 offsets,
200 sorted,
201 m as u32,
202 n_out,
203 k,
204 stream,
205 )?;
206 gpu.synchronize(stream)?;
207 let mut raw = vec![0u8; m * n_out as usize * 2];
208 gpu.copy_d2h(c, &mut raw)?;
209 Ok(bf16_to_f32(&raw))
210 })();
211 for p in hold {
212 let _ = gpu.free(p);
213 }
214 run
215}
216
217fn upload_ptr_table(
218 gpu: &dyn GpuBackend,
219 packed: &[QuantizedWeight],
220 hold: &mut Vec<DevicePtr>,
221) -> Result<(DevicePtr, DevicePtr, DevicePtr)> {
222 let n = packed.len();
223 let packed_bytes: Vec<u8> = packed
224 .iter()
225 .flat_map(|w| w.weight.0.to_le_bytes())
226 .collect();
227 let scale_bytes: Vec<u8> = packed
228 .iter()
229 .flat_map(|w| w.weight_scale.0.to_le_bytes())
230 .collect();
231 let scale2_bytes: Vec<u8> = packed
232 .iter()
233 .flat_map(|w| w.weight_scale_2.to_le_bytes())
234 .collect();
235 let packed_ptrs = gpu.alloc((n * 8).max(1))?;
236 hold.push(packed_ptrs);
237 gpu.copy_h2d(&packed_bytes, packed_ptrs)?;
238 let scale_ptrs = gpu.alloc((n * 8).max(1))?;
239 hold.push(scale_ptrs);
240 gpu.copy_h2d(&scale_bytes, scale_ptrs)?;
241 let scale2_vals = gpu.alloc((n * 4).max(1))?;
242 hold.push(scale2_vals);
243 gpu.copy_h2d(&scale2_bytes, scale2_vals)?;
244 Ok((packed_ptrs, scale_ptrs, scale2_vals))
245}
246
247fn index_packed(
248 packed: &[(String, QuantizedWeight)],
249) -> Result<HashMap<(usize, usize), QuantizedWeight>> {
250 let mut t = HashMap::new();
251 for (prefix, w) in packed {
252 let (id, proj) = parse_expert_proj(prefix)?;
253 if t.insert((id, proj), *w).is_some() {
254 bail!("K3 MXFP4: duplicate packed {prefix}");
255 }
256 }
257 Ok(t)
258}
259
260fn parse_expert_proj(prefix: &str) -> Result<(usize, usize)> {
261 let (head, proj) = prefix
262 .rsplit_once('.')
263 .with_context(|| format!("K3 MXFP4: expert prefix {prefix}"))?;
264 let proj_i = match proj {
265 "w1" => 0,
266 "w2" => 1,
267 "w3" => 2,
268 _ => bail!("K3 MXFP4: expected w1|w2|w3 in {prefix}"),
269 };
270 let id_s = head
271 .rsplit_once(".experts.")
272 .map(|(_, id)| id)
273 .with_context(|| format!("K3 MXFP4: experts.id in {prefix}"))?;
274 let id: usize = id_s
275 .parse()
276 .with_context(|| format!("K3 MXFP4: expert id {id_s}"))?;
277 Ok((id, proj_i))
278}
279
280fn gather_proj(
281 table: &HashMap<(usize, usize), QuantizedWeight>,
282 ids: &[usize],
283 proj: usize,
284 name: &str,
285) -> Result<Vec<QuantizedWeight>> {
286 ids.iter()
287 .map(|&id| {
288 table.get(&(id, proj)).copied().with_context(|| {
289 format!(
290 "K3 MXFP4: expert {id} {name} missing; packed experts cannot silently run host F32"
291 )
292 })
293 })
294 .collect()
295}
296
297fn up_bf16(gpu: &dyn GpuBackend, v: &[f32], hold: &mut Vec<DevicePtr>) -> Result<DevicePtr> {
298 let b: Vec<u8> = v
299 .iter()
300 .flat_map(|&f| bf16::from_f32(f).to_le_bytes())
301 .collect();
302 let p = gpu.alloc(b.len().max(1))?;
303 hold.push(p);
304 gpu.copy_h2d(&b, p)?;
305 Ok(p)
306}
307
308fn up_i32(gpu: &dyn GpuBackend, v: &[i32], hold: &mut Vec<DevicePtr>) -> Result<DevicePtr> {
309 let b: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
310 let p = gpu.alloc(b.len().max(1))?;
311 hold.push(p);
312 gpu.copy_h2d(&b, p)?;
313 Ok(p)
314}
315
316fn bf16_to_f32(raw: &[u8]) -> Vec<f32> {
317 raw.chunks_exact(2)
318 .map(|b| bf16::from_le_bytes([b[0], b[1]]).to_f32())
319 .collect()
320}
321
322#[cfg(test)]
323mod tests {
324 use super::*;
325 use spark_runtime::gpu::mock::{MockArg, MockGpuBackend};
326 use spark_runtime::kernel_args::div_ceil;
327
328 fn dummy_qw(gpu: &MockGpuBackend) -> QuantizedWeight {
329 QuantizedWeight {
330 weight: gpu.alloc(16).unwrap(),
331 weight_scale: gpu.alloc(1).unwrap(),
332 weight_scale_2: 1.0,
333 input_scale: DevicePtr::NULL,
334 weight_scale_2_vec: DevicePtr::NULL,
335 }
336 }
337
338 #[test]
339 fn resolve_looks_up_e8m0_ptrtable() {
340 let gpu = MockGpuBackend::new();
341 let _ = K3MoeGemmKernels::resolve(&gpu).unwrap();
342 assert_eq!(
343 gpu.kernel_lookups_snapshot(),
344 vec![(MODULE.to_string(), PTRTABLE_E8M0.to_string())]
345 );
346 }
347
348 #[test]
349 fn deny_kernel_resolve_bails_not_silent_cpu() {
350 let gpu = MockGpuBackend::new();
351 gpu.deny_kernel(MODULE, PTRTABLE_E8M0);
352 let err = K3MoeGemmKernels::resolve(&gpu).unwrap_err().to_string();
353 assert!(
354 err.contains(PTRTABLE_E8M0) && err.contains("cannot silently run host F32"),
355 "{err}"
356 );
357 assert_eq!(gpu.launch_count(), 0);
358 }
359
360 #[test]
361 fn mock_launch_contract_one_expert() {
362 let gpu = MockGpuBackend::new();
363 let k = K3MoeGemmKernels::resolve(&gpu).unwrap();
364 let packed = dummy_qw(&gpu);
365 let a = gpu.alloc(64).unwrap();
366 let c = gpu.alloc(128).unwrap();
367 let off = gpu.alloc(8).unwrap();
368 let ids = gpu.alloc(4).unwrap();
369 let n_out = 64u32;
370 let kk = 32u32;
371 launch_k3_moe_e8m0_ptrtable(&gpu, k.ptrtable, a, &[packed], c, off, ids, 1, n_out, kk, 3)
372 .unwrap();
373 let launches = gpu.launches_snapshot();
374 assert_eq!(launches.len(), 1);
375 assert_eq!(launches[0].grid, [div_ceil(n_out, 64), 1, 1]);
376 assert_eq!(launches[0].block, [128, 1, 1]);
377 assert_eq!(launches[0].stream, 3);
378 assert_eq!(launches[0].args.len(), 10);
379 assert_eq!(
380 launches[0].args[7],
381 MockArg::Bytes(1u32.to_le_bytes().to_vec())
382 );
383 assert_eq!(
384 launches[0].args[8],
385 MockArg::Bytes(n_out.to_le_bytes().to_vec())
386 );
387 assert_eq!(
388 launches[0].args[9],
389 MockArg::Bytes(kk.to_le_bytes().to_vec())
390 );
391 }
392
393 #[test]
394 fn parse_k3_expert_prefix() {
395 let p = "language_model.model.layers.12.block_sparse_moe.experts.7.w1";
396 assert_eq!(parse_expert_proj(p).unwrap(), (7, 0));
397 assert_eq!(
398 parse_expert_proj("model.layers.1.block_sparse_moe.experts.0.w3").unwrap(),
399 (0, 2)
400 );
401 }
402
403 #[test]
404 fn empty_ids_does_not_launch() {
405 let gpu = MockGpuBackend::new();
406 let k = K3MoeGemmKernels::resolve(&gpu).unwrap();
407 let cfg = LatentMoeConfig {
408 hidden: 2,
409 latent: 2,
410 expert_hidden: 2,
411 n_routed: 1,
412 top_k: 1,
413 n_shared: 0,
414 situ_beta: 4.0,
415 situ_linear_beta: 25.0,
416 use_norm: false,
417 renormalize: true,
418 };
419 let y =
420 launch_k3_latent_moe_experts(&gpu, &k, &[], &[1.0, 0.0], &[], &[], &cfg, 0).unwrap();
421 assert_eq!(y, vec![0.0, 0.0]);
422 assert_eq!(gpu.launch_count(), 0);
423 }
424
425 #[test]
426 fn missing_packed_proj_bails_not_cpu() {
427 let gpu = MockGpuBackend::new();
428 let k = K3MoeGemmKernels::resolve(&gpu).unwrap();
429 let cfg = LatentMoeConfig {
430 hidden: 2,
431 latent: 2,
432 expert_hidden: 2,
433 n_routed: 1,
434 top_k: 1,
435 n_shared: 0,
436 situ_beta: 4.0,
437 situ_linear_beta: 25.0,
438 use_norm: false,
439 renormalize: true,
440 };
441 let packed = [(
442 "model.layers.1.block_sparse_moe.experts.0.w1".to_string(),
443 dummy_qw(&gpu),
444 )];
445 let before = gpu.launch_count();
446 let err =
447 launch_k3_latent_moe_experts(&gpu, &k, &packed, &[1.0, 0.0], &[0], &[1.0], &cfg, 0)
448 .unwrap_err()
449 .to_string();
450 assert!(
451 err.contains("w2") && err.contains("cannot silently run host F32"),
452 "{err}"
453 );
454 assert_eq!(gpu.launch_count(), before);
455 }
456}