spark_model/layers/
nemotron_moe.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Nemotron-H standalone MoE FFN layer.
4//!
5//! Supports two variants:
6//!   - **Nano 30B**: Direct MoE — experts operate on full hidden_size.
7//!   - **Super 120B**: LatentMoE — routed experts operate in latent space `[moe_latent_size]`,
8//!     with fc1/fc2 latent projections bridging hidden↔latent.
9//!
10//! Forward: RMS norm → gate → sigmoid topK routing → (fc1_latent if latent) →
11//!          batched up GEMV → fused relu²+down → weighted_sum → (fc2_latent if latent) →
12//!          shared expert up+relu²+down → sum routed+shared → residual add.
13//!
14//! All expert dispatch is device-side (pointer tables) — zero D2H sync.
15
16use anyhow::Result;
17use atlas_core::config::ModelConfig;
18use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
19use spark_runtime::kv_cache::PagedKvCache;
20
21use crate::layer::{EmptyLayerState, ForwardContext, LayerState, TransformerLayer};
22use crate::layers::ops;
23use crate::weight_map::{DenseWeight, NemotronMoeWeights, QuantizedWeight};
24
25/// Device-side pointer table for one projection across all experts.
26struct ExpertPtrTable {
27    packed_ptrs: DevicePtr,
28    scale_ptrs: DevicePtr,
29    scale2_vals: DevicePtr,
30}
31
32/// Nemotron-H standalone MoE FFN layer.
33pub struct NemotronMoeLayer {
34    weights: NemotronMoeWeights,
35    input_norm: DenseWeight,
36    /// LatentMoE dimension (0 = direct, >0 = latent).
37    moe_latent_size: usize,
38    /// Routed expert intermediate size for this layer (Puzzle: per-block).
39    moe_inter: usize,
40    /// Top-K experts activated per token for this layer (Puzzle: per-block).
41    top_k: usize,
42    // Kernel handles — decode (single token)
43    rms_norm_residual_k: KernelHandle,
44    dense_gemv_k: KernelHandle,
45    topk_sigmoid_k: KernelHandle,
46    moe_expert_gemv_k: KernelHandle,
47    w4a16_gemv_k: KernelHandle,
48    /// Single-warp `w4a16_gemv_sw`. `KernelHandle(0)` on miss → base GEMV.
49    w4a16_gemv_sw_k: KernelHandle,
50    /// Native-FP8 decode GEMV for the shared-expert up_proj (see
51    /// `NemotronMoeWeights::shared_up_fp8`). 0 when unavailable.
52    w8a16_gemv_k: KernelHandle,
53    /// Native-FP8 prefill GEMM for the shared expert. 0 when unavailable.
54    w8a16_gemm_k: KernelHandle,
55    w8a16_gemm_pipelined_k: KernelHandle,
56    relu2_down_shared_k: KernelHandle,
57    weighted_sum_scale_k: KernelHandle,
58    residual_add_k: KernelHandle,
59    // Kernel handles — prefill (batched GEMM)
60    dense_gemm_k: KernelHandle,
61    /// Pipelined tensor-core BF16 GEMM (mma.sync.m16n8k16 + cp.async 2-stage,
62    /// 128x128 tile). `dense_gemm_bf16` is a SCALAR 16x16 kernel — on the
63    /// large-M prefill shapes it is ~40x slower, and the three dense GEMMs of a
64    /// LatentMoE layer (gate, fc1_latent, fc2_latent) were the single largest
65    /// prefill cost on Puzzle (34% of all GPU time). Same math (cosine=1.0).
66    dense_gemm_pipelined_k: KernelHandle,
67    w4a16_gemm_k: KernelHandle,
68    // Batched N-token MoE prefill kernels
69    topk_sigmoid_batched_k: KernelHandle,
70    moe_up_prefill_k: KernelHandle,
71    moe_relu2_down_prefill_k: KernelHandle,
72    moe_weighted_sum_prefill_k: KernelHandle,
73    // Sorted grouped GEMM (Qwen pattern — proven to work)
74    moe_sort_k: KernelHandle,
75    moe_grouped_gemm_k: KernelHandle,
76    moe_relu2_elementwise_k: KernelHandle,
77    moe_grouped_gemm_relu2_k: KernelHandle,
78    moe_w4a4_grouped_k: KernelHandle,
79    moe_unpermute_reduce_k: KernelHandle,
80    moe_grouped_gemm_n128_k: KernelHandle,
81    up_ptrs: ExpertPtrTable,
82    down_ptrs: ExpertPtrTable,
83    // Transposed expert pointer tables (for N128 grouped GEMM)
84    up_ptrs_t: Option<ExpertPtrTable>,
85    down_ptrs_t: Option<ExpertPtrTable>,
86    // Transposed shared expert weights
87    shared_up_t: Option<QuantizedWeight>,
88    shared_down_t: Option<QuantizedWeight>,
89    // Pre-dequantized FP8 E4M3 [N, K] copies of the shared-expert projections.
90    // Consumed by `fp8_gemm_t_m128_mfast` (no dequant phase); see the SSM layer.
91    shared_up_pd_fp8: Option<DevicePtr>,
92    shared_down_pd_fp8: Option<DevicePtr>,
93    // FP8 E4M3 copies of the BF16 latent projections, so prefill runs the tuned
94    // FP8 GEMM instead of dense_gemm_bf16_pipelined (and halves their bytes).
95    fc1_pd_fp8: Option<DevicePtr>,
96    fc2_pd_fp8: Option<DevicePtr>,
97    // Transposed SSM GEMM kernel handle (for shared expert)
98    w4a16_gemm_t_k: KernelHandle,
99    w4a16_gemm_t_m128_k: KernelHandle,
100    fp8_gemm_m128_k: KernelHandle,
101    w4a4_gemm_k: KernelHandle,
102    quantize_nvfp4_k: KernelHandle,
103}
104
105impl NemotronMoeLayer {
106    pub fn new(
107        weights: NemotronMoeWeights,
108        input_norm: DenseWeight,
109        config: &ModelConfig,
110        gpu: &dyn GpuBackend,
111        moe_inter: usize,
112        top_k: usize,
113    ) -> Result<Self> {
114        let up_ptrs = build_ptr_table(&weights.experts, |e| &e.up_proj, gpu)?;
115        let down_ptrs = build_ptr_table(&weights.experts, |e| &e.down_proj, gpu)?;
116        let moe_inter = if moe_inter > 0 {
117            moe_inter
118        } else {
119            config.moe_intermediate_size
120        };
121        let top_k = if top_k > 0 {
122            top_k
123        } else {
124            config.num_experts_per_tok
125        };
126        // Nemotron's `top_k` is PER LAYER (`num_experts_per_tok_for`), so this
127        // runs once per MoE layer and a single outlying block config cannot
128        // slip through on the model-wide value. `MoeLayer::new` carries the
129        // same pair of bounds; `NemotronMoeLayer` had no check at all and its
130        // decode kernel is the one whose shadows were capped at 24.
131        let num_experts = weights.experts.len();
132        anyhow::ensure!(
133            top_k > 0
134                && top_k <= num_experts
135                && top_k <= crate::layers::ops::MOE_TOPK_SIGMOID_MAX_TOP_K
136                && num_experts <= crate::layers::ops::MOE_TOPK_SIGMOID_MAX_EXPERTS,
137            "Nemotron MoE config invalid: top_k={} must be in 1..={} and within \
138             the routing kernels' bounds (top_k max {}, num_experts={} max {})",
139            top_k,
140            num_experts,
141            crate::layers::ops::MOE_TOPK_SIGMOID_MAX_TOP_K,
142            num_experts,
143            crate::layers::ops::MOE_TOPK_SIGMOID_MAX_EXPERTS,
144        );
145
146        Ok(Self {
147            weights,
148            input_norm,
149            moe_latent_size: config.moe_latent_size,
150            moe_inter,
151            top_k,
152            rms_norm_residual_k: gpu.kernel("norm", "rms_norm_residual")?,
153            dense_gemv_k: gpu.kernel("gemv", "dense_gemv_bf16")?,
154            topk_sigmoid_k: gpu.kernel("moe_topk_sig", "moe_topk_sigmoid")?,
155            moe_expert_gemv_k: gpu.kernel("moe_expert_gemv", "moe_expert_gemv")?,
156            w4a16_gemv_k: gpu.kernel("w4a16_gemv", "w4a16_gemv")?,
157            w4a16_gemv_sw_k: super::try_kernel(gpu, "w4a16_gemv", "w4a16_gemv_sw"),
158            w8a16_gemv_k: super::try_kernel(gpu, "w8a16_gemv", "w8a16_gemv"),
159            w8a16_gemm_k: super::try_kernel(gpu, "w8a16_gemm", "w8a16_gemm"),
160            w8a16_gemm_pipelined_k: super::try_kernel(
161                gpu,
162                "w8a16_gemm_pipelined",
163                "w8a16_gemm_pipelined",
164            ),
165            relu2_down_shared_k: gpu.kernel("moe_relu2_fused", "moe_expert_relu2_down_shared")?,
166            weighted_sum_scale_k: gpu.kernel("relu2", "moe_weighted_sum_scale")?,
167            residual_add_k: gpu.kernel("residual_add", "bf16_residual_add")?,
168            dense_gemm_k: gpu.kernel("gemm", "dense_gemm_bf16")?,
169            dense_gemm_pipelined_k: super::try_kernel(gpu, "gemm", "dense_gemm_bf16_pipelined"),
170            w4a16_gemm_k: gpu.kernel("w4a16", "w4a16_gemm")?,
171            topk_sigmoid_batched_k: super::try_kernel(
172                gpu,
173                "nemotron_moe_prefill",
174                "nemotron_moe_topk_sigmoid_batched",
175            ),
176            moe_up_prefill_k: super::try_kernel(
177                gpu,
178                "nemotron_moe_prefill",
179                "nemotron_moe_up_prefill",
180            ),
181            moe_relu2_down_prefill_k: super::try_kernel(
182                gpu,
183                "nemotron_moe_prefill",
184                "nemotron_moe_relu2_down_prefill",
185            ),
186            moe_weighted_sum_prefill_k: super::try_kernel(
187                gpu,
188                "nemotron_moe_prefill",
189                "nemotron_moe_weighted_sum_prefill",
190            ),
191            moe_sort_k: super::try_kernel(gpu, "moe", "moe_sort_by_expert"),
192            moe_grouped_gemm_k: super::try_kernel(
193                gpu,
194                "moe_w4a16",
195                "moe_w4a16_grouped_gemm_ptrtable",
196            ),
197            moe_relu2_elementwise_k: super::try_kernel(gpu, "relu2", "relu_squared_inplace"),
198            moe_grouped_gemm_relu2_k: super::try_kernel(
199                gpu,
200                "moe_w4a16",
201                "moe_w4a16_grouped_gemm_ptrtable_relu2",
202            ),
203            moe_w4a4_grouped_k: super::try_kernel(gpu, "moe_w4a4", "moe_w4a4_grouped_gemm_relu2"),
204            moe_unpermute_reduce_k: super::try_kernel(gpu, "moe", "moe_unpermute_reduce_indexed"),
205            moe_grouped_gemm_n128_k: super::try_kernel(
206                gpu,
207                "moe_w4a16",
208                "moe_w4a16_grouped_gemm_ptrtable_t",
209            ),
210            up_ptrs,
211            down_ptrs,
212            up_ptrs_t: None,
213            down_ptrs_t: None,
214            shared_up_t: None,
215            shared_down_t: None,
216            shared_up_pd_fp8: None,
217            shared_down_pd_fp8: None,
218            fc1_pd_fp8: None,
219            fc2_pd_fp8: None,
220            w4a16_gemm_t_k: super::try_kernel(gpu, "w4a16", "w4a16_gemm_t"),
221            w4a16_gemm_t_m128_k: super::try_kernel(gpu, "w4a16", "w4a16_gemm_t_m128"),
222            fp8_gemm_m128_k: super::try_kernel(gpu, "w4a16", "fp8_gemm_t_m128_mfast"),
223            w4a4_gemm_k: super::try_kernel(gpu, "w4a4", "w4a4_gemm_mfast"),
224            quantize_nvfp4_k: super::try_kernel(gpu, "quantize_nvfp4", "quantize_bf16_to_nvfp4"),
225        })
226    }
227}
228
229mod decode_helpers;
230mod prefill_fallback;
231mod prefill_shared_up;
232mod prefill_sorted;
233mod prefill_weights;
234mod ptr_tables;
235
236use prefill_sorted::SortedPrefillCtx;
237use ptr_tables::{build_ptr_table, build_ptr_table_from_weights};
238
239impl TransformerLayer for NemotronMoeLayer {
240    fn decode(
241        &self,
242        hidden: DevicePtr,
243        residual: DevicePtr,
244        _state: &mut dyn LayerState,
245        _kv_cache: &mut PagedKvCache,
246        _seq_len: usize,
247        _block_table: &mut Vec<u32>,
248        _disk_block_ids: &mut Vec<u32>,
249        _disk_last_offloaded_per_layer: &mut Vec<u32>,
250        ctx: &ForwardContext,
251        stream: u64,
252    ) -> Result<()> {
253        self.decode_inner(hidden, residual, ctx, stream)
254    }
255
256    /// Batched MoE prefill: uses GEMM for gate/fc1/fc2/shared, per-token for routing + experts.
257    ///
258    /// For Super 120B with 40 MoE layers, this replaces O(N * 7 kernel_launches) decode calls
259    /// with O(4 GEMMs + N * 3 kernel_launches), cutting TTFT by 30-50%.
260    #[allow(clippy::overly_complex_bool_expr)]
261    fn prefill(
262        &self,
263        hidden: DevicePtr,
264        residual: DevicePtr,
265        num_tokens: usize,
266        _state: &mut dyn LayerState,
267        _kv_cache: &mut PagedKvCache,
268        _seq_len_start: usize,
269        _block_table: &mut Vec<u32>,
270        _disk_block_ids: &mut Vec<u32>,
271        _disk_last_offloaded_per_layer: &mut Vec<u32>,
272        _kv_write_start: usize,
273        ctx: &ForwardContext,
274        stream: u64,
275    ) -> Result<()> {
276        let h = ctx.config.hidden_size;
277        let inter = self.moe_inter as u32;
278        let shared_inter = ctx.config.shared_expert_intermediate_size as u32;
279        let num_experts = ctx.config.num_experts as u32;
280        let top_k = self.top_k as u32;
281        let eps = ctx.config.rms_norm_eps as f32;
282        let scale = ctx.config.routed_scaling_factor as f32;
283        let n = num_tokens as u32;
284
285        // ── 1. Batched RMS norm: [N, H] → normed[N, H] + residual update ──
286        let normed = ctx.buffers.norm_output();
287        ops::rms_norm_residual(
288            ctx.gpu,
289            self.rms_norm_residual_k,
290            hidden,
291            &self.input_norm,
292            normed,
293            residual,
294            n,
295            h as u32,
296            eps,
297            stream,
298        )?;
299
300        // ── 2. Batched Gate GEMM: [N, H] x [H, num_experts]^T → [N, num_experts] ──
301        let gate_logits = ctx.buffers.gate_logits();
302        self.dense_gemm_prefill(
303            ctx.gpu,
304            normed,
305            &self.weights.gate,
306            gate_logits,
307            n,
308            num_experts,
309            h as u32,
310            stream,
311        )?;
312
313        // Check if batched MoE prefill kernels are available
314        let has_batched = self.topk_sigmoid_batched_k.0 != 0
315            && self.moe_up_prefill_k.0 != 0
316            && self.moe_relu2_down_prefill_k.0 != 0
317            && self.moe_weighted_sum_prefill_k.0 != 0;
318
319        // ── 3. Shared expert UP ──
320        // When batched MoE prefill is available, the shared expert UP is handled
321        // inside the batched UP kernel (step 5b). We only pre-compute here for
322        // the per-token fallback path or LatentMoE.
323        let shared_up_out_base = ctx.buffers.ssm_qkvz();
324        let use_batched_moe = has_batched && num_tokens > 1;
325        // Always compute shared expert UP — even when batched path overwrites it later.
326        // The batched UP kernel writes shared_up_out for shared blocks, but we need
327        // this result for the per-token fallback path AND it's harmless to overwrite.
328        // Arm selection (native FP8 → W4A4 → pre-dequant FP8 → transposed NVFP4
329        // → plain W4A16) lives in `prefill_shared_up.rs` (500-LoC cap split).
330        self.prefill_shared_up(normed, shared_up_out_base, n, h, shared_inter, ctx, stream)?;
331
332        // ── 4. LatentMoE: batched fc1_latent GEMM [N, H] → [N, L] ──
333        // Use attn_output as temp buffer (m*max_dim*2, large enough for [N, L]).
334        // Cannot use ssm_ba (too small) or moe_output (used later for unpermute).
335        let latent = self.moe_latent_size as u32;
336        let latent_base = if latent > 0 {
337            let latent_buf = ctx.buffers.attn_output();
338            if let Some(w_fp8) = self.fc1_pd_fp8 {
339                ops::fp8_gemm_m128_mfast(
340                    ctx.gpu,
341                    self.fp8_gemm_m128_k,
342                    normed,
343                    w_fp8,
344                    latent_buf,
345                    n,
346                    latent,
347                    h as u32,
348                    stream,
349                )?;
350            } else {
351                let fc1 = self.weights.fc1_latent_proj.as_ref().unwrap();
352                self.dense_gemm_prefill(
353                    ctx.gpu, normed, fc1, latent_buf, n, latent, h as u32, stream,
354                )?;
355            }
356            Some(latent_buf)
357        } else {
358            None
359        };
360
361        // ── 5. Batched routing + expert dispatch (N tokens, 4 kernel launches) ──
362        // When batched prefill kernels are available, replace the per-token loop
363        // (N × 5 launches = 10k+ launches) with 4 batched launches.
364        let scratch = ctx.buffers.scratch();
365        let indices_dev = scratch;
366        let weights_dev = scratch.offset(n as usize * top_k as usize * 4);
367
368        // Sorted MoE prefill: sort tokens by expert, then grouped GEMM.
369        // This is the proven Qwen pattern — avoids the crashing batched UP/DOWN kernels.
370        let use_sorted = use_batched_moe
371            && self.moe_sort_k.0 != 0
372            && self.moe_grouped_gemm_k.0 != 0
373            && self.moe_unpermute_reduce_k.0 != 0;
374
375        let p = SortedPrefillCtx {
376            n,
377            num_tokens,
378            h,
379            inter,
380            shared_inter,
381            num_experts,
382            top_k,
383            scale,
384            latent,
385            gate_logits,
386            indices_dev,
387            weights_dev,
388            normed,
389            hidden,
390            latent_base,
391            shared_up_out_base,
392        };
393        if use_sorted {
394            self.prefill_sorted_path(&p, ctx, stream)?;
395        } else {
396            self.prefill_fallback_path(&p, ctx, stream)?;
397        }
398
399        Ok(())
400    }
401
402    fn alloc_state(&self, _gpu: &dyn GpuBackend) -> Result<Box<dyn LayerState>> {
403        Ok(Box::new(EmptyLayerState))
404    }
405}