use anyhow::Result;
use crate::inference::models::qwen35::kv_cache::HybridKvCache;
use crate::inference::models::qwen35::model::Qwen35Model;
use crate::serve::gpu::GpuContext;
use crate::serve::multi_seq_kv::SlotId;
use super::hidden_capture::DFlashCaptureSession;
use super::target::DFlashTarget;
pub struct Qwen35DFlashTarget<'a> {
pub model: &'a mut Qwen35Model,
pub kv_cache: &'a mut HybridKvCache,
pub dflash_capture: Option<DFlashCaptureSession>,
pub slot_id: SlotId,
}
impl<'a> Qwen35DFlashTarget<'a> {
pub fn new(model: &'a mut Qwen35Model, kv_cache: &'a mut HybridKvCache) -> Self {
Self::new_with_slot(model, kv_cache, SlotId(0))
}
pub fn new_with_slot(
model: &'a mut Qwen35Model,
kv_cache: &'a mut HybridKvCache,
slot_id: SlotId,
) -> Self {
Self {
model,
kv_cache,
dflash_capture: None,
slot_id,
}
}
pub fn with_slot_id(mut self, slot_id: SlotId) -> Self {
self.slot_id = slot_id;
self
}
}
impl<'a> DFlashTarget for Qwen35DFlashTarget<'a> {
fn install_dflash_capture(&mut self, session: DFlashCaptureSession) {
self.dflash_capture = Some(session);
}
fn take_dflash_capture(&mut self) -> Option<DFlashCaptureSession> {
self.dflash_capture.take()
}
fn has_dflash_capture(&self) -> bool {
self.dflash_capture.is_some()
}
fn rollback_kv(&mut self, trim: usize) {
if trim == 0 {
return;
}
let trim_u32 = trim as u32;
let slot = self.slot_id;
let cur_full = self
.kv_cache
.full_attn
.first()
.and_then(|s| s.current_len.get(slot.0 as usize).copied())
.unwrap_or(0);
let new_len_full = cur_full.saturating_sub(trim_u32);
if let Err(e) = self
.kv_cache
.truncate_full_attn_to_for_slot(slot, new_len_full)
{
eprintln!(
"[Qwen35DFlashTarget] truncate_full_attn_to_for_slot({:?}, {}) \
failed: {} — full-attn cursor may be stale by {} positions",
slot, new_len_full, e, trim
);
}
let cur_mtp = self
.kv_cache
.mtp_slot
.as_ref()
.and_then(|s| s.current_len.get(slot.0 as usize).copied())
.unwrap_or(0);
let new_len_mtp = cur_mtp.saturating_sub(trim_u32);
if let Err(e) = self.kv_cache.truncate_mtp_to_for_slot(slot, new_len_mtp) {
eprintln!(
"[Qwen35DFlashTarget] truncate_mtp_to_for_slot({:?}, {}) \
failed: {} — MTP cursor may be stale by {} positions",
slot, new_len_mtp, e, trim
);
}
if !self.kv_cache.linear_attn.is_empty() {
let first = &self.kv_cache.linear_attn[0];
if let Some(capture) = first.capture_states.as_ref() {
let recurrent_elems = first.recurrent.element_count();
if recurrent_elems > 0 {
let capture_elems = capture.element_count();
let n_tokens_max = capture_elems / recurrent_elems;
if (trim_u32 as usize) < n_tokens_max && n_tokens_max > 0 {
let accepted_idx = (n_tokens_max as u32) - 1 - trim_u32;
if let Err(e) = self.kv_cache.rollback_la_to(slot, accepted_idx) {
eprintln!(
"[Qwen35DFlashTarget] rollback_la_to({:?}, {}) failed: {} \
— LA state may be stale by {} positions",
slot, accepted_idx, e, trim
);
}
}
}
}
}
}
fn forward_decode_verify_batched(
&mut self,
tokens: &[u32],
start_seq_pos: usize,
_gpu: &mut GpuContext,
) -> Result<Vec<u32>> {
if tokens.is_empty() {
return Ok(Vec::new());
}
let positions_flat = crate::inference::models::qwen35::spec_decode::positions_for_range(
start_seq_pos as i32,
tokens.len(),
);
let seq_len = tokens.len();
let (logits, _hidden) = self.model.forward_gpu_with_hidden_dflash(
tokens,
&positions_flat,
self.kv_cache,
self.dflash_capture.as_mut(),
self.slot_id,
)?;
let vocab = self.model.cfg.vocab_size as usize;
if logits.len() != seq_len * vocab {
anyhow::bail!(
"forward_decode_verify_batched: expected logits len {} (seq_len={} × vocab={}), got {}",
seq_len * vocab,
seq_len,
vocab,
logits.len()
);
}
let mut argmaxes = Vec::with_capacity(seq_len);
for row in logits.chunks_exact(vocab) {
let mut best_idx = 0u32;
let mut best_val = f32::NEG_INFINITY;
for (i, &v) in row.iter().enumerate() {
if v > best_val {
best_val = v;
best_idx = i as u32;
}
}
argmaxes.push(best_idx);
}
crate::inference::spec_decode::emit_acceptance_metric(
crate::inference::spec_decode::SpecDecodeAcceptanceMetric::new(
self.slot_id,
seq_len as u32,
seq_len as u32,
0,
),
);
Ok(argmaxes)
}
}