spark_model/layers/ops/
moe_atomic_c4.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Atomic-add C=4 MoE decode experiment.
4
5#![allow(unused_imports)]
6
7use anyhow::Result;
8use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
9use spark_runtime::kernel_args::{KernelLaunch, div_ceil};
10
11use crate::weight_map::QuantizedWeight;
12
13#[allow(clippy::too_many_arguments)]
14pub fn moe_decode_atomic_c4_silu_down_accum(
15    gpu: &dyn GpuBackend,
16    kernel: KernelHandle,
17    gate_out: DevicePtr,
18    up_out: DevicePtr,
19    packed_ptrs: DevicePtr,
20    scale_ptrs: DevicePtr,
21    scale2_vals: DevicePtr,
22    expert_indices: DevicePtr,
23    expert_weights: DevicePtr,
24    routed_accum: DevicePtr,
25    sh_gate_in: DevicePtr,
26    sh_up_in: DevicePtr,
27    sh_down: &QuantizedWeight,
28    sh_down_out: DevicePtr,
29    hidden: u32,
30    inter: u32,
31    top_k: u32,
32    num_tokens: u32,
33    stream: u64,
34) -> Result<()> {
35    let smem_bytes = (inter as usize * std::mem::size_of::<f32>()) as u32;
36    KernelLaunch::new(gpu, kernel)
37        .grid([div_ceil(hidden, 8), num_tokens * (top_k + 1), 1])
38        .block([128, 1, 1])
39        .shared_mem(smem_bytes)
40        .arg_ptr(gate_out)
41        .arg_ptr(up_out)
42        .arg_ptr(packed_ptrs)
43        .arg_ptr(scale_ptrs)
44        .arg_ptr(scale2_vals)
45        .arg_ptr(expert_indices)
46        .arg_ptr(expert_weights)
47        .arg_ptr(routed_accum)
48        .arg_ptr(sh_gate_in)
49        .arg_ptr(sh_up_in)
50        .arg_ptr(sh_down.weight)
51        .arg_ptr(sh_down.weight_scale)
52        .arg_f32(sh_down.weight_scale_2)
53        .arg_ptr(sh_down_out)
54        .arg_u32(hidden)
55        .arg_u32(inter)
56        .arg_u32(top_k)
57        .arg_u32(num_tokens)
58        .launch(stream)
59}
60
61#[allow(clippy::too_many_arguments)]
62pub fn moe_decode_atomic_c4_finalize(
63    gpu: &dyn GpuBackend,
64    kernel: KernelHandle,
65    output: DevicePtr,
66    routed_accum: DevicePtr,
67    shared_out: DevicePtr,
68    input: DevicePtr,
69    gate_weight: DevicePtr,
70    hidden: u32,
71    num_tokens: u32,
72    include_shared: bool,
73    stream: u64,
74) -> Result<()> {
75    KernelLaunch::new(gpu, kernel)
76        .grid([div_ceil(hidden, 256), num_tokens, 1])
77        .block([256, 1, 1])
78        .arg_ptr(output)
79        .arg_ptr(routed_accum)
80        .arg_ptr(shared_out)
81        .arg_ptr(input)
82        .arg_ptr(gate_weight)
83        .arg_u32(hidden)
84        .arg_u32(num_tokens)
85        .arg_u32(if include_shared { 1 } else { 0 })
86        .launch(stream)
87}