use anyhow::{anyhow, Context, Result};
use mlx_native::DType;
use super::config::DFlashConfig;
use super::hidden_capture::{
append_capture_positions, extract_drafter_concat, DFlashCaptureSession,
};
use super::kv_cache::DFlashKvCache;
use super::orchestrator::step_round_from_argmaxes;
use super::qwen35_target::Qwen35DFlashTarget;
use super::target::DFlashTarget;
use super::tensors::DFlashModelTensors;
use crate::serve::gpu::GpuContext;
pub fn dispatch_qwen35_dflash_generate(
target: &mut Qwen35DFlashTarget<'_>,
drafter_tensors: &DFlashModelTensors,
drafter_cache: &mut DFlashKvCache,
drafter_cfg: &DFlashConfig,
prompt_tokens: &[u32],
max_new_tokens: usize,
block_size: u32,
eos_token_ids: &[u32],
gpu: &mut GpuContext,
) -> Result<Vec<u32>> {
if prompt_tokens.is_empty() {
anyhow::bail!("dispatch_qwen35_dflash_generate: empty prompt");
}
if block_size < 2 {
anyhow::bail!(
"dispatch_qwen35_dflash_generate: block_size must be >= 2 (got {})",
block_size,
);
}
if max_new_tokens == 0 {
return Ok(prompt_tokens.to_vec());
}
let hs = target.model.cfg.hidden_size as usize;
let final_layer_idx = target.model.layers.len() - 1;
let mut combined_capture_ids: Vec<usize> = drafter_cfg.target_layer_ids.clone();
if !combined_capture_ids.contains(&final_layer_idx) {
combined_capture_ids.push(final_layer_idx);
}
combined_capture_ids.sort_unstable();
combined_capture_ids.dedup();
let n_layers = target.model.layers.len();
for &lid in &combined_capture_ids {
if lid >= n_layers {
anyhow::bail!(
"dispatch_qwen35_dflash_generate: combined_capture_ids contains \
layer {} >= n_layers={}",
lid,
n_layers,
);
}
}
target
.model
.ensure_gpu_cache_primed()
.context("ensure_gpu_cache_primed")?;
let mut output: Vec<u32> = prompt_tokens.to_vec();
target.install_dflash_capture(DFlashCaptureSession::new(
combined_capture_ids.clone(),
prompt_tokens.len(),
hs,
false,
));
let initial_argmaxes = target
.forward_decode_verify_batched(prompt_tokens, 0, gpu)
.context("initial prefill forward")?;
let first_token = *initial_argmaxes
.last()
.ok_or_else(|| anyhow!("initial prefill: empty argmaxes"))?;
let mut prior_captured = target
.take_dflash_capture()
.ok_or_else(|| anyhow!("initial prefill: capture vanished"))?;
debug_assert_eq!(prior_captured.seq_len, prompt_tokens.len());
output.push(first_token);
if eos_token_ids.contains(&first_token) || max_new_tokens == 0 {
return Ok(output);
}
let mut last_token = first_token;
let drafter_target_layer_ids = drafter_cfg.target_layer_ids.clone();
let n_target_layers = drafter_target_layer_ids.len();
let row_stride = n_target_layers * hs;
target.model.with_gpu_cache_mut(|device, _reg| {
target
.kv_cache
.ensure_la_capture(&target.model.cfg, device, block_size)
})?;
let profile_on = std::env::var("HF2Q_DFLASH_PROFILE").as_deref() == Ok("1");
let mut rounds_count = 0usize;
let mut t_embed_ms = 0.0f64;
let mut t_extract_ms = 0.0f64;
let mut t_drafter_fwd_ms = 0.0f64;
let mut t_drafter_argmax_ms = 0.0f64;
let mut t_verify_ms = 0.0f64;
let mut t_trim_ms = 0.0f64;
while output.len() - prompt_tokens.len() < max_new_tokens {
rounds_count += 1;
let t0 = profile_on.then(std::time::Instant::now);
let mut block: Vec<u32> = Vec::with_capacity(block_size as usize);
block.push(last_token);
block.extend(std::iter::repeat(drafter_cfg.mask_token_id).take((block_size - 1) as usize));
let h = target
.model
.embed_tokens_gpu(&block)
.context("drafter embed")?;
if let Some(t) = t0 {
t_embed_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let prior_ctx_len = prior_captured.seq_len;
debug_assert_eq!(
prior_ctx_len,
output.len() - 1,
"prior_captured stale: seq_len={} but output.len()-1={}",
prior_ctx_len,
output.len() - 1,
);
let drafter_cached_seq_len = drafter_cache.layers[0].seq_len as usize;
debug_assert!(
prior_ctx_len >= drafter_cached_seq_len,
"drafter cache regressed: cached={} prior_ctx_len={}",
drafter_cached_seq_len,
prior_ctx_len,
);
let drafter_new_rows = prior_ctx_len - drafter_cached_seq_len;
let t0 = profile_on.then(std::time::Instant::now);
let drafter_concat_full = extract_drafter_concat(
&prior_captured.hidden_output,
&combined_capture_ids,
&drafter_target_layer_ids,
prior_ctx_len,
hs,
)?;
let new_rows_start = drafter_cached_seq_len * row_stride;
let drafter_concat_new: &[f32] = &drafter_concat_full[new_rows_start..];
debug_assert_eq!(
drafter_concat_new.len(),
drafter_new_rows * row_stride,
"drafter_concat_new length mismatch",
);
let target_hidden_concat = target.model.with_gpu_cache_mut(|device, _reg| {
let mut buf = device
.alloc_buffer(
drafter_concat_new.len() * 4,
DType::F32,
vec![drafter_new_rows.max(1), row_stride],
)
.map_err(|e| anyhow!("alloc target_hidden_concat: {e}"))?;
if drafter_new_rows > 0 {
buf.as_mut_slice::<f32>()
.map_err(|e| anyhow!("target_hidden_concat slice: {e}"))?
.copy_from_slice(drafter_concat_new);
}
Ok(buf)
})?;
if let Some(t) = t0 {
t_extract_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let t0 = profile_on.then(std::time::Instant::now);
let h_final = target
.model
.with_gpu_cache_mut(|device, registry| {
super::forward::dispatch_dflash_model_forward(
registry,
device,
&h,
&target_hidden_concat,
drafter_tensors,
drafter_cache,
drafter_cfg,
block_size,
drafter_new_rows as u32,
)
})
.context("drafter forward")?;
if let Some(t) = t0 {
t_drafter_fwd_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let t0 = profile_on.then(std::time::Instant::now);
let h_final_host: Vec<f32> = {
let slice = h_final
.as_slice::<f32>()
.map_err(|e| anyhow!("h_final slice: {e}"))?;
slice.to_vec()
};
let expected_h_final_len = (block_size as usize) * hs;
if h_final_host.len() != expected_h_final_len {
anyhow::bail!(
"drafter h_final length {} != block_size({}) * hs({}) = {}",
h_final_host.len(),
block_size,
hs,
expected_h_final_len,
);
}
let all_argmaxes = target
.model
.per_position_argmax_from_normed_hidden(&h_final_host, block_size)
.context("drafter argmax")?;
let drafts: Vec<u32> = all_argmaxes[1..].to_vec();
debug_assert_eq!(drafts.len(), (block_size - 1) as usize);
if let Some(t) = t0 {
t_drafter_argmax_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let t0 = profile_on.then(std::time::Instant::now);
let verify_seq_len = block_size as usize;
target.install_dflash_capture(DFlashCaptureSession::new(
combined_capture_ids.clone(),
verify_seq_len,
hs,
false,
));
let mut verify_input = Vec::with_capacity(verify_seq_len);
verify_input.push(last_token);
verify_input.extend(drafts.iter().copied());
let start_pos = output.len() - 1;
let target_argmaxes = target
.forward_decode_verify_batched(&verify_input, start_pos, gpu)
.context("verify forward")?;
let verify_captured = target
.take_dflash_capture()
.ok_or_else(|| anyhow!("verify capture vanished"))?;
if let Some(t) = t0 {
t_verify_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let round = step_round_from_argmaxes(&drafts, &target_argmaxes, eos_token_ids);
if profile_on {
eprintln!(
"[HF2Q_DFLASH_ACCEPT qwen35] round={} accept_count={}/{} \
drafts={:?} target_argmaxes={:?} committed={:?}",
rounds_count,
round.accept_count,
drafts.len(),
drafts,
target_argmaxes,
round.committed_tokens,
);
}
let rollback = drafts.len().saturating_sub(round.accept_count);
if rollback > 0 {
target.rollback_kv(rollback);
}
let n_committed = round.committed_tokens.len();
output.extend(round.committed_tokens.iter().copied());
last_token = *round
.committed_tokens
.last()
.ok_or_else(|| anyhow!("committed_tokens empty (round invariant violated)"))?;
if round.hit_eos {
break;
}
if output.len() - prompt_tokens.len() >= max_new_tokens {
break;
}
let t0 = profile_on.then(std::time::Instant::now);
prior_captured = append_capture_positions(&prior_captured, &verify_captured, n_committed)
.context("append accepted positions")?;
if let Some(t) = t0 {
t_trim_ms += t.elapsed().as_secs_f64() * 1000.0;
}
}
if profile_on && rounds_count > 0 {
let n = rounds_count as f64;
eprintln!(
"[HF2Q_DFLASH_PROFILE qwen35] rounds={} per-round-ms: \
embed={:.2} extract={:.2} drafter_fwd={:.2} \
drafter_argmax={:.2} verify={:.2} trim={:.2} TOTAL={:.2}",
rounds_count,
t_embed_ms / n,
t_extract_ms / n,
t_drafter_fwd_ms / n,
t_drafter_argmax_ms / n,
t_verify_ms / n,
t_trim_ms / n,
(t_embed_ms
+ t_extract_ms
+ t_drafter_fwd_ms
+ t_drafter_argmax_ms
+ t_verify_ms
+ t_trim_ms)
/ n,
);
}
let max_total = prompt_tokens.len() + max_new_tokens;
if output.len() > max_total {
output.truncate(max_total);
}
Ok(output)
}
#[cfg(test)]
mod tests {
}