1use 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
25struct ExpertPtrTable {
27 packed_ptrs: DevicePtr,
28 scale_ptrs: DevicePtr,
29 scale2_vals: DevicePtr,
30}
31
32pub struct NemotronMoeLayer {
34 weights: NemotronMoeWeights,
35 input_norm: DenseWeight,
36 moe_latent_size: usize,
38 moe_inter: usize,
40 top_k: usize,
42 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 w4a16_gemv_sw_k: KernelHandle,
50 w8a16_gemv_k: KernelHandle,
53 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 dense_gemm_k: KernelHandle,
61 dense_gemm_pipelined_k: KernelHandle,
67 w4a16_gemm_k: KernelHandle,
68 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 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 up_ptrs_t: Option<ExpertPtrTable>,
85 down_ptrs_t: Option<ExpertPtrTable>,
86 shared_up_t: Option<QuantizedWeight>,
88 shared_down_t: Option<QuantizedWeight>,
89 shared_up_pd_fp8: Option<DevicePtr>,
92 shared_down_pd_fp8: Option<DevicePtr>,
93 fc1_pd_fp8: Option<DevicePtr>,
96 fc2_pd_fp8: Option<DevicePtr>,
97 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 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 #[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 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 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 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 let shared_up_out_base = ctx.buffers.ssm_qkvz();
324 let use_batched_moe = has_batched && num_tokens > 1;
325 self.prefill_shared_up(normed, shared_up_out_base, n, h, shared_inter, ctx, stream)?;
331
332 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 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 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}