spark_model/layers/nemotron_mamba2/
trait_impl.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! `impl TransformerLayer for NemotronMamba2Layer`.
4
5use anyhow::Result;
6use spark_runtime::gpu::{DevicePtr, GpuBackend};
7use spark_runtime::kv_cache::PagedKvCache;
8
9use super::NemotronMamba2Layer;
10use crate::layer::{ForwardContext, LayerState, SsmLayerState, TransformerLayer};
11use crate::layers::ops;
12
13impl TransformerLayer for NemotronMamba2Layer {
14    fn decode(
15        &self,
16        hidden: DevicePtr,
17        residual: DevicePtr,
18        state: &mut dyn LayerState,
19        _kv_cache: &mut spark_runtime::kv_cache::PagedKvCache,
20        _seq_len: usize,
21        _block_table: &mut Vec<u32>,
22        _disk_block_ids: &mut Vec<u32>,
23        _disk_last_offloaded_per_layer: &mut Vec<u32>,
24        ctx: &ForwardContext,
25        stream: u64,
26    ) -> Result<()> {
27        let h = ctx.config.hidden_size;
28        let eps = ctx.config.rms_norm_eps as f32;
29
30        let ssm_state = state
31            .as_any_mut()
32            .downcast_mut::<SsmLayerState>()
33            .ok_or_else(|| anyhow::anyhow!("Expected SsmLayerState"))?;
34
35        // 1. RMS norm + save residual
36        let normed = ctx.buffers.norm_output();
37        ops::rms_norm_residual(
38            ctx.gpu,
39            self.rms_norm_residual_k,
40            hidden,
41            &self.input_norm,
42            normed,
43            residual,
44            1,
45            h as u32,
46            eps,
47            stream,
48        )?;
49
50        // 2. in_proj GEMV: normed[hidden_size] -> proj[in_proj_size]
51        //    Layout: [z(d_inner) | xBC(d_xbc) | dt(num_heads)]
52        let proj = ctx.buffers.ssm_qkvz();
53        // Use FP8 GEMV if available (skips double-quantization lossy path)
54        if let Some(ref w) = self.in_proj_bf16 {
55            // Native BF16: `ssm.in_proj` is NULL here (never quantized).
56            ops::dense_gemv(
57                ctx.gpu,
58                self.dense_gemv_bf16_k,
59                normed,
60                w,
61                proj,
62                self.in_proj_size as u32,
63                h as u32,
64                stream,
65            )?;
66        } else if let Some(ref fp8w) = self.in_proj_fp8 {
67            ops::w8a16_gemv(
68                ctx.gpu,
69                self.w8a16_gemv_k,
70                normed,
71                fp8w.weight,
72                fp8w.row_scale,
73                proj,
74                self.in_proj_size as u32,
75                h as u32,
76                stream,
77            )?;
78        } else {
79            ops::w4a16_decode_gemv(
80                ctx.gpu,
81                self.w4a16_gemv_k,
82                self.w4a16_gemv_sw_k,
83                ctx.levers.gemv_sw,
84                normed,
85                &self.ssm.in_proj,
86                proj,
87                self.in_proj_size as u32,
88                h as u32,
89                stream,
90            )?;
91        }
92
93        // Pointers into projection output (BF16, 2 bytes per element)
94        let z_ptr = proj;
95        let xbc_ptr = proj.offset(self.d_inner * 2);
96        let dt_ptr = proj.offset((self.d_inner + self.d_xbc) * 2);
97
98        // 3. Conv1d update on xBC (with bias, fused SiLU)
99        let xbc_out = ctx.buffers.ssm_deinterleaved();
100        self.conv1d_update_biased(
101            ctx.gpu,
102            ssm_state.conv_state,
103            xbc_ptr,
104            xbc_out,
105            self.d_xbc as u32,
106            self.d_conv as u32,
107            1,
108            stream,
109        )?;
110
111        // 4. Split xBC_out into x, B, C (BF16 offsets)
112        let x_ptr = xbc_out;
113        let gs = self.n_groups * self.state_size;
114        let b_ptr = xbc_out.offset(self.d_inner * 2);
115        let c_ptr = xbc_out.offset((self.d_inner + gs) * 2);
116
117        // 5. SSM decode: state update + y output
118        let y_ptr = ctx.buffers.attn_output();
119        self.ssm_decode(
120            ctx.gpu,
121            ssm_state.h_state,
122            x_ptr,
123            b_ptr,
124            c_ptr,
125            dt_ptr,
126            y_ptr,
127            1,
128            stream,
129        )?;
130
131        // 6. Gated RMS norm: rms_norm(y, ssm_norm) * silu(z)
132        //    y is [d_inner], z is [d_inner], gate_stride = in_proj_size (z at start of proj)
133        let gated_out = ctx.buffers.norm_output();
134        let group_size = (self.d_inner / self.n_groups) as u32;
135        ops::gated_rms_norm(
136            ctx.gpu,
137            self.gated_rms_norm_k,
138            y_ptr,
139            z_ptr,
140            &self.ssm.ssm_norm,
141            gated_out,
142            1,
143            self.d_inner as u32,
144            self.in_proj_size as u32,
145            eps,
146            group_size,
147            stream,
148        )?;
149
150        // 7. out_proj GEMV: gated_out[d_inner] -> out[hidden_size]
151        // Use qkv_output (NOT ssm_qkvz) — ssm_qkvz still holds z_ptr being read
152        // by gated_rms_norm above. Writing out_proj to the same buffer creates a
153        // write-after-read race that corrupts the gate signal → all-zero output.
154        let out = ctx.buffers.qkv_output();
155        if let Some(ref w) = self.out_proj_bf16 {
156            ops::dense_gemv(
157                ctx.gpu,
158                self.dense_gemv_bf16_k,
159                gated_out,
160                w,
161                out,
162                h as u32,
163                self.d_inner as u32,
164                stream,
165            )?;
166        } else if let Some(ref fp8w) = self.out_proj_fp8 {
167            ops::w8a16_gemv(
168                ctx.gpu,
169                self.w8a16_gemv_k,
170                gated_out,
171                fp8w.weight,
172                fp8w.row_scale,
173                out,
174                h as u32,
175                self.d_inner as u32,
176                stream,
177            )?;
178        } else {
179            ops::w4a16_decode_gemv(
180                ctx.gpu,
181                self.w4a16_gemv_k,
182                self.w4a16_gemv_sw_k,
183                ctx.levers.gemv_sw,
184                gated_out,
185                &self.ssm.out_proj,
186                out,
187                h as u32,
188                self.d_inner as u32,
189                stream,
190            )?;
191        }
192
193        // 8. Residual add: hidden += out_proj_result (hidden unchanged by rms_norm_residual)
194        ops::residual_add(ctx.gpu, self.residual_add_k, hidden, out, h as u32, stream)?;
195
196        Ok(())
197    }
198
199    fn prefill(
200        &self,
201        hidden: DevicePtr,
202        residual: DevicePtr,
203        num_tokens: usize,
204        state: &mut dyn LayerState,
205        _kv_cache: &mut PagedKvCache,
206        _seq_len_start: usize,
207        _block_table: &mut Vec<u32>,
208        _disk_block_ids: &mut Vec<u32>,
209        _disk_last_offloaded_per_layer: &mut Vec<u32>,
210        _kv_write_start: usize,
211        ctx: &ForwardContext,
212        stream: u64,
213    ) -> Result<()> {
214        self.prefill_ssm(hidden, residual, num_tokens, state, ctx, stream)
215    }
216
217    fn alloc_state(&self, gpu: &dyn GpuBackend) -> Result<Box<dyn LayerState>> {
218        let h_state = gpu.alloc(self.h_state_bytes)?;
219        gpu.memset(h_state, 0, self.h_state_bytes)?;
220        let conv_state = gpu.alloc(self.conv_state_bytes)?;
221        gpu.memset(conv_state, 0, self.conv_state_bytes)?;
222        Ok(Box::new(SsmLayerState {
223            h_state,
224            conv_state,
225            h_state_checkpoint: None,
226            conv_state_checkpoint: None,
227            h_state_intermediates: Vec::new(),
228            conv_state_intermediates: Vec::new(),
229            h_is_f16: false,
230            // Stage-3 narrowing is GDN-only (`ssm_h_fp16_preconditions`
231            // refuses a non-GDN SSM stack), so this state is always FP32-wide.
232            h_prefill_stage: None,
233            ple: None,
234        }))
235    }
236}