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