spark_model/model/trait_impl/
mod.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! `impl Model for TransformerModel` — thin trait impl that delegates to
4//! `<method>_dispatch` helpers split across sibling files for the ≤500
5//! LoC cap. Each sibling adds methods to the `TransformerModel`
6//! inherent impl. The trait impl below is purely one-line delegators.
7
8#![allow(unused_imports, dead_code, clippy::too_many_arguments)]
9
10use anyhow::Result;
11use spark_runtime::gpu::DevicePtr;
12use spark_runtime::kv_cache::PagedKvCache;
13
14use super::types::{PinnedMetaStaging, TransformerModel};
15use crate::layer::{AttnMetadataDev, LayerState};
16use crate::speculative::DraftProposer;
17use crate::traits::{ChunkedPrefillPageMetadata, Model, PrefillSlice, SequenceState};
18use crate::weight_map::{DenseWeight, MtpWeights};
19
20mod async_chkpt;
21mod decode_a;
22mod decode_a2;
23mod decode_a3;
24mod decode_a_diag;
25mod decode_b;
26mod decode_b2;
27mod decode_checkpoint;
28mod decode_graph_key;
29mod decode_multi_seq_gate;
30mod drafter_prefill;
31mod ep_misc;
32mod graph_borrow;
33mod lm_head_batched;
34mod meta;
35mod prefill_a;
36mod prefill_b;
37mod prefill_c;
38mod prefill_d;
39mod prefix_reuse;
40mod sequence;
41mod speculative;
42pub(in crate::model) mod ssm_fault_in;
43mod verify_a;
44mod verify_b;
45mod verify_c;
46mod verify_c2;
47mod verify_d;
48mod verify_e;
49pub(in crate::model) mod verify_e2;
50mod verify_fused;
51
52impl Model for TransformerModel {
53    fn teardown(&mut self) -> Result<()> {
54        self.release_pools()
55    }
56
57    /// Poll this model's own InnerQ driver. A miss is logged, never fatal — it
58    /// is a diagnostic lever, not part of serving.
59    #[cfg(feature = "cuda")]
60    fn poll_innerq(&self) {
61        if let Some(driver) = self.innerq.as_ref()
62            && let Err(e) = driver.maybe_finalize(128)
63        {
64            tracing::warn!("InnerQ maybe_finalize failed: {e:#}");
65        }
66    }
67
68    fn prepare_vision_embed(&self, images: &[crate::VisionItem]) -> Result<()> {
69        self.prepare_vision_embed_dispatch(images)
70    }
71    fn prepare_vision_embed_batched(
72        &self,
73        per_request: &[Vec<crate::VisionItem>],
74    ) -> Result<Vec<(usize, usize, usize, usize)>> {
75        self.prepare_vision_embed_batched_dispatch(per_request)
76    }
77    fn set_vision_slice_base(&self, row_base: usize, grid_base: usize, owned_images: usize) {
78        *self.vision_row_base.lock() = row_base;
79        *self.vision_grid_base.lock() = grid_base;
80        *self.vision_owned_images.lock() = owned_images;
81    }
82    // The four prefill entry points each end with `try_eager_drafter_prefill`:
83    // the whole-prompt drafter capture is a single shared slot, so it must be
84    // consumed while THIS sequence still owns it — one tick later, at the
85    // first propose, a concurrent sequence's prefill has already restarted it
86    // and every sequence but the last-prefilled drafts blind. See
87    // `drafter_prefill.rs`. Kill switch `ATLAS_NO_MTP_EAGER_DRAFTER`.
88    fn tokens_contain_vision_pad(&self, tokens: &[u32]) -> bool {
89        self.tokens_have_vision_pad(tokens)
90    }
91    fn prefill(&self, tokens: &[u32], seq: &mut SequenceState, stream: u64) -> Result<DevicePtr> {
92        self.stamp_overlay_route(seq.adapter_slot);
93        let logits = self.prefill_dispatch(tokens, seq, stream)?;
94        self.try_eager_drafter_prefill(seq, true, stream);
95        Ok(logits)
96    }
97    fn prefill_chunk(
98        &self,
99        tokens: &[u32],
100        seq: &mut SequenceState,
101        chunk_start: usize,
102        chunk_len: usize,
103        is_last_chunk: bool,
104        stream: u64,
105    ) -> Result<DevicePtr> {
106        self.stamp_overlay_route(seq.adapter_slot);
107        let logits = self.prefill_chunk_dispatch(
108            tokens,
109            seq,
110            chunk_start,
111            chunk_len,
112            is_last_chunk,
113            stream,
114        )?;
115        self.try_eager_drafter_prefill(seq, is_last_chunk, stream);
116        Ok(logits)
117    }
118    fn prefill_twophase(
119        &self,
120        tokens: &[u32],
121        seq: &mut SequenceState,
122        chunk_size: usize,
123        stream: u64,
124    ) -> Result<DevicePtr> {
125        self.stamp_overlay_route(seq.adapter_slot);
126        let logits = self.prefill_twophase_dispatch(tokens, seq, chunk_size, stream)?;
127        self.try_eager_drafter_prefill(seq, true, stream);
128        Ok(logits)
129    }
130    fn decode(&self, token: u32, seq: &mut SequenceState, _stream: u64) -> Result<DevicePtr> {
131        self.stamp_overlay_route(seq.adapter_slot);
132        self.stamp_decode_moe_single(seq.adapter_slot);
133        self.decode_dispatch(token, seq, _stream)
134    }
135    fn decode_batch(
136        &self,
137        tokens: &[u32],
138        seqs: &mut [&mut SequenceState],
139        stream: u64,
140    ) -> Result<DevicePtr> {
141        self.stamp_overlay_route_batch(seqs);
142        self.stamp_decode_moe_batch(seqs);
143        let r = self.decode_batch_dispatch(tokens, seqs, stream);
144        if r.is_err() {
145            // A mid-capture refuse (MoE LoRA router/mixed/non-active) in the
146            // batched-decode compute leaves the capture stream recording; release
147            // it so the caller's sequence cleanup doesn't hit
148            // STREAM_CAPTURE_UNSUPPORTED and poison every later op (a single
149            // refused concurrent request would otherwise brick the server). The
150            // batched path captures on the default stream (decode_a2).
151            self.gpu.abort_capture_if_active(self.gpu.default_stream());
152        }
153        r
154    }
155    fn mixed_forward(
156        &self,
157        decode_tokens: &[u32],
158        decode_seqs: &mut [&mut SequenceState],
159        prefill_tokens: &[u32],
160        prefill_seq: &mut SequenceState,
161        prefill_chunk_start: usize,
162        prefill_chunk_len: usize,
163        prefill_is_last: bool,
164        stream: u64,
165    ) -> Result<crate::traits::MixedForwardResult> {
166        // Mixed decode+prefill batch spans multiple adapters ⇒ mark mixed so the
167        // overlay hooks skip (per-token seq_slot routing is SOLID Incr-4).
168        self.overlay_route_slot
169            .store(i32::MIN, std::sync::atomic::Ordering::Relaxed);
170        // Decode portion: Skip only if every decode seq is base, else refuse.
171        self.stamp_decode_moe_batch(decode_seqs);
172        let r = self.mixed_forward_dispatch(
173            decode_tokens,
174            decode_seqs,
175            prefill_tokens,
176            prefill_seq,
177            prefill_chunk_start,
178            prefill_chunk_len,
179            prefill_is_last,
180            stream,
181        );
182        if r.is_err() {
183            // Same brick guard as decode_batch: a refuse in the captured decode
184            // portion must not leave the default stream recording.
185            self.gpu.abort_capture_if_active(self.gpu.default_stream());
186        }
187        let out = r?;
188        self.try_eager_drafter_prefill(prefill_seq, prefill_is_last, stream);
189        Ok(out)
190    }
191
192    /// Q12 Phase 4b override. The concrete dispatcher routes ineligible
193    /// batches to its sequential path before state mutation. Errors from an
194    /// admitted kernel batch must propagate: retrying sequentially can
195    /// reapply prefix-cache and KV state.
196    fn prefill_batch_chunk(
197        &self,
198        streams: &mut [PrefillSlice<'_>],
199        stream: u64,
200    ) -> Result<Vec<DevicePtr>> {
201        self.prefill_batch_chunk_rows(streams, stream, 0)
202    }
203    /// Mixed-step variant: shift the finishing streams' logits rows clear of
204    /// the decode lanes. See the trait docs for the aliasing this prevents.
205    fn prefill_batch_chunk_rows(
206        &self,
207        streams: &mut [PrefillSlice<'_>],
208        stream: u64,
209        row_base: usize,
210    ) -> Result<Vec<DevicePtr>> {
211        self.prefill_batch_chunk_dispatch(streams, stream, row_base)
212    }
213    fn vocab_size(&self) -> usize {
214        self.vocab_size_dispatch()
215    }
216    fn set_active_lora(&mut self, name: &str) -> Result<()> {
217        self.rotate_lora_to(name)
218    }
219    fn adapter_id_for(&self, slot: i32) -> u64 {
220        self.adapter_id_for_slot(slot)
221    }
222    fn acquire_adapter_slot(&self, slot: i32) -> i32 {
223        TransformerModel::acquire_adapter_slot(self, slot)
224    }
225    fn release_adapter_slot(&self, resolved: i32) {
226        TransformerModel::release_adapter_slot(self, resolved)
227    }
228    fn swap_lora_from_disk(
229        &mut self,
230        dir: &std::path::Path,
231        name: &str,
232        slot: usize,
233    ) -> Result<()> {
234        // Disk staging is plain file I/O and is portable; only the PEER path
235        // needs RDMA. Still cuda-gated, since it lands into a device pool.
236        #[cfg(feature = "cuda")]
237        {
238            self.swap_lora_slot_from_disk(dir, name, slot)
239        }
240        #[cfg(not(feature = "cuda"))]
241        {
242            let _ = (dir, name, slot);
243            anyhow::bail!("LoRA disk swap requires the cuda feature")
244        }
245    }
246    fn promote_lora_from_peer(
247        &mut self,
248        peer_addr: &str,
249        adapter_id: &str,
250        name: &str,
251        peft: atlas_core::config::PeftAdapterConfig,
252    ) -> Result<(usize, Option<String>)> {
253        #[cfg(all(feature = "cuda", unix))]
254        {
255            self.promote_lora_slot_from_peer(peer_addr, adapter_id, name, peft)
256        }
257        #[cfg(not(all(feature = "cuda", unix)))]
258        {
259            let _ = (peer_addr, adapter_id, name, peft);
260            anyhow::bail!("LoRA peer promotion stages over RDMA (rdma-core); unix-only")
261        }
262    }
263    fn promote_lora_from_disk(
264        &mut self,
265        dir: &std::path::Path,
266        name: &str,
267    ) -> Result<(usize, Option<String>)> {
268        #[cfg(feature = "cuda")]
269        {
270            self.promote_lora_slot_from_disk(dir, name)
271        }
272        #[cfg(not(feature = "cuda"))]
273        {
274            let _ = (dir, name);
275            anyhow::bail!("LoRA disk promotion requires the cuda feature")
276        }
277    }
278    fn high_speed_swap_dims(&self) -> Option<spark_storage::ModelDims> {
279        self.high_speed_swap_dims_dispatch()
280    }
281    fn normalize_ssm_states(&self, seq: &SequenceState, stream: u64) -> Result<()> {
282        self.normalize_ssm_states_dispatch(seq, stream)
283    }
284    fn bind_gpu_to_thread(&self) -> Result<()> {
285        self.bind_gpu_to_thread_dispatch()
286    }
287    fn alloc_sequence(&self) -> Result<SequenceState> {
288        self.alloc_sequence_dispatch(usize::MAX)
289    }
290
291    fn alloc_sequence_for(&self, budget_tokens: usize) -> Result<SequenceState> {
292        self.alloc_sequence_dispatch(budget_tokens)
293    }
294    fn copy_logits_to_host(&self, logits_ptr: DevicePtr, dst: &mut [u8]) -> Result<()> {
295        self.copy_logits_to_host_dispatch(logits_ptr, dst)
296    }
297    fn logits_ptr_is_fp32(&self, logits_ptr: DevicePtr) -> bool {
298        self.logits_ptr_is_fp32_dispatch(logits_ptr)
299    }
300    fn logits_buffer_ptr(&self) -> DevicePtr {
301        self.logits_buffer_ptr_dispatch()
302    }
303    fn argmax_on_device(&self, logits_ptr: DevicePtr, _stream: u64) -> Result<u32> {
304        self.argmax_on_device_dispatch(logits_ptr, _stream)
305    }
306    fn argmax_batch(&self, logits_ptr: DevicePtr, n: usize, _stream: u64) -> Result<Vec<u32>> {
307        self.argmax_batch_dispatch(logits_ptr, n, _stream)
308    }
309    fn hidden_after_norm(&self) -> DevicePtr {
310        self.hidden_after_norm_dispatch()
311    }
312    fn decode_verify(
313        &self,
314        tokens: &[u32],
315        seq: &mut SequenceState,
316        stream: u64,
317    ) -> Result<Vec<u32>> {
318        self.ssm_pool.require_verify_rollback_supported()?;
319        let r = self.decode_verify_dispatch(tokens, seq, stream);
320        if r.is_err() {
321            // Same brick guard as decode_batch: a refuse mid-verify-capture
322            // (MTP/spec) must not leave the default stream recording. No-op when
323            // not capturing. Verify captures on default_stream (verify_a/b/…).
324            self.gpu.abort_capture_if_active(self.gpu.default_stream());
325        }
326        r
327    }
328    fn checkpoint_ssm_states(&self, seq: &mut SequenceState) -> Result<()> {
329        self.checkpoint_ssm_states_dispatch(seq)
330    }
331    fn rollback_ssm_states(&self, seq: &mut SequenceState, num_accepted: usize) -> Result<()> {
332        self.rollback_ssm_states_dispatch(seq, num_accepted)
333    }
334    fn has_ssm_layers(&self) -> bool {
335        self.ssm_pool.num_ssm_layers > 0
336    }
337    fn mtp_slot_draft_capacity(&self, slot_idx: usize) -> usize {
338        self.ssm_pool.verify_draft_capacity(slot_idx)
339    }
340    fn decode_rollback_ring_slots(&self) -> usize {
341        if self.ssm_snapshots.decode_rollback_enabled() {
342            self.ssm_snapshots.decode_ring_slots
343        } else {
344            0
345        }
346    }
347    fn save_decode_ssm_snapshot(&self, seq: &SequenceState, ring_slot: usize) -> Result<()> {
348        self.save_decode_ssm_snapshot_dispatch(seq, ring_slot)
349    }
350    fn restore_decode_ssm_snapshot(&self, seq: &SequenceState, ring_slot: usize) -> Result<()> {
351        self.restore_decode_ssm_snapshot_dispatch(seq, ring_slot)
352    }
353    fn generate_speculative(
354        &self,
355        prompt_tokens: &[u32],
356        params: &spark_runtime::sampler::SamplingParams,
357        num_drafts: usize,
358    ) -> Result<crate::engine::GenerateResult> {
359        self.generate_speculative_dispatch(prompt_tokens, params, num_drafts)
360    }
361    fn has_proposer(&self) -> bool {
362        self.has_proposer_dispatch()
363    }
364    fn dflash_gamma(&self) -> Option<usize> {
365        self.proposer.as_ref().and_then(|p| p.block_gamma())
366    }
367    fn has_self_speculative(&self) -> bool {
368        self.has_self_speculative_dispatch()
369    }
370    fn decode_draft(&self, token: u32, seq: &mut SequenceState, stream: u64) -> Result<DevicePtr> {
371        self.decode_draft_dispatch(token, seq, stream)
372    }
373    fn cache_sequence(&self, seq: &SequenceState) {
374        self.cache_sequence_dispatch(seq)
375    }
376    fn decode_marconi_checkpoint(&self, seq: &mut SequenceState) {
377        self.decode_marconi_checkpoint_dispatch(seq)
378    }
379    fn free_sequence(&self, seq: &mut SequenceState) -> Result<()> {
380        self.free_sequence_dispatch(seq)
381    }
382    fn decode_verify_graphed(
383        &self,
384        tokens: &[u32; 2],
385        seq: &mut SequenceState,
386        _stream: u64,
387    ) -> Result<[u32; 2]> {
388        self.ssm_pool.require_verify_rollback_supported()?;
389        self.decode_verify_graphed_dispatch(tokens, seq, _stream)
390    }
391    fn decode_verify_graphed_k3(
392        &self,
393        tokens: &[u32; 3],
394        seq: &mut SequenceState,
395        _stream: u64,
396    ) -> Result<[u32; 3]> {
397        self.ssm_pool.require_verify_rollback_supported()?;
398        self.decode_verify_graphed_k3_dispatch(tokens, seq, _stream)
399    }
400    fn decode_verify_graphed_k4(
401        &self,
402        tokens: &[u32; 4],
403        seq: &mut SequenceState,
404        _stream: u64,
405    ) -> Result<[u32; 4]> {
406        self.ssm_pool.require_verify_rollback_supported()?;
407        self.decode_verify_graphed_k4_dispatch(tokens, seq, _stream)
408    }
409    fn can_batch_verify(&self, ks: &[usize]) -> bool {
410        self.can_batch_verify_dispatch(ks)
411    }
412    fn decode_verify_batched(
413        &self,
414        tokens: &[u32],
415        ks: &[usize],
416        seqs: &mut [&mut SequenceState],
417        _stream: u64,
418    ) -> Result<Vec<u32>> {
419        self.ssm_pool.require_verify_rollback_supported()?;
420        self.decode_verify_batched_dispatch(tokens, ks, seqs, _stream)
421    }
422    fn stash_verify_hidden_rows(&self, rows: &[usize], _stream: u64) -> Result<()> {
423        self.stash_verify_hidden_rows_dispatch(rows, _stream)
424    }
425    fn save_hidden_for_mtp_from_stash(&self, idx: usize, _stream: u64) -> Result<()> {
426        self.save_hidden_for_mtp_from_stash_dispatch(idx, _stream)
427    }
428    fn run_mtp_propose_batched(
429        &self,
430        tokens: &[u32],
431        positions: &[usize],
432        stash_idx: &[usize],
433        num_drafts: usize,
434        seqs: &mut [&mut SequenceState],
435        _stream: u64,
436        out_conf: Option<&mut Vec<Vec<f32>>>,
437    ) -> Result<Option<Vec<Vec<u32>>>> {
438        self.run_mtp_propose_batched_dispatch(
439            tokens, positions, stash_idx, num_drafts, seqs, out_conf,
440        )
441    }
442    fn mtp_propose_batch_max(&self) -> usize {
443        match &self.proposer {
444            Some(p) => p.propose_batch_max(&self.buffers, &self.config),
445            None => 1,
446        }
447    }
448    fn decode_verify_graphed_kgamma(
449        &self,
450        tokens: &[u32],
451        seq: &mut SequenceState,
452        _stream: u64,
453    ) -> Result<Vec<u32>> {
454        self.ssm_pool.require_verify_rollback_supported()?;
455        self.decode_verify_graphed_kgamma_dispatch(tokens, seq, _stream)
456    }
457    fn decode_and_verify_fused(
458        &self,
459        tokens: &[u32],
460        seq: &mut SequenceState,
461        _stream: u64,
462    ) -> Result<Vec<u32>> {
463        self.ssm_pool.require_verify_rollback_supported()?;
464        self.decode_and_verify_fused_dispatch(tokens, seq, _stream)
465    }
466    fn save_hidden_for_catchup(&self, token_idx: usize, pos: usize) -> Result<()> {
467        self.save_hidden_for_catchup_dispatch(token_idx, pos)
468    }
469
470    fn save_hidden_for_mtp(&self, token_idx: usize, _stream: u64) -> Result<()> {
471        self.save_hidden_for_mtp_dispatch(token_idx, _stream)
472    }
473    fn save_dflash_hidden_for_propose(&self, token_idx: usize, _stream: u64) -> Result<()> {
474        self.save_dflash_hidden_dispatch(token_idx, _stream)
475    }
476
477    fn dflash_accept_append(&self, seq: &mut SequenceState) -> Result<()> {
478        let base = match self.dflash_hidden_save {
479            Some(p) => p,
480            None => return Ok(()),
481        };
482        let prop = match seq.proposer_state.as_mut() {
483            Some(p) => p.as_mut(),
484            None => return Ok(()),
485        };
486        let d = prop
487            .as_any_mut()
488            .downcast_mut::<crate::layers::DflashProposerState>()
489            .ok_or_else(|| anyhow::anyhow!("not DFlash proposer state"))?;
490        let n_layers = self.dflash_capture_layers.len();
491        if n_layers == 0 {
492            return Ok(());
493        }
494        let ctx_slot_bytes = n_layers * self.config.hidden_size * 2;
495        let save_1 = base.offset(ctx_slot_bytes);
496        let dst = d.ctx_hidden_acc.offset(d.ctx_len * ctx_slot_bytes);
497        self.gpu
498            .copy_d2d_async(save_1, dst, ctx_slot_bytes, self.gpu.default_stream())?;
499        d.ctx_positions.push((seq.seq_len as i32).saturating_sub(1));
500        d.ctx_len += 1;
501        Ok(())
502    }
503
504    fn dflash_eagle_accept_append(&self, seq: &mut SequenceState) -> Result<()> {
505        let base = match self.dflash_hidden_save {
506            Some(p) => p,
507            None => return Ok(()),
508        };
509        let prop = match seq.proposer_state.as_mut() {
510            Some(p) => p.as_mut(),
511            None => return Ok(()),
512        };
513        let d = prop
514            .as_any_mut()
515            .downcast_mut::<crate::layers::DflashProposerState>()
516            .ok_or_else(|| anyhow::anyhow!("not DFlash proposer state"))?;
517        let n_layers = self.dflash_capture_layers.len();
518        if n_layers == 0 {
519            return Ok(());
520        }
521        let ctx_slot_bytes = n_layers * self.config.hidden_size * 2;
522        let stream = self.gpu.default_stream();
523        let pos_row0 = (seq.seq_len as i32).saturating_sub(2);
524        let pos_row1 = (seq.seq_len as i32).saturating_sub(1);
525        // Row 0 @ N
526        let save_0 = base;
527        let dst_0 = d.ctx_hidden_acc.offset(d.ctx_len * ctx_slot_bytes);
528        self.gpu
529            .copy_d2d_async(save_0, dst_0, ctx_slot_bytes, stream)?;
530        d.ctx_positions.push(pos_row0);
531        d.ctx_len += 1;
532        // Row 1 @ N+1
533        let save_1 = base.offset(ctx_slot_bytes);
534        let dst_1 = d.ctx_hidden_acc.offset(d.ctx_len * ctx_slot_bytes);
535        self.gpu
536            .copy_d2d_async(save_1, dst_1, ctx_slot_bytes, stream)?;
537        d.ctx_positions.push(pos_row1);
538        d.ctx_len += 1;
539        d.skip_next_decode_append = true;
540        Ok(())
541    }
542
543    fn dflash_eagle_kgamma_append(
544        &self,
545        seq: &mut SequenceState,
546        num_accepted: usize,
547        base_pos: usize,
548    ) -> Result<()> {
549        let base = match self.dflash_hidden_save {
550            Some(p) => p,
551            None => return Ok(()),
552        };
553        let prop = match seq.proposer_state.as_mut() {
554            Some(p) => p.as_mut(),
555            None => return Ok(()),
556        };
557        let d = prop
558            .as_any_mut()
559            .downcast_mut::<crate::layers::DflashProposerState>()
560            .ok_or_else(|| anyhow::anyhow!("not DFlash proposer state"))?;
561        let n_layers = self.dflash_capture_layers.len();
562        if n_layers == 0 {
563            return Ok(());
564        }
565        let ctx_slot_bytes = n_layers * self.config.hidden_size * 2;
566        let stream = self.gpu.default_stream();
567        for t in 0..=num_accepted {
568            let row = base.offset(t * ctx_slot_bytes);
569            let dst = d.ctx_hidden_acc.offset(d.ctx_len * ctx_slot_bytes);
570            self.gpu.copy_d2d_async(row, dst, ctx_slot_bytes, stream)?;
571            let pos = (base_pos + t) as i32;
572            d.ctx_positions.push(pos);
573            d.ctx_len += 1;
574        }
575        d.skip_next_decode_append = true;
576        Ok(())
577    }
578
579    fn dflash_capture_band(&self) -> usize {
580        self.dflash_kgamma
581    }
582
583    fn commit_ctx(
584        &self,
585        seq: &mut SequenceState,
586        num_committed: usize,
587        base_pos: usize,
588        scratch_row: usize,
589    ) -> Result<()> {
590        if num_committed == 0 {
591            return Ok(());
592        }
593        let base = match self.dflash_hidden_save {
594            Some(p) => p,
595            None => return Ok(()),
596        };
597        let prop = match seq.proposer_state.as_mut() {
598            Some(p) => p.as_mut(),
599            None => return Ok(()),
600        };
601        // Graceful no-op for non-DFlash proposers (shared bootstrap path).
602        let d = match prop
603            .as_any_mut()
604            .downcast_mut::<crate::layers::DflashProposerState>()
605        {
606            Some(d) => d,
607            None => return Ok(()),
608        };
609        let n_layers = self.dflash_capture_layers.len();
610        if n_layers == 0 {
611            return Ok(());
612        }
613        let ctx_slot_bytes = n_layers * self.config.hidden_size * 2;
614        let stream = self.gpu.default_stream();
615
616        // Scratch-capacity guard: `try_dflash_capture_all` caps its writes at
617        // `dflash_hidden_save_rows` (γ+1), so a batch row beyond that was
618        // never captured — committing it would append a STALE row (poisoned
619        // ctx is worse than a hole). Skip with a warning; only reachable if
620        // --max-num-seqs exceeds γ+1 on a DFlash serve.
621        if scratch_row + num_committed > self.dflash_hidden_save_rows {
622            tracing::warn!(
623                "commit_ctx: scratch rows {}..{} exceed capture capacity {} — skipping (ctx hole)",
624                scratch_row,
625                scratch_row + num_committed,
626                self.dflash_hidden_save_rows,
627            );
628            return Ok(());
629        }
630
631        // Watermark slide FIRST, on the ctx_len (row-index) axis. If the
632        // incoming rows would exceed capacity, keep the NEWEST rows and drop
633        // the oldest (mirrors dflash_serial_ctx_append). keep is clamped so
634        // drop_n >= keep — the single D2D copy's src/dst can never overlap.
635        // ctx_committed resets to 0 (next propose re-precomputes the slid
636        // rows chunk-wise); ctx_positions values (absolute RoPE positions)
637        // are preserved by the drain, so stamps stay exact across the slide.
638        if d.ctx_len + num_committed > d.max_ctx_len {
639            let keep = (d.max_ctx_len / 2).min(d.max_ctx_len.saturating_sub(num_committed));
640            let drop_n = d.ctx_len.saturating_sub(keep);
641            if drop_n > 0 {
642                let src = d.ctx_hidden_acc.offset(drop_n * ctx_slot_bytes);
643                let dst0 = d.ctx_hidden_acc.offset(0);
644                self.gpu
645                    .copy_d2d_async(src, dst0, keep * ctx_slot_bytes, stream)?;
646                d.ctx_positions.drain(..drop_n);
647                d.ctx_len = keep;
648                d.ctx_committed = 0;
649                tracing::info!(
650                    "DFlash UNIFIED_CTX watermark: slid ctx window (dropped {} oldest, keep {})",
651                    drop_n,
652                    keep,
653                );
654            }
655        }
656
657        // Append num_committed rows at the TAIL (ctx_len axis). dst uses
658        // ctx_len (acc row index); base_pos stamps ctx_positions (RoPE axis).
659        // Conflating the two axes is the DDD §4.1 landmine: they coincide
660        // only until the first slide — and the sliding prompts ARE the reds.
661        debug_assert_eq!(d.ctx_positions.len(), d.ctx_len);
662        for t in 0..num_committed {
663            let row = base.offset((scratch_row + t) * ctx_slot_bytes);
664            let dst = d.ctx_hidden_acc.offset(d.ctx_len * ctx_slot_bytes);
665            self.gpu.copy_d2d_async(row, dst, ctx_slot_bytes, stream)?;
666            d.ctx_positions.push((base_pos + t) as i32);
667            d.ctx_len += 1;
668        }
669        // Freshest ctx slot = row (num_committed-1) = the bonus generator
670        // (EAGLE order, matches kgamma_append). Block the next propose()'s
671        // internal decode-append so this capture is never double-appended.
672        d.skip_next_decode_append = true;
673
674        // Per-commit ledger (debug): the C>=2 GAP diagnosis reads this to
675        // find steps whose commits under-append vs the positions committed.
676        tracing::debug!(
677            "CTX_COMMIT slot={} rows={} base_pos={} ctx_len_after={}",
678            seq.slot_idx,
679            num_committed,
680            base_pos,
681            d.ctx_len,
682        );
683        // One-shot activation log so A/B runs can confirm the path is live.
684        if self.stats.once("log:dflash_unified_ctx") {
685            tracing::info!(
686                "DFlash UNIFIED_CTX ACTIVE: first commit_ctx rows={} base_pos={} ctx_len={}",
687                num_committed,
688                base_pos,
689                d.ctx_len,
690            );
691        }
692        Ok(())
693    }
694
695    fn dflash_serial_ctx_append(&self, seq: &mut SequenceState) -> Result<()> {
696        // Ctx-holes fix: append the serial-decoded token's captured hidden.
697        // The decode layer loop (decode_a.rs try_dflash_capture) already
698        // filled `dflash_hidden_save` row 0 with this token's per-layer
699        // hiddens — the same [slot0|..|slot4] layout as one accumulator row.
700        let base = match self.dflash_hidden_save {
701            Some(p) => p,
702            None => return Ok(()),
703        };
704        let prop = match seq.proposer_state.as_mut() {
705            Some(p) => p.as_mut(),
706            None => return Ok(()),
707        };
708        // Graceful no-op for non-DFlash proposers (this bootstrap path is
709        // shared with EAGLE/MTP, unlike the DFlash-only eagle append above).
710        let d = match prop
711            .as_any_mut()
712            .downcast_mut::<crate::layers::DflashProposerState>()
713        {
714            Some(d) => d,
715            None => return Ok(()),
716        };
717        let n_layers = self.dflash_capture_layers.len();
718        if n_layers == 0 {
719            return Ok(());
720        }
721        let ctx_slot_bytes = n_layers * self.config.hidden_size * 2;
722        let stream = self.gpu.default_stream();
723        // Bounded watermark: accumulator full → slide the window. Keep the
724        // NEWEST keep = max/2 rows, drop the oldest (dropping the newest
725        // would starve the drafter of exactly the tokens that drive
726        // acceptance — the 846-token think overrun). drop_n >= keep holds
727        // whenever ctx_len >= max_ctx_len, so src/dst regions of the single
728        // D2D copy can never overlap — no ring arithmetic, no status-1.
729        // ctx_committed resets to 0: the next propose re-precomputes the
730        // slid rows chunk-wise (ctx_window rows/pass) and rewrites their
731        // paged K/V at the new slot indices; ctx_positions values (absolute
732        // positions) are preserved by the drain, so RoPE stamps stay exact.
733        if d.ctx_len >= d.max_ctx_len {
734            let keep = d.max_ctx_len / 2;
735            let drop_n = d.ctx_len - keep;
736            let src = d.ctx_hidden_acc.offset(drop_n * ctx_slot_bytes);
737            let dst0 = d.ctx_hidden_acc.offset(0);
738            self.gpu
739                .copy_d2d_async(src, dst0, keep * ctx_slot_bytes, stream)?;
740            d.ctx_positions.drain(..drop_n);
741            d.ctx_len = keep;
742            d.ctx_committed = 0;
743            tracing::info!(
744                "DFlash SERIAL_APPEND watermark: slid ctx window (dropped {} oldest, keep {})",
745                drop_n,
746                keep,
747            );
748        }
749        let dst = d.ctx_hidden_acc.offset(d.ctx_len * ctx_slot_bytes);
750        self.gpu.copy_d2d_async(base, dst, ctx_slot_bytes, stream)?;
751        // One-shot activation log so A/B runs can confirm the fix is live.
752        if self.stats.once("log:dflash_serial_append") {
753            tracing::info!(
754                "DFlash SERIAL_APPEND ACTIVE: first serial ctx append at ctx_len={} pos={}",
755                d.ctx_len,
756                seq.seq_len.saturating_sub(1),
757            );
758        }
759        // Position convention: decode() advanced seq_len past the token we
760        // just processed, so its true absolute position is seq_len - 1 —
761        // identical to propose.rs's `position.saturating_sub(1)` stamp.
762        debug_assert_eq!(d.ctx_positions.len(), d.ctx_len);
763        d.ctx_positions.push(seq.seq_len.saturating_sub(1) as i32);
764        d.ctx_len += 1;
765        // The latest capture is now in ctx; a propose() firing later (e.g.
766        // adaptive re-probe) must not decode-append it again.
767        d.skip_next_decode_append = true;
768        Ok(())
769    }
770    fn run_mtp_propose(
771        &self,
772        token: u32,
773        position: usize,
774        seq: &mut SequenceState,
775        _stream: u64,
776    ) -> Result<Option<u32>> {
777        self.run_mtp_propose_dispatch(token, position, seq, _stream)
778    }
779    fn run_mtp_propose_multi(
780        &self,
781        token: u32,
782        position: usize,
783        num_drafts: usize,
784        seq: &mut SequenceState,
785        _stream: u64,
786        grammar_bitmask: Option<&[i32]>,
787    ) -> Result<Vec<u32>> {
788        self.run_mtp_propose_multi_dispatch(
789            token,
790            position,
791            num_drafts,
792            seq,
793            _stream,
794            grammar_bitmask,
795        )
796    }
797    fn read_deferred_draft_token(&self) -> Result<u32> {
798        self.read_deferred_draft_token_dispatch()
799    }
800    fn trim_proposer_state(
801        &self,
802        seq: &mut SequenceState,
803        num_accepted: usize,
804        _stream: u64,
805    ) -> Result<()> {
806        self.trim_proposer_state_dispatch(seq, num_accepted, _stream)
807    }
808    fn compact_sequence(&self, seq: &mut SequenceState, new_slot: usize) -> Result<()> {
809        self.compact_sequence_dispatch(seq, new_slot)
810    }
811    fn detach_slot_for_reuse(&self, seq: &mut SequenceState) {
812        self.detach_slot_for_reuse_dispatch(seq)
813    }
814    fn save_sequence_state(
815        &self,
816        seq: &SequenceState,
817        writer: &mut dyn std::io::Write,
818    ) -> Result<()> {
819        self.save_sequence_state_dispatch(seq, writer)
820    }
821    fn restore_sequence_state(
822        &self,
823        seq: &mut SequenceState,
824        num_blocks: usize,
825        reader: &mut dyn std::io::Read,
826    ) -> Result<()> {
827        self.restore_sequence_state_dispatch(seq, num_blocks, reader)
828    }
829    fn num_free_blocks(&self) -> usize {
830        self.num_free_blocks_dispatch()
831    }
832    fn num_total_blocks(&self) -> usize {
833        self.num_total_blocks_dispatch()
834    }
835    fn reclaim_prefix_blocks(&self, num_blocks: usize) -> usize {
836        self.reclaim_prefix_blocks_dispatch(num_blocks)
837    }
838    fn start_checkpoint_async(&self, seq: &mut SequenceState) -> Result<()> {
839        self.start_checkpoint_async_dispatch(seq)
840    }
841    fn start_rollback_and_checkpoint_async(
842        &self,
843        seq: &mut SequenceState,
844        num_accepted: usize,
845    ) -> Result<()> {
846        self.start_rollback_and_checkpoint_async_dispatch(seq, num_accepted)
847    }
848    fn sync_secondary(&self) -> Result<()> {
849        self.sync_secondary_dispatch()
850    }
851    fn commit_accepted_prefix(
852        &self,
853        seq: &mut SequenceState,
854        num_accepted: usize,
855        k: usize,
856    ) -> Result<()> {
857        self.commit_accepted_prefix_dispatch(seq, num_accepted, k)
858    }
859    fn ep_worker_step(&self, slots: &mut [Option<SequenceState>]) -> Result<bool> {
860        self.ep_worker_step_dispatch(slots)
861    }
862    fn is_ep(&self) -> bool {
863        self.is_ep_dispatch()
864    }
865    fn hc_mult(&self) -> usize {
866        self.config.hc_mult
867    }
868
869    fn is_mla(&self) -> bool {
870        self.is_mla_dispatch()
871    }
872
873    fn kv_block_size(&self) -> Option<usize> {
874        Some(self.kv_cache.lock().block_size())
875    }
876    fn decode_logits_fp32(&self) -> bool {
877        self.decode_logits_fp32_dispatch()
878    }
879    fn decode_logits_ptr(&self) -> DevicePtr {
880        self.decode_logits_ptr_dispatch()
881    }
882    fn ep_broadcast_cmd(&self, cmd: u32) -> Result<()> {
883        self.ep_broadcast_cmd_dispatch(cmd)
884    }
885    fn ep_broadcast_cmd_for_seq(&self, seq_id: u32, cmd: u32) -> Result<()> {
886        // Routes to the helper added in 21e2130. Behaviour depends on the
887        // ep_protocol_v2 field set at construction from ATLAS_EP_PROTOCOL.
888        self.ep_broadcast_seq_and_cmd(seq_id, cmd, self.ep_protocol_v2)
889    }
890    fn ep_protocol_v2(&self) -> bool {
891        self.ep_protocol_v2
892    }
893    fn ep_broadcast_tokens(&self, tokens: &[u32]) -> Result<Vec<u32>> {
894        self.ep_broadcast_tokens_dispatch(tokens)
895    }
896    fn default_stream(&self) -> u64 {
897        self.default_stream_dispatch()
898    }
899    fn create_stream(&self) -> Result<u64> {
900        self.create_stream_dispatch()
901    }
902    fn create_event(&self) -> Result<u64> {
903        self.create_event_dispatch()
904    }
905    fn record_event(&self, event: u64, stream: u64) -> Result<()> {
906        self.record_event_dispatch(event, stream)
907    }
908    fn stream_wait_event(&self, stream: u64, event: u64) -> Result<()> {
909        self.stream_wait_event_dispatch(stream, event)
910    }
911    fn synchronize(&self, stream: u64) -> Result<()> {
912        self.synchronize_dispatch(stream)
913    }
914}
915
916impl TransformerModel {
917    /// Collect chunk-boundary aux layer state (PLE, QSA) for a Marconi
918    /// snapshot. Returns the blobs to attach; empty when no layer carries
919    /// aux state.
920    pub(in crate::model) fn collect_aux_states(
921        &self,
922        seq: &SequenceState,
923        stream: u64,
924    ) -> Result<Vec<(u32, Vec<u8>)>> {
925        let mut out = Vec::new();
926        for (i, l) in self.layers.iter().enumerate() {
927            if let Some(blob) =
928                l.snapshot_aux(seq.layer_states[i].as_ref(), self.gpu.as_ref(), stream)?
929            {
930                out.push((i as u32, blob));
931            }
932        }
933        Ok(out)
934    }
935
936    /// Whether restoring a snapshot WITHOUT aux blobs would be unsound for
937    /// this model (some layer carries per-sequence aux state).
938    pub(in crate::model) fn requires_aux_state(&self) -> bool {
939        self.layers.iter().any(|l| l.has_aux_state())
940    }
941
942    /// Apply a snapshot's aux blobs to the owning layers.
943    pub(in crate::model) fn apply_aux_states(
944        &self,
945        seq: &mut SequenceState,
946        blobs: &[(u32, Vec<u8>)],
947        stream: u64,
948    ) -> Result<()> {
949        for (i, blob) in blobs {
950            self.layers[*i as usize].restore_aux(
951                seq.layer_states[*i as usize].as_mut(),
952                blob,
953                self.gpu.as_ref(),
954                stream,
955            )?;
956        }
957        Ok(())
958    }
959}