1use anyhow::Result;
8use spark_runtime::gpu::{DevicePtr, GpuBackend};
9use spark_runtime::kv_cache::PagedKvCache;
10
11use super::Qwen3SsmLayer;
12use crate::layer::{ForwardContext, GdnPrefillBuffers, LayerState, TransformerLayer};
13
14impl TransformerLayer for Qwen3SsmLayer {
15 fn as_any_mut(&mut self) -> Option<&mut dyn std::any::Any> {
18 Some(self)
19 }
20
21 fn decode_prestage(
24 &self,
25 token: u32,
26 state: &mut dyn LayerState,
27 gpu: &dyn GpuBackend,
28 stream: u64,
29 ) -> Result<()> {
30 if let Some(ple) = self.ple.as_ref() {
31 let st = ple_seq_state(ple, state, gpu)?;
32 ple.prestage(st, &[token], gpu, stream)?;
33 }
34 Ok(())
35 }
36
37 fn has_aux_state(&self) -> bool {
38 self.ple.is_some()
39 }
40
41 fn decode_graph_unsupported(&self) -> bool {
45 self.ple.is_some()
46 }
47
48 fn snapshot_aux(
49 &self,
50 state: &dyn LayerState,
51 gpu: &dyn GpuBackend,
52 stream: u64,
53 ) -> Result<Option<Vec<u8>>> {
54 let Some(ple) = self.ple.as_ref() else {
55 return Ok(None);
56 };
57 let ssm = state
58 .as_any()
59 .downcast_ref::<crate::layer::SsmLayerState>()
60 .ok_or_else(|| anyhow::anyhow!("PLE host layer state is not SsmLayerState"))?;
61 match ssm.ple.as_ref() {
62 Some(st) => Ok(Some(ple.snapshot_aux(st, gpu, stream)?)),
63 None => Ok(None),
66 }
67 }
68
69 fn restore_aux(
70 &self,
71 state: &mut dyn LayerState,
72 blob: &[u8],
73 gpu: &dyn GpuBackend,
74 stream: u64,
75 ) -> Result<()> {
76 let ple = self
77 .ple
78 .as_ref()
79 .ok_or_else(|| anyhow::anyhow!("restore_aux: no PLE on this layer"))?;
80 let st = ple_seq_state(ple, state, gpu)?;
81 ple.restore_aux(st, blob, gpu, stream)
82 }
83
84 fn decode_prestage_rearm(&self, state: &mut dyn LayerState) {
85 if let Some(ple) = self.ple.as_ref()
86 && let Some(ssm) = state
87 .as_any_mut()
88 .downcast_mut::<crate::layer::SsmLayerState>()
89 && let Some(st) = ssm.ple.as_mut()
90 {
91 ple.rearm(st);
92 }
93 }
94
95 fn decode(
96 &self,
97 hidden: DevicePtr,
98 residual: DevicePtr,
99 state: &mut dyn LayerState,
100 kv_cache: &mut PagedKvCache,
101 seq_len: usize,
102 block_table: &mut Vec<u32>,
103 disk_block_ids: &mut Vec<u32>,
104 disk_last_offloaded_per_layer: &mut Vec<u32>,
105 ctx: &ForwardContext,
106 stream: u64,
107 ) -> Result<()> {
108 if self.hc.is_some() {
109 return self.decode_inner_hc(hidden, state, ctx, stream);
110 }
111 self.decode_inner(
112 hidden,
113 residual,
114 state,
115 kv_cache,
116 seq_len,
117 block_table,
118 disk_block_ids,
119 disk_last_offloaded_per_layer,
120 ctx,
121 stream,
122 )
123 }
124
125 fn decode_batched(
126 &self,
127 hidden: DevicePtr,
128 residual: DevicePtr,
129 num_tokens: usize,
130 state: &mut dyn LayerState,
131 _kv_cache: &mut PagedKvCache,
132 _seq_len: usize,
133 _block_table: &mut Vec<u32>,
134 _disk_block_ids: &mut Vec<u32>,
135 _disk_last_offloaded_per_layer: &mut Vec<u32>,
136 ctx: &ForwardContext,
137 stream: u64,
138 ) -> Result<()> {
139 self.refuse_batched_under_hc("decode_batched")?;
144 self.decode_batched_inner(
145 hidden,
146 residual,
147 num_tokens,
148 super::trait_decode_batched::GdnStates::Single(state),
149 ctx,
150 stream,
151 )
152 }
153
154 fn decode_verify_multi<'a, 'b: 'a>(
155 &self,
156 hidden: DevicePtr,
157 residual: DevicePtr,
158 n_seqs: usize,
159 ks: &[usize],
160 states: &'a mut [&'b mut (dyn LayerState + 'static)],
161 _kv_cache: &mut PagedKvCache,
162 wy_tables: DevicePtr,
163 ctx: &ForwardContext,
164 stream: u64,
165 ) -> Result<()> {
166 self.refuse_batched_under_hc("decode_verify_multi")?;
167 anyhow::ensure!(
168 states.len() == n_seqs && ks.len() == n_seqs,
169 "decode_verify_multi: states/ks/n mismatch"
170 );
171 let num_tokens: usize = ks.iter().sum();
172 self.decode_batched_inner(
173 hidden,
174 residual,
175 num_tokens,
176 super::trait_decode_batched::GdnStates::Multi {
177 states,
178 ks,
179 wy_tables,
180 },
181 ctx,
182 stream,
183 )
184 }
185
186 fn decode_multi_seq<'a, 'b: 'a>(
187 &self,
188 hidden: DevicePtr,
189 residual: DevicePtr,
190 num_seqs: usize,
191 states: &'a mut [&'b mut (dyn LayerState + 'static)],
192 kv_cache: &mut PagedKvCache,
193 seq_lens: &[usize],
194 block_tables: &[Vec<u32>],
195 ctx: &ForwardContext,
196 stream: u64,
197 ) -> Result<()> {
198 if self.hc.is_some() {
199 return self.decode_multi_seq_inner_hc(hidden, num_seqs, states, seq_lens, ctx, stream);
203 }
204 self.decode_multi_seq_inner(
205 hidden,
206 residual,
207 num_seqs,
208 states,
209 kv_cache,
210 seq_lens,
211 block_tables,
212 ctx,
213 stream,
214 )
215 }
216
217 fn prefill(
218 &self,
219 hidden: DevicePtr,
220 residual: DevicePtr,
221 num_tokens: usize,
222 state: &mut dyn LayerState,
223 kv_cache: &mut PagedKvCache,
224 seq_len_start: usize,
225 block_table: &mut Vec<u32>,
226 disk_block_ids: &mut Vec<u32>,
227 disk_last_offloaded_per_layer: &mut Vec<u32>,
228 kv_write_start: usize,
229 ctx: &ForwardContext,
230 stream: u64,
231 ) -> Result<()> {
232 if self.hc.is_some() {
236 return self.prefill_inner_hc(hidden, num_tokens, state, seq_len_start, ctx, stream);
237 }
238 self.prefill_inner(
239 hidden,
240 residual,
241 num_tokens,
242 state,
243 kv_cache,
244 seq_len_start,
245 block_table,
246 disk_block_ids,
247 disk_last_offloaded_per_layer,
248 kv_write_start,
249 ctx,
250 stream,
251 )
252 }
253
254 fn is_ssm_layer(&self) -> bool {
255 self.is_ssm_layer_inner()
256 }
257
258 fn prefill_phase1(
259 &self,
260 hidden: DevicePtr,
261 residual: DevicePtr,
262 num_tokens: usize,
263 state: &mut dyn LayerState,
264 kv_cache: &mut PagedKvCache,
265 seq_len_start: usize,
266 block_table: &mut Vec<u32>,
267 disk_block_ids: &mut Vec<u32>,
268 disk_last_offloaded_per_layer: &mut Vec<u32>,
269 kv_write_start: usize,
270 gdn_bufs: &GdnPrefillBuffers,
271 token_offset: usize,
272 ctx: &ForwardContext,
273 stream: u64,
274 ) -> Result<()> {
275 self.prefill_phase1_inner(
276 hidden,
277 residual,
278 num_tokens,
279 state,
280 kv_cache,
281 seq_len_start,
282 block_table,
283 disk_block_ids,
284 disk_last_offloaded_per_layer,
285 kv_write_start,
286 gdn_bufs,
287 token_offset,
288 ctx,
289 stream,
290 )
291 }
292
293 fn prefill_phase1_proj_batched(
294 &self,
295 hidden_stacked: DevicePtr,
296 residual_stacked: DevicePtr,
297 total_tokens: usize,
298 gdn_bufs: &GdnPrefillBuffers,
299 ctx: &ForwardContext,
300 stream: u64,
301 ) -> Result<()> {
302 self.prefill_phase1_proj_batched_inner(
303 hidden_stacked,
304 residual_stacked,
305 total_tokens,
306 gdn_bufs,
307 ctx,
308 stream,
309 )
310 }
311
312 fn prefill_phase1_conv1d_one(
313 &self,
314 state: &mut dyn LayerState,
315 token_offset: usize,
316 len: usize,
317 gdn_bufs: &GdnPrefillBuffers,
318 ctx: &ForwardContext,
319 stream: u64,
320 ) -> Result<()> {
321 self.prefill_phase1_conv1d_one_inner(state, token_offset, len, gdn_bufs, ctx, stream)
322 }
323
324 fn prefill_phase1_l2_batched(
325 &self,
326 total_tokens: usize,
327 gdn_bufs: &GdnPrefillBuffers,
328 ctx: &ForwardContext,
329 stream: u64,
330 ) -> Result<()> {
331 self.prefill_phase1_l2_batched_inner(total_tokens, gdn_bufs, ctx, stream)
332 }
333
334 fn prefill_gdn_full(
335 &self,
336 state: &mut dyn LayerState,
337 gdn_bufs: &GdnPrefillBuffers,
338 ctx: &ForwardContext,
339 stream: u64,
340 ) -> Result<()> {
341 self.prefill_gdn_full_inner(state, gdn_bufs, ctx, stream)
342 }
343
344 fn prefill_gdn_full_batched(
345 &self,
346 h_state_ptrs: DevicePtr,
347 gdn_bufs: &GdnPrefillBuffers,
348 batch_size: u32,
349 chunk_len: u32,
350 ctx: &ForwardContext,
351 stream: u64,
352 ) -> Result<()> {
353 self.prefill_gdn_full_batched_inner(
354 h_state_ptrs,
355 gdn_bufs,
356 batch_size,
357 chunk_len,
358 ctx,
359 stream,
360 )
361 }
362
363 fn prefill_gdn_full_batched_fla_varlen(
364 &self,
365 h_state_ptrs: DevicePtr,
366 gdn_bufs: &GdnPrefillBuffers,
367 batch_size: u32,
368 cu_seqlens: DevicePtr,
369 max_num_chunks: u32,
370 total_nt: usize,
371 max_seqlen: u32,
372 ctx: &ForwardContext,
373 stream: u64,
374 ) -> Result<bool> {
375 self.prefill_gdn_full_batched_fla_varlen_inner(
376 h_state_ptrs,
377 gdn_bufs,
378 batch_size,
379 cu_seqlens,
380 max_num_chunks,
381 total_nt,
382 max_seqlen,
383 ctx,
384 stream,
385 )
386 }
387
388 fn prefill_phase3(
389 &self,
390 hidden: DevicePtr,
391 residual: DevicePtr,
392 num_tokens: usize,
393 gdn_bufs: &GdnPrefillBuffers,
394 token_offset: usize,
395 ctx: &ForwardContext,
396 stream: u64,
397 ) -> Result<()> {
398 self.prefill_phase3_inner(
399 hidden,
400 residual,
401 num_tokens,
402 gdn_bufs,
403 token_offset,
404 ctx,
405 stream,
406 )
407 }
408
409 fn alloc_state(&self, gpu: &dyn GpuBackend) -> Result<Box<dyn LayerState>> {
410 self.alloc_state_inner(gpu)
411 }
412
413 fn release_state(&self, state: &mut dyn LayerState, gpu: &dyn GpuBackend) -> Result<()> {
420 let Some(ssm) = state
421 .as_any_mut()
422 .downcast_mut::<crate::layer::SsmLayerState>()
423 else {
424 return Ok(());
425 };
426 let Some(mut st) = ssm.ple.take() else {
427 return Ok(());
428 };
429 let Some(ple) = self.ple.as_ref() else {
430 anyhow::bail!("release_state: PLE seq state present but layer has no PLE");
431 };
432 ple.release_seq_state(&mut st, gpu)
433 }
434}
435
436fn ple_seq_state<'a>(
439 ple: &crate::layers::ple::PleLayer,
440 state: &'a mut dyn LayerState,
441 gpu: &dyn GpuBackend,
442) -> Result<&'a mut crate::layers::ple::PleSeqState> {
443 let ssm = state
444 .as_any_mut()
445 .downcast_mut::<crate::layer::SsmLayerState>()
446 .ok_or_else(|| anyhow::anyhow!("PLE host layer state is not SsmLayerState"))?;
447 if ssm.ple.is_none() {
448 ssm.ple = Some(ple.new_seq_state(gpu)?);
449 }
450 Ok(ssm.ple.as_mut().expect("just created"))
451}