spark_model/layers/qwen3_ssm/
trait_layer.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! `TransformerLayer` impl for [`Qwen3SsmLayer`] — the trait surface that
4//! forwards into the `trait_*` sibling modules holding the actual phases.
5//! Split out of `mod.rs` to keep it under the 500-LoC cap.
6
7use 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    /// Downcast hook so the LoRA install walk can reach this layer's MoE FFN
16    /// (Feature-1: routed-expert/router deltas exist on GDN layers too).
17    fn as_any_mut(&mut self) -> Option<&mut dyn std::any::Any> {
18        Some(self)
19    }
20
21    /// PLE's host half (hash + NVMe fault-in + slot upload), hoisted before
22    /// graph replay/capture. No-op on the 47 layers without a PLE site.
23    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    /// PLE's per-seq host hash on the hc multi-seq decode path is
42    /// capture-illegal (pageable reads); the single-decode path prestages
43    /// around it, the batched path does not — veto batched graphs.
44    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            // Sequence never ran this layer (snapshot before first pass):
64            // nothing to carry, and restore-side declines aux-less slots.
65            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        // v1 is C=1 only under an mHC highway: these paths keep their own
140        // residual bookkeeping, which the highway replaces. Refusing is the
141        // point — a batched GDN step running on an unmixed stream produces
142        // plausible, wrong activations. Avarok #753.
143        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            // #753 item B milestone 2: the highway replaces the residual the
200            // non-hc path folds into its fused norm kernels; run the
201            // hc-bracketed variant instead of refusing.
202            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        // Under an mHC highway the residual bookkeeping is completely
233        // different — the highway IS the residual — so this is a second entry
234        // path, not a flag on the first. See `trait_prefill_hc.rs`.
235        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    /// Free the PLE carry this sequence lazily attached.
414    ///
415    /// Only the `ple` field — the h/conv state in `SsmLayerState` is pooled
416    /// and released by slot in `free_sequence_dispatch`, so freeing it here
417    /// would be a double free. The PLE conv buffer is the one piece that is
418    /// allocated per sequence and owned by nothing.
419    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
436/// The PLE per-seq carry from a sequence's [`SsmLayerState`], lazily created
437/// on first use. Errors if the state is not an `SsmLayerState`.
438fn 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}