1use anyhow::Result;
16use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
17use spark_runtime::kernel_args::{KernelLaunch, div_ceil};
18
19use crate::weight_map::{DenseWeight, Fp8Weight, NemotronSsmWeights, QuantizedWeight};
20
21mod prefill;
22mod prefill_proj;
23mod trait_impl;
24
25#[allow(dead_code)]
26pub struct NemotronMamba2Layer {
27 input_norm: DenseWeight,
28 ssm: NemotronSsmWeights,
29 in_proj_fp8: Option<Fp8Weight>,
31 out_proj_fp8: Option<Fp8Weight>,
32 native_fp8_prefill: bool,
38 in_proj_t: Option<QuantizedWeight>,
40 out_proj_t: Option<QuantizedWeight>,
41 in_proj_pd_fp8: Option<DevicePtr>,
44 out_proj_pd_fp8: Option<DevicePtr>,
45 in_proj_bf16: Option<DenseWeight>,
52 out_proj_bf16: Option<DenseWeight>,
53 rms_norm_residual_k: KernelHandle,
55 w4a16_gemv_k: KernelHandle,
56 w4a16_gemv_sw_k: KernelHandle,
58 w8a16_gemv_k: KernelHandle,
59 conv1d_update_k: KernelHandle,
60 mamba2_ssm_k: KernelHandle,
61 gated_rms_norm_k: KernelHandle,
62 residual_add_k: KernelHandle,
63 w4a16_gemm_k: KernelHandle,
65 w8a16_gemm_k: KernelHandle,
67 w8a16_gemm_pipelined_k: KernelHandle,
68 w4a16_gemm_t_k: KernelHandle,
69 w4a16_gemm_t_m128_k: KernelHandle,
70 fp8_gemm_t_k: KernelHandle,
71 fp8_fp8_gemm_t_k: KernelHandle,
72 dense_gemm_bf16_k: KernelHandle,
74 dense_gemv_bf16_k: KernelHandle,
75 bf16_to_fp8_k: KernelHandle,
76 w4a4_gemm_k: KernelHandle,
77 quantize_nvfp4_k: KernelHandle,
78 conv1d_prefill_k: KernelHandle,
79 conv1d_prefill_tp_k: KernelHandle,
80 mamba2_ssm_prefill_k: KernelHandle,
81 mamba2_ssm_prefill_persistent_k: KernelHandle,
82 ssd_cumsum_k: KernelHandle,
84 ssd_bmm_k: KernelHandle,
85 ssd_scan_k: KernelHandle,
86 d_inner: usize,
88 d_xbc: usize,
89 in_proj_size: usize,
90 num_heads: usize,
91 head_dim: usize,
92 state_size: usize,
93 n_groups: usize,
94 d_conv: usize,
95 h_state_bytes: usize,
96 conv_state_bytes: usize,
97 layer_idx: usize,
98}
99
100impl NemotronMamba2Layer {
101 pub fn new(
102 input_norm: DenseWeight,
103 ssm: NemotronSsmWeights,
104 config: &atlas_core::config::ModelConfig,
105 gpu: &dyn GpuBackend,
106 layer_idx: usize,
107 ) -> Result<Self> {
108 let num_heads = config.mamba_num_heads;
109 let head_dim = config.mamba_head_dim;
110 let state_size = config.ssm_state_size;
111 let n_groups = config.n_groups;
112 let d_conv = config.linear_conv_kernel_dim;
113 let d_inner = config.mamba2_d_inner();
114 let d_xbc = config.mamba2_d_xbc();
115 let in_proj_size = config.mamba2_in_proj_size();
116
117 Ok(Self {
118 input_norm,
119 ssm,
120 in_proj_fp8: None,
121 out_proj_fp8: None,
122 native_fp8_prefill: false,
123 in_proj_t: None,
124 out_proj_t: None,
125 in_proj_pd_fp8: None,
126 out_proj_pd_fp8: None,
127 in_proj_bf16: None,
128 out_proj_bf16: None,
129 rms_norm_residual_k: gpu.kernel("norm", "rms_norm_residual")?,
130 w4a16_gemv_k: gpu.kernel("w4a16_gemv", "w4a16_gemv")?,
131 w4a16_gemv_sw_k: super::try_kernel(gpu, "w4a16_gemv", "w4a16_gemv_sw"),
132 w8a16_gemv_k: super::try_kernel(gpu, "w8a16_gemv", "w8a16_gemv"),
133 conv1d_update_k: gpu.kernel("causal_conv1d", "causal_conv1d_update")?,
134 mamba2_ssm_k: gpu.kernel("mamba2_ssm", "mamba2_ssm_decode")?,
135 gated_rms_norm_k: gpu.kernel("norm", "gated_rms_norm")?,
136 residual_add_k: gpu.kernel("residual_add", "bf16_residual_add")?,
137 w4a16_gemm_k: gpu.kernel("w4a16", "w4a16_gemm")?,
138 w8a16_gemm_k: super::try_kernel(gpu, "w8a16_gemm", "w8a16_gemm"),
139 w8a16_gemm_pipelined_k: super::try_kernel(
140 gpu,
141 "w8a16_gemm_pipelined",
142 "w8a16_gemm_pipelined",
143 ),
144 w4a16_gemm_t_k: super::try_kernel(gpu, "w4a16", "w4a16_gemm_t"),
145 w4a16_gemm_t_m128_k: super::try_kernel(gpu, "w4a16", "w4a16_gemm_t_m128"),
146 fp8_gemm_t_k: super::try_kernel(gpu, "w4a16", "fp8_gemm_t_m128_mfast"),
147 fp8_fp8_gemm_t_k: super::try_kernel(gpu, "w4a16", "fp8_fp8_gemm_t_m128_mfast"),
148 dense_gemm_bf16_k: super::try_kernel(gpu, "gemm", "dense_gemm_bf16_pipelined"),
149 dense_gemv_bf16_k: super::try_kernel(gpu, "gemv", "dense_gemv_bf16"),
150 bf16_to_fp8_k: super::try_kernel(gpu, "w4a16", "bf16_to_fp8"),
151 w4a4_gemm_k: super::try_kernel(gpu, "w4a4", "w4a4_gemm_mfast"),
152 quantize_nvfp4_k: super::try_kernel(gpu, "quantize_nvfp4", "quantize_bf16_to_nvfp4"),
153 conv1d_prefill_k: gpu.kernel("causal_conv1d", "causal_conv1d_update_prefill")?,
154 conv1d_prefill_tp_k: super::try_kernel(
155 gpu,
156 "causal_conv1d",
157 "causal_conv1d_update_prefill_tp",
158 ),
159 mamba2_ssm_prefill_k: gpu.kernel("mamba2_ssm", "mamba2_ssm_prefill")?,
160 ssd_cumsum_k: super::try_kernel(gpu, "mamba2_ssd_chunk", "mamba2_ssd_cumsum"),
161 ssd_bmm_k: super::try_kernel(gpu, "mamba2_ssd_chunk", "mamba2_ssd_bmm"),
162 ssd_scan_k: super::try_kernel(gpu, "mamba2_ssd_chunk", "mamba2_ssd_scan"),
163 mamba2_ssm_prefill_persistent_k: super::try_kernel(
164 gpu,
165 "mamba2_ssm",
166 "mamba2_ssm_prefill_persistent",
167 ),
168 d_inner,
169 d_xbc,
170 in_proj_size,
171 num_heads,
172 head_dim,
173 state_size,
174 n_groups,
175 d_conv,
176 h_state_bytes: num_heads * head_dim * state_size * 4, conv_state_bytes: d_xbc * d_conv * 4, layer_idx,
179 })
180 }
181
182 pub fn set_fp8_weights(
199 &mut self,
200 in_proj: Option<Fp8Weight>,
201 out_proj: Option<Fp8Weight>,
202 prefill: bool,
203 ) -> Result<()> {
204 use crate::weight_map::WeightQuantFormat;
205 if let Some(ref w) = in_proj {
206 w.scale_format.expect(
207 WeightQuantFormat::Fp8BlockScaled,
208 "nemotron mamba2 in_proj (w8a16 expects [ceil(N/128),ceil(K/128)] FP32 block scales)",
209 );
210 }
211 if let Some(ref w) = out_proj {
212 w.scale_format.expect(
213 WeightQuantFormat::Fp8BlockScaled,
214 "nemotron mamba2 out_proj (w8a16 expects [ceil(N/128),ceil(K/128)] FP32 block scales)",
215 );
216 }
217 anyhow::ensure!(
218 self.w8a16_gemv_k.0 != 0,
219 "native FP8 SSM requires the w8a16_gemv kernel (decode)"
220 );
221 anyhow::ensure!(
222 !prefill || self.w8a16_gemm_pipelined_k.0 != 0 || self.w8a16_gemm_k.0 != 0,
223 "native FP8 SSM requires w8a16_gemm[_pipelined] (prefill)"
224 );
225 self.in_proj_fp8 = in_proj;
226 self.out_proj_fp8 = out_proj;
227 self.native_fp8_prefill = prefill;
228 Ok(())
229 }
230
231 pub fn ssm_weights(&self) -> &NemotronSsmWeights {
233 &self.ssm
234 }
235
236 pub fn set_prefill_weights(
240 &mut self,
241 in_proj_t: Option<QuantizedWeight>,
242 out_proj_t: Option<QuantizedWeight>,
243 ) {
244 self.in_proj_t = in_proj_t;
245 self.out_proj_t = out_proj_t;
246 }
247
248 pub fn set_bf16_weights(&mut self, in_proj: DenseWeight, out_proj: DenseWeight) {
260 self.in_proj_bf16 = Some(in_proj);
261 self.out_proj_bf16 = Some(out_proj);
262 }
263
264 pub fn bf16_native_ready(&self) -> bool {
267 self.in_proj_bf16.is_some()
268 && self.out_proj_bf16.is_some()
269 && self.dense_gemm_bf16_k.0 != 0
270 && self.dense_gemv_bf16_k.0 != 0
271 }
272
273 pub fn set_fp8_prefill_weights(&mut self, in_proj: DevicePtr, out_proj: DevicePtr) {
274 self.in_proj_pd_fp8 = Some(in_proj);
275 self.out_proj_pd_fp8 = Some(out_proj);
276 }
277
278 fn conv1d_update_biased(
283 &self,
284 gpu: &dyn GpuBackend,
285 conv_state: DevicePtr,
286 input: DevicePtr,
287 output: DevicePtr,
288 d_inner: u32,
289 d_conv: u32,
290 batch_size: u32,
291 stream: u64,
292 ) -> Result<()> {
293 KernelLaunch::new(gpu, self.conv1d_update_k)
294 .grid([div_ceil(d_inner, 256), batch_size, 1])
295 .block([256, 1, 1])
296 .arg_ptr(conv_state)
297 .arg_ptr(input)
298 .arg_ptr(self.ssm.conv1d_weight.weight)
299 .arg_ptr(self.ssm.conv1d_bias.weight)
300 .arg_ptr(output)
301 .arg_u32(batch_size)
302 .arg_u32(d_inner)
303 .arg_u32(d_conv)
304 .launch(stream)
305 }
306
307 #[allow(clippy::too_many_arguments)]
311 fn ssm_decode(
312 &self,
313 gpu: &dyn GpuBackend,
314 h_state: DevicePtr,
315 x: DevicePtr,
316 b_proj: DevicePtr,
317 c_proj: DevicePtr,
318 dt_raw: DevicePtr,
319 output: DevicePtr,
320 batch_size: u32,
321 stream: u64,
322 ) -> Result<()> {
323 KernelLaunch::new(gpu, self.mamba2_ssm_k)
324 .grid([self.num_heads as u32, batch_size, 1])
325 .block([self.state_size as u32, 1, 1])
326 .arg_ptr(h_state)
327 .arg_ptr(x)
328 .arg_ptr(b_proj)
329 .arg_ptr(c_proj)
330 .arg_ptr(dt_raw)
331 .arg_ptr(self.ssm.a_log.weight)
332 .arg_ptr(self.ssm.d_param.weight)
333 .arg_ptr(self.ssm.dt_bias.weight)
334 .arg_ptr(output)
335 .arg_u32(batch_size)
336 .arg_u32(self.num_heads as u32)
337 .arg_u32(self.head_dim as u32)
338 .arg_u32(self.state_size as u32)
339 .arg_u32(self.n_groups as u32)
340 .arg_f32(1e-9) .arg_f32(1e9) .launch(stream)
343 }
344}