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