spark_model/layers/nemotron_mamba2/
trait_impl.rs1use 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 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 let proj = ctx.buffers.ssm_qkvz();
53 if let Some(ref w) = self.in_proj_bf16 {
55 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 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 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 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 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 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 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 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 h_prefill_stage: None,
233 ple: None,
234 }))
235 }
236}