spark_model/kimi_k3/
moe_cuda.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Host launch for K3 packed LatentMoE experts (`moe_w4a16` E8M0 ptrtable).
4//!
5//! Required handle: [`PTRTABLE_E8M0`] from DSV4 extra_cu. Lookup-fail bails;
6//! packed tensors must not silently dequant on the host F32 twin path.
7
8use 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
18/// PTX module from kimi-k3 extra_cu of DSV4 `moe_w4a16_grouped_gemm.cu`.
19pub 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/// One grouped E8M0 GEMM: `C[M, N] = A[M, K] @ B_packed[N, K/2]` via ptrtable.
42///
43/// Each expert in `packed` owns one output row (`expert_offsets` 0..=num_experts).
44#[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, // decode: one token per selected expert
80            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/// Packed w1/w2/w3 SiTU mix. Same contract as [`atlas_core::kimi_k3::mix_routed_experts`].
91#[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}