1#![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 #[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 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 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 self.overlay_route_slot
169 .store(i32::MIN, std::sync::atomic::Ordering::Relaxed);
170 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 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 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 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 #[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 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 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 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 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 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 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 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 d.skip_next_decode_append = true;
673
674 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 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 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 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 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 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 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 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 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 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 pub(in crate::model) fn requires_aux_state(&self) -> bool {
939 self.layers.iter().any(|l| l.has_aux_state())
940 }
941
942 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}