use crate::inference::spec_decode::verifier::accept_prefix_argmax;
use anyhow::Context;
#[derive(Debug, Clone, PartialEq)]
pub struct RoundResult {
pub committed_tokens: Vec<u32>,
pub accept_count: usize,
pub hit_eos: bool,
}
pub fn step_round_from_argmaxes(
drafts: &[u32],
target_argmaxes: &[u32],
eos_token_ids: &[u32],
) -> RoundResult {
let (accept_count, model_token) = accept_prefix_argmax(drafts, target_argmaxes);
let mut effective_accept = accept_count;
let mut effective_final: u32 = model_token;
let mut hit_eos = false;
for (i, &t) in drafts.iter().take(accept_count).enumerate() {
if eos_token_ids.contains(&t) {
effective_accept = i;
effective_final = t;
hit_eos = true;
break;
}
}
if !hit_eos && eos_token_ids.contains(&model_token) {
hit_eos = true;
}
let mut committed = drafts[..effective_accept].to_vec();
committed.push(effective_final);
RoundResult {
committed_tokens: committed,
accept_count: effective_accept,
hit_eos,
}
}
pub fn dispatch_dflash_spec_decode_round_target_side<T: super::target::DFlashTarget>(
target: &mut T,
last_committed_token: u32,
drafts: &[u32],
current_seq_pos: usize,
eos_token_ids: &[u32],
gpu: &mut crate::serve::gpu::GpuContext,
) -> anyhow::Result<RoundResult> {
let mut verify_input = Vec::with_capacity(drafts.len() + 1);
verify_input.push(last_committed_token);
verify_input.extend_from_slice(drafts);
let argmaxes = target
.forward_decode_verify_batched(&verify_input, current_seq_pos, gpu)
.map_err(|e| anyhow::anyhow!("spec_decode_round: verify_batched: {e}"))?;
let round = step_round_from_argmaxes(drafts, &argmaxes, eos_token_ids);
let rollback = drafts.len().saturating_sub(round.accept_count);
if rollback > 0 {
target.rollback_kv(rollback);
}
Ok(round)
}
pub fn dispatch_dflash_one_round(
target: &mut crate::inference::models::gemma4::MlxModelWeights,
drafter_tensors: &super::tensors::DFlashModelTensors,
drafter_cache: &mut super::kv_cache::DFlashKvCache,
drafter_cfg: &super::config::DFlashConfig,
last_committed_token: u32,
target_hidden_concat: &mlx_native::MlxBuffer,
ctx_chunk_size: u32,
current_seq_pos: usize,
block_size: u32,
eos_token_ids: &[u32],
gpu: &mut crate::serve::gpu::GpuContext,
) -> anyhow::Result<RoundResult> {
if block_size < 2 {
anyhow::bail!("dispatch_dflash_one_round: block_size must be >= 2; got {block_size}");
}
let mut block: Vec<u32> = Vec::with_capacity(block_size as usize);
block.push(last_committed_token);
block.extend(std::iter::repeat(drafter_cfg.mask_token_id).take((block_size - 1) as usize));
let h = target
.embed_tokens(&block, gpu)
.map_err(|e| anyhow::anyhow!("one_round: embed_tokens: {e}"))?;
let h_final = {
let (exec, reg) = gpu.split();
let device = exec.device();
super::forward::dispatch_dflash_model_forward(
reg,
device,
&h,
target_hidden_concat,
drafter_tensors,
drafter_cache,
drafter_cfg,
block_size,
ctx_chunk_size,
)
.map_err(|e| anyhow::anyhow!("one_round: drafter forward: {e}"))?
};
let all_argmaxes: Vec<u32> = {
let h_final_slice: &[f32] = h_final
.as_slice::<f32>()
.map_err(|e| anyhow::anyhow!("one_round: h_final slice: {e}"))?;
let host_copy: Vec<f32> = h_final_slice.to_vec();
target
.per_position_argmax_from_hidden_opt(&host_copy, block_size, false, gpu)
.map_err(|e| anyhow::anyhow!("one_round: per_position argmax: {e}"))?
};
let drafts: Vec<u32> = all_argmaxes[1..].to_vec();
let round = dispatch_dflash_spec_decode_round_target_side(
target,
last_committed_token,
&drafts,
current_seq_pos,
eos_token_ids,
gpu,
)?;
Ok(round)
}
pub fn dispatch_dflash_generate_one_round_with_initial_capture(
target: &mut crate::inference::models::gemma4::MlxModelWeights,
drafter_tensors: &super::tensors::DFlashModelTensors,
drafter_cache: &mut super::kv_cache::DFlashKvCache,
drafter_cfg: &super::config::DFlashConfig,
prompt_tokens: &[u32],
block_size: u32,
eos_token_ids: &[u32],
gpu: &mut crate::serve::gpu::GpuContext,
) -> anyhow::Result<RoundResult> {
use super::hidden_capture::{DFlashCaptureSession, PrefillCapture};
let hs = target.hidden_size;
let num_target_layers = target.layers.len();
let mut capture_layer_ids: Vec<usize> = drafter_cfg.target_layer_ids.clone();
capture_layer_ids.sort_unstable();
capture_layer_ids.dedup();
for &i in &capture_layer_ids {
if i >= num_target_layers {
anyhow::bail!(
"generate_one_round: target_layer_id {} >= num_target_layers {}",
i,
num_target_layers
);
}
}
let session = DFlashCaptureSession::new(
capture_layer_ids.clone(),
prompt_tokens.len(),
hs,
false, );
target.install_dflash_capture(session);
let last_committed_token = target
.forward_prefill_batched(prompt_tokens, 0, 0, gpu)
.map_err(|e| anyhow::anyhow!("generate: initial prompt forward: {e}"))?;
let captured = target.take_dflash_capture().ok_or_else(|| {
anyhow::anyhow!("generate: capture session vanished after prompt forward")
})?;
let concat_vec: Vec<f32> = {
let view = PrefillCapture {
target_layer_ids: &captured.target_layer_ids,
hidden_output: &mut captured.hidden_output.clone(),
per_position_argmaxes: None,
};
view.permute_to_concat(prompt_tokens.len(), hs)
};
let target_hidden_concat = {
let (exec, _reg) = gpu.split();
let dev = exec.device();
let mut buf = dev
.alloc_buffer(
concat_vec.len() * 4,
mlx_native::DType::F32,
vec![prompt_tokens.len(), capture_layer_ids.len() * hs],
)
.map_err(|e| anyhow::anyhow!("generate: alloc target_hidden_concat: {e}"))?;
buf.as_mut_slice::<f32>()
.map_err(|e| anyhow::anyhow!("generate: target_hidden_concat slice: {e}"))?
.copy_from_slice(&concat_vec);
buf
};
let current_seq_pos = prompt_tokens.len();
let round = dispatch_dflash_one_round(
target,
drafter_tensors,
drafter_cache,
drafter_cfg,
last_committed_token,
&target_hidden_concat,
prompt_tokens.len() as u32,
current_seq_pos,
block_size,
eos_token_ids,
gpu,
)?;
Ok(round)
}
pub fn dispatch_dflash_generate(
target: &mut crate::inference::models::gemma4::MlxModelWeights,
drafter_tensors: &super::tensors::DFlashModelTensors,
drafter_cache: &mut super::kv_cache::DFlashKvCache,
drafter_cfg: &super::config::DFlashConfig,
prompt_tokens: &[u32],
max_new_tokens: usize,
block_size: u32,
eos_token_ids: &[u32],
gpu: &mut crate::serve::gpu::GpuContext,
) -> anyhow::Result<Vec<u32>> {
use super::hidden_capture::{
extract_drafter_concat, extract_final_layer_slab, DFlashCaptureSession,
};
if prompt_tokens.is_empty() {
anyhow::bail!("dispatch_dflash_generate: empty prompt");
}
if block_size < 2 {
anyhow::bail!("dispatch_dflash_generate: block_size must be >= 2");
}
let hs = target.hidden_size;
let final_layer_idx = target.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 mut output: Vec<u32> = prompt_tokens.to_vec();
let xlen_sdpa = std::env::var("HF2Q_DFLASH_XLEN_SDPA").as_deref() == Ok("1");
let session =
DFlashCaptureSession::new(combined_capture_ids.clone(), prompt_tokens.len(), hs, false);
target.install_dflash_capture(session);
let max_decode_for_alloc = max_new_tokens + block_size as usize - 1;
let first_token = target
.forward_prefill_batched(prompt_tokens, max_decode_for_alloc, 0, gpu)
.map_err(|e| anyhow::anyhow!("generate: initial prompt forward: {e}"))?;
let mut prior_captured = target
.take_dflash_capture()
.ok_or_else(|| anyhow::anyhow!("generate: initial 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 profile_on = std::env::var("HF2Q_DFLASH_PROFILE").as_deref() == Ok("1");
let mut t_embed_ms = 0.0f64;
let mut t_extract_concat_ms = 0.0f64;
let mut t_drafter_fwd_ms = 0.0f64;
let mut t_drafter_argmax_ms = 0.0f64;
let mut t_verify_prefill_ms = 0.0f64;
let mut t_target_argmax_ms = 0.0f64;
let mut t_trim_ms = 0.0f64;
let mut rounds_count = 0usize;
while output.len() - prompt_tokens.len() < max_new_tokens {
rounds_count += 1;
let t0_embed = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
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
.embed_tokens(&block, gpu)
.map_err(|e| anyhow::anyhow!("generate: embed_tokens: {e}"))?;
if let Some(t) = t0_embed {
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 t0_extract = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let drafter_cached_seq_len = drafter_cache.layers[0].seq_len as usize;
debug_assert!(
prior_ctx_len >= drafter_cached_seq_len,
"drafter cache state regressed: cached={} prior_ctx_len={}",
drafter_cached_seq_len,
prior_ctx_len,
);
let drafter_new_rows = prior_ctx_len - drafter_cached_seq_len;
let n_target_layers = drafter_cfg.target_layer_ids.len();
let row_stride = n_target_layers * hs;
let drafter_concat_vec_full = extract_drafter_concat(
&prior_captured.hidden_output,
&combined_capture_ids,
&drafter_cfg.target_layer_ids,
prior_ctx_len,
hs,
)?;
let new_rows_start = drafter_cached_seq_len * row_stride;
let drafter_concat_new: &[f32] = &drafter_concat_vec_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 = {
let (exec, _reg) = gpu.split();
let dev = exec.device();
let mut buf = dev
.alloc_buffer(
drafter_concat_new.len() * 4,
mlx_native::DType::F32,
vec![drafter_new_rows.max(1), row_stride],
)
.map_err(|e| anyhow::anyhow!("generate: alloc target_hidden_concat: {e}"))?;
if drafter_new_rows > 0 {
buf.as_mut_slice::<f32>()
.map_err(|e| anyhow::anyhow!("generate: target_hidden_concat slice: {e}"))?
.copy_from_slice(drafter_concat_new);
}
buf
};
if let Some(t) = t0_extract {
t_extract_concat_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let t0_drafter = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let h_final = {
let (exec, reg) = gpu.split();
let device = exec.device();
super::forward::dispatch_dflash_model_forward(
reg,
device,
&h,
&target_hidden_concat,
drafter_tensors,
drafter_cache,
drafter_cfg,
block_size,
drafter_new_rows as u32,
)
.context("generate: drafter forward")?
};
if let Some(t) = t0_drafter {
t_drafter_fwd_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let t0_drafter_argmax = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let drafts: Vec<u32> = {
let h_final_slice: &[f32] = h_final
.as_slice::<f32>()
.map_err(|e| anyhow::anyhow!("generate: h_final slice: {e}"))?;
let host_copy: Vec<f32> = h_final_slice.to_vec();
if std::env::var("HF2Q_DFLASH_DRAFTER_DUMP").as_deref() == Ok("1") {
let bs = block_size as usize;
let mut max_pairwise_diff = 0.0f32;
for i in 1..bs {
let row_i = &host_copy[i * hs..(i + 1) * hs];
let row_im1 = &host_copy[(i - 1) * hs..i * hs];
let diff = row_i
.iter()
.zip(row_im1.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
max_pairwise_diff = max_pairwise_diff.max(diff);
}
let nan_count = host_copy.iter().filter(|x| x.is_nan()).count();
let inf_count = host_copy.iter().filter(|x| x.is_infinite()).count();
let mut row_max_abs = Vec::with_capacity(bs);
for i in 0..bs {
let row = &host_copy[i * hs..(i + 1) * hs];
let m = row
.iter()
.filter(|x| x.is_finite())
.map(|x| x.abs())
.fold(0.0f32, f32::max);
row_max_abs.push(m);
}
eprintln!(
"[DRAFTER_DUMP round={} block_size={} hs={}]\n \
h_final[pos=0,d=0..8] = {:?}\n \
h_final[pos=1,d=0..8] = {:?}\n \
h_final[pos={},d=0..8] = {:?}\n \
max_adj_pairwise_abs_diff(rows 1..{}) = {:.6e}\n \
nan_count = {} inf_count = {}\n \
per_row_max_abs = {:?}",
rounds_count,
bs,
hs,
&host_copy[0..8.min(hs)],
&host_copy[hs..hs + 8.min(hs)],
bs - 1,
&host_copy[(bs - 1) * hs..(bs - 1) * hs + 8.min(hs)],
bs,
max_pairwise_diff,
nan_count,
inf_count,
row_max_abs,
);
}
let all_argmaxes = target
.per_position_argmax_from_hidden_batched_impl(&host_copy, block_size, false, gpu)
.map_err(|e| anyhow::anyhow!("generate: drafter argmax: {e}"))?;
all_argmaxes[1..].to_vec()
};
if let Some(t) = t0_drafter_argmax {
t_drafter_argmax_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let (verify_captured, target_argmaxes, verify_seq_len_for_path) = if xlen_sdpa {
let verify_input: Vec<u32> = std::iter::once(last_token)
.chain(drafts.iter().copied())
.collect();
let verify_seq_len = verify_input.len(); let start_pos = output.len() - 1;
let verify_session =
DFlashCaptureSession::new(combined_capture_ids.clone(), verify_seq_len, hs, false);
target.install_dflash_capture(verify_session);
let t0_verify = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let xlen_max_decode = start_pos;
let _verify_last_argmax = target
.forward_prefill_batched(&verify_input, xlen_max_decode, start_pos, gpu)
.map_err(|e| anyhow::anyhow!("generate: verify forward (xlen): {e}"))?;
let captured = target
.take_dflash_capture()
.ok_or_else(|| anyhow::anyhow!("generate: verify capture vanished (xlen)"))?;
if let Some(t) = t0_verify {
t_verify_prefill_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let t0_target_argmax = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let final_slab = extract_final_layer_slab(
&captured.hidden_output,
&combined_capture_ids,
final_layer_idx,
verify_seq_len,
hs,
)?;
let argmaxes = target
.per_position_argmax_from_hidden_batched_impl(
&final_slab,
verify_seq_len as u32,
true,
gpu,
)
.map_err(|e| anyhow::anyhow!("generate: target argmax (xlen): {e}"))?;
if let Some(t) = t0_target_argmax {
t_target_argmax_ms += t.elapsed().as_secs_f64() * 1000.0;
}
(captured, argmaxes, verify_seq_len)
} else {
let mut verify_prefix: Vec<u32> = output.clone();
verify_prefix.extend(drafts.iter().copied());
let verify_prefix_len = verify_prefix.len();
let verify_session = DFlashCaptureSession::new(
combined_capture_ids.clone(),
verify_prefix_len,
hs,
false,
);
target.install_dflash_capture(verify_session);
let t0_verify = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let _verify_last_argmax = target
.forward_prefill_batched(&verify_prefix, max_decode_for_alloc, 0, gpu)
.map_err(|e| anyhow::anyhow!("generate: verify forward: {e}"))?;
let captured = target
.take_dflash_capture()
.ok_or_else(|| anyhow::anyhow!("generate: verify capture vanished"))?;
if let Some(t) = t0_verify {
t_verify_prefill_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let t0_target_argmax = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let final_slab = extract_final_layer_slab(
&captured.hidden_output,
&combined_capture_ids,
final_layer_idx,
verify_prefix_len,
hs,
)?;
let verify_start = output.len() - 1;
let verify_end = verify_start + block_size as usize;
debug_assert_eq!(verify_end, verify_prefix_len);
let verify_slab_tail: &[f32] = &final_slab[verify_start * hs..verify_end * hs];
let argmaxes = target
.per_position_argmax_from_hidden_batched_impl(
verify_slab_tail,
block_size,
true,
gpu,
)
.map_err(|e| anyhow::anyhow!("generate: target argmax: {e}"))?;
if let Some(t) = t0_target_argmax {
t_target_argmax_ms += t.elapsed().as_secs_f64() * 1000.0;
}
(captured, argmaxes, verify_prefix_len)
};
let round = step_round_from_argmaxes(&drafts, &target_argmaxes, eos_token_ids);
if std::env::var("HF2Q_DFLASH_PROFILE").as_deref() == Ok("1") {
eprintln!(
"[HF2Q_DFLASH_ACCEPT] round={rounds_count} accept_count={}/{} \
drafts={drafts:?} target_argmaxes={target_argmaxes:?} \
committed={:?}",
round.accept_count,
drafts.len(),
round.committed_tokens,
);
}
if std::env::var("HF2Q_DFLASH_HIDDEN_DEBUG").as_deref() == Ok("1") {
let final_combined_idx = combined_capture_ids
.iter()
.position(|&c| c == final_layer_idx)
.expect("final layer in combined capture set");
let path = if xlen_sdpa { "OptA" } else { "OptC" };
let logical_pos = output.len() - 1; let capture_row = if xlen_sdpa { 0 } else { logical_pos };
let capture_seq_len = verify_captured.seq_len;
let layer_base = final_combined_idx * capture_seq_len * hs;
let row_base = layer_base + capture_row * hs;
let row_slice = &verify_captured.hidden_output[row_base..row_base + 8.min(hs)];
let max_abs = verify_captured.hidden_output
[layer_base..layer_base + capture_seq_len * hs]
.iter()
.fold(0.0f32, |a, &b| a.max(b.abs()));
eprintln!(
"[HIDDEN_DEBUG path={} round={} logical_pos={} capture_row={} \
hidden_final_layer[d=0..8]={:?} max_abs={:.4e}",
path, rounds_count, logical_pos, capture_row, row_slice, max_abs,
);
}
if xlen_sdpa {
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().unwrap();
if round.hit_eos {
break;
}
if output.len() - prompt_tokens.len() >= max_new_tokens {
break;
}
let t0_trim = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
if xlen_sdpa {
prior_captured = super::hidden_capture::append_capture_positions(
&prior_captured,
&verify_captured,
n_committed,
)?;
} else {
let mut next_captured = verify_captured;
let next_prior_ctx_len = output.len() - 1;
debug_assert!(
next_prior_ctx_len <= next_captured.seq_len,
"trim target {} > current {}",
next_prior_ctx_len,
next_captured.seq_len
);
super::hidden_capture::trim_capture_to(&mut next_captured, next_prior_ctx_len);
prior_captured = next_captured;
}
if let Some(t) = t0_trim {
t_trim_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let _ = verify_seq_len_for_path; }
if profile_on && rounds_count > 0 {
let n = rounds_count as f64;
eprintln!(
"[HF2Q_DFLASH_PROFILE] rounds={} per-round-ms: embed={:.2} extract={:.2} drafter_fwd={:.2} drafter_argmax={:.2} verify_prefill={:.2} target_argmax={:.2} trim={:.2} TOTAL={:.2}",
rounds_count,
t_embed_ms / n,
t_extract_concat_ms / n,
t_drafter_fwd_ms / n,
t_drafter_argmax_ms / n,
t_verify_prefill_ms / n,
t_target_argmax_ms / n,
t_trim_ms / n,
(t_embed_ms + t_extract_concat_ms + t_drafter_fwd_ms + t_drafter_argmax_ms
+ t_verify_prefill_ms + t_target_argmax_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 {
use super::*;
use crate::inference::spec_decode::dflash::{
config::DFlashConfig,
forward::dispatch_dflash_model_forward,
kv_cache::DFlashKvCache,
tensors::DFlashModelTensors,
weights::{DFlashWeights, DFlashWeightsFile},
};
use mlx_native::{DType, KernelRegistry, MlxDevice};
#[test]
fn k0_empty_drafts_degrades_to_single_token() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafts: Vec<u32> = vec![];
let target_argmaxes: Vec<u32> = vec![42];
let result = step_round_from_argmaxes(&drafts, &target_argmaxes, &[]);
assert_eq!(result.committed_tokens, vec![42]);
assert_eq!(result.accept_count, 0);
assert!(!result.hit_eos);
}
#[test]
fn full_accept_returns_all_drafts_plus_model_token() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafts = vec![10, 20, 30];
let target = vec![10, 20, 30, 99];
let result = step_round_from_argmaxes(&drafts, &target, &[]);
assert_eq!(result.committed_tokens, vec![10, 20, 30, 99]);
assert_eq!(result.accept_count, 3);
assert!(!result.hit_eos);
}
#[test]
fn partial_accept_truncates_at_first_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafts = vec![10, 20, 30, 40];
let target = vec![10, 20, 88, 99]; let result = step_round_from_argmaxes(&drafts, &target, &[]);
assert_eq!(result.committed_tokens, vec![10, 20, 88]);
assert_eq!(result.accept_count, 2);
assert!(!result.hit_eos);
}
#[test]
fn eos_in_accepted_prefix_stops_generation() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafts = vec![10, 7, 30];
let target = vec![10, 7, 30];
let result = step_round_from_argmaxes(&drafts, &target, &[7]);
assert_eq!(result.committed_tokens, vec![10, 7]);
assert_eq!(result.accept_count, 1); assert!(result.hit_eos);
}
#[test]
fn eos_as_model_free_token_sets_hit_eos() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafts = vec![10, 20];
let target = vec![10, 1, 30];
let result = step_round_from_argmaxes(&drafts, &target, &[1]);
assert_eq!(result.committed_tokens, vec![10, 1]);
assert_eq!(result.accept_count, 1);
assert!(result.hit_eos);
}
#[test]
fn eos_check_handles_full_accept_with_eos_continuation() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafts = vec![10, 20];
let target = vec![10, 20, 1];
let result = step_round_from_argmaxes(&drafts, &target, &[1]);
assert_eq!(result.committed_tokens, vec![10, 20, 1]);
assert_eq!(result.accept_count, 2);
assert!(result.hit_eos);
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_orchestrator_drafter_loop_with_simulated_target() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = DFlashConfig::from_json_str(
crate::inference::spec_decode::dflash::config::tests::GEMMA4_26B_A4B_DFLASH_CONFIG,
)
.expect("config parse");
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let mut cache = DFlashKvCache::new(&device, &cfg, 128).expect("cache");
let block_size = 8u32;
let ctx_chunk = 4u32;
let hidden = cfg.hidden_size as u32;
let fc_in = cfg.fc_input_dim() as u32;
let h_elem = (block_size as usize) * (hidden as usize);
let mut h = device
.alloc_buffer(
h_elem * 4,
DType::F32,
vec![block_size as usize, hidden as usize],
)
.expect("alloc h");
{
let s = h.as_mut_slice::<f32>().expect("h slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let thc_elem = (ctx_chunk as usize) * (fc_in as usize);
let mut target_hidden = device
.alloc_buffer(
thc_elem * 4,
DType::F32,
vec![ctx_chunk as usize, fc_in as usize],
)
.expect("alloc target_hidden");
{
let s = target_hidden
.as_mut_slice::<f32>()
.expect("target_hidden slice");
for (i, v) in s.iter_mut().enumerate() {
*v = 0.1 + ((i % 17) as f32) / 170.0;
}
}
let h_final = dispatch_dflash_model_forward(
&mut registry,
&device,
&h,
&target_hidden,
&tensors,
&mut cache,
&cfg,
block_size,
ctx_chunk,
)
.expect("drafter forward");
assert_eq!(
h_final.element_count(),
(block_size as usize) * (hidden as usize)
);
let drafts: Vec<u32> = (1..=7).collect();
let target_argmaxes = vec![1, 2, 3, 4, 99, 100, 101, 102];
let round = step_round_from_argmaxes(&drafts, &target_argmaxes, &[]);
assert_eq!(round.accept_count, 4);
assert_eq!(round.committed_tokens, vec![1, 2, 3, 4, 99]);
assert!(!round.hit_eos);
for (i, l) in cache.layers.iter().enumerate() {
assert_eq!(
l.seq_len, ctx_chunk,
"layer {i} cache should advance by ctx_chunk"
);
}
}
#[test]
#[ignore = "requires gemma-4-26b GGUF + DFlash drafter HF cache + ~22GB RAM"]
fn e2e_dispatch_dflash_generate_gemma4_26b() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::models::gemma4::MlxModelWeights;
use crate::inference::spec_decode::dflash::{
kv_cache::DFlashKvCache,
tensors::DFlashModelTensors,
weights::{DFlashWeights, DFlashWeightsFile},
};
use crate::serve::{config::Gemma4Config, gpu::GpuContext, header::LoadProgress};
use std::path::PathBuf;
let target_gguf = PathBuf::from(
"/opt/hf2q/models/gemma-4-26b-a4b-it-ara-abliterated/\
gemma4-ara-2pass-APEX-Q5_K_M.gguf",
);
let tokenizer_path =
PathBuf::from("/opt/hf2q/models/gemma-4-26b-a4b-it-ara-abliterated/tokenizer.json");
let home = std::env::var("HOME").expect("HOME env set");
let drafter_dir = format!(
"{home}/.cache/huggingface/hub/\
models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/\
77d4202772dfe50b2396ec7bac9cfffc7b9e7057"
);
let drafter_cfg_path = format!("{drafter_dir}/config.json");
let drafter_safetensors_path = format!("{drafter_dir}/model.safetensors");
for p in [
&target_gguf,
&tokenizer_path,
&PathBuf::from(&drafter_cfg_path),
&PathBuf::from(&drafter_safetensors_path),
] {
if !p.exists() {
panic!(
"required artifact missing: {} — see test #[ignore] note",
p.display()
);
}
}
let mut gpu = GpuContext::new().expect("Metal device available");
let gguf = mlx_native::gguf::GgufFile::open(&target_gguf).expect("open target GGUF");
let target_cfg = Gemma4Config::from_gguf(&gguf).expect("gemma4 cfg from gguf");
let mut progress = LoadProgress::new(false, 0, 0);
let mut target =
MlxModelWeights::load_from_gguf(&gguf, &target_cfg, &mut gpu, &mut progress)
.expect("load target weights from GGUF");
let drafter_cfg =
DFlashConfig::from_json_path(&drafter_cfg_path).expect("drafter config.json");
let drafter_file =
DFlashWeightsFile::open(&drafter_safetensors_path).expect("drafter safetensors open");
let drafter_weights = DFlashWeights::load(drafter_file.bytes(), &drafter_cfg)
.expect("drafter validated load");
let drafter_tensors = {
let (exec, _reg) = gpu.split();
DFlashModelTensors::upload(exec.device(), &drafter_cfg, &drafter_weights)
.expect("drafter GPU upload")
};
let drafter_cache_cap: u32 = 4096;
let mut drafter_cache = {
let (exec, _reg) = gpu.split();
DFlashKvCache::new(exec.device(), &drafter_cfg, drafter_cache_cap)
.expect("drafter cache alloc")
};
let tokenizer =
tokenizers::Tokenizer::from_file(&tokenizer_path).expect("load tokenizer.json");
let prompt_text =
std::env::var("HF2Q_TEST_PROMPT").unwrap_or_else(|_| "Q: What is 2+2?\nA:".to_string());
let encoding = tokenizer
.encode(prompt_text.as_str(), false)
.expect("encode");
let prompt_tokens: Vec<u32> = encoding.get_ids().to_vec();
assert!(!prompt_tokens.is_empty(), "prompt encoding empty");
eprintln!(
"[e2e] prompt={prompt_text:?} prompt_tokens.len()={}",
prompt_tokens.len()
);
let prompt_len = prompt_tokens.len();
let max_new_tokens = 16usize; let block_size = 8u32;
eprintln!("[e2e] BASELINE: single-token decode for N={max_new_tokens} tokens");
let t_baseline = std::time::Instant::now();
let max_decode_for_alloc = max_new_tokens + block_size as usize - 1;
let first_token_baseline = target
.forward_prefill_batched(&prompt_tokens, max_decode_for_alloc, 0, &mut gpu)
.expect("baseline initial prefill");
let mut baseline_new: Vec<u32> = vec![first_token_baseline];
let mut last_tok = first_token_baseline;
for step in 0..(max_new_tokens - 1) {
let seq_pos = prompt_len + step;
let mut prof: Option<crate::inference::models::gemma4::TokenProfile> = None;
let next = target
.forward_decode(last_tok, seq_pos, &mut gpu, &mut prof)
.expect("baseline forward_decode");
baseline_new.push(next);
last_tok = next;
}
let baseline_elapsed = t_baseline.elapsed();
eprintln!(
"[e2e] BASELINE done: {:.2}s, tokens={baseline_new:?}",
baseline_elapsed.as_secs_f64(),
);
let rollback_count = prompt_len + max_new_tokens - 1;
target.rollback_kv(rollback_count);
eprintln!("[e2e] target.rollback_kv({rollback_count}) → seq_len=0");
let eos_token_ids: Vec<u32> = vec![1, 106];
let t_spec = std::time::Instant::now();
let spec_output = match dispatch_dflash_generate(
&mut target,
&drafter_tensors,
&mut drafter_cache,
&drafter_cfg,
&prompt_tokens,
max_new_tokens,
block_size,
&eos_token_ids,
&mut gpu,
) {
Ok(toks) => toks,
Err(e) => {
eprintln!("[e2e] dispatch_dflash_generate FAILED chain:");
for (i, cause) in e.chain().enumerate() {
eprintln!("[e2e] #{i}: {cause}");
}
panic!("dispatch_dflash_generate end-to-end (see chain above)");
}
};
let spec_elapsed = t_spec.elapsed();
eprintln!(
"[e2e] SPEC done: {:.2}s, output.len()={} (prompt={prompt_len}, new<={max_new_tokens})",
spec_elapsed.as_secs_f64(),
spec_output.len(),
);
assert!(
spec_output.len() > prompt_len,
"spec output must emit ≥ 1 new token; got len={}",
spec_output.len()
);
let spec_new = &spec_output[prompt_len..];
assert!(
spec_new.len() <= max_new_tokens,
"spec must not exceed max_new_tokens={max_new_tokens}; got {}",
spec_new.len()
);
let n_compare = baseline_new.len().min(spec_new.len());
eprintln!("[e2e] baseline_new = {baseline_new:?}");
eprintln!("[e2e] spec_new = {spec_new:?}");
for i in 0..n_compare {
let mark = if baseline_new[i] == spec_new[i] {
"✓"
} else {
"✗"
};
eprintln!(
"[e2e] pos {i}: baseline={} spec={} {mark}",
baseline_new[i], spec_new[i]
);
}
assert_eq!(
spec_new.len(),
baseline_new.len(),
"spec emitted {} tokens, baseline emitted {} — length mismatch",
spec_new.len(),
baseline_new.len(),
);
for (i, (b, s)) in baseline_new.iter().zip(spec_new.iter()).enumerate() {
assert_eq!(
b, s,
"coherence gate FAILED at new-token position {i}: baseline={b} spec={s} \
(first {i} tokens matched). Investigate orchestrator at this round.",
);
}
let decoded = tokenizer
.decode(spec_new, false)
.unwrap_or_else(|e| format!("<decode failed: {e}>"));
eprintln!("[e2e] COHERENCE PASS: spec_new == baseline_new for all {n_compare} tokens");
eprintln!("[e2e] decoded = {decoded:?}");
}
#[test]
#[ignore = "iter-67 dual-axis diagnostic; requires gemma-4-26b GGUF + DFlash drafter"]
fn e2e_coherence_gemma4_chat_templated_prompt() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::models::gemma4::MlxModelWeights;
use crate::inference::spec_decode::dflash::{
kv_cache::DFlashKvCache,
tensors::DFlashModelTensors,
weights::{DFlashWeights, DFlashWeightsFile},
};
use crate::serve::{config::Gemma4Config, gpu::GpuContext, header::LoadProgress};
use std::path::PathBuf;
let target_gguf = PathBuf::from(
"/opt/hf2q/models/gemma-4-26b-a4b-it-ara-abliterated/\
gemma4-ara-2pass-APEX-Q5_K_M.gguf",
);
let home = std::env::var("HOME").expect("HOME env set");
let drafter_dir = format!(
"{home}/.cache/huggingface/hub/\
models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/\
77d4202772dfe50b2396ec7bac9cfffc7b9e7057"
);
let drafter_cfg_path = format!("{drafter_dir}/config.json");
let drafter_safetensors_path = format!("{drafter_dir}/model.safetensors");
for p in [
target_gguf.to_string_lossy().to_string(),
drafter_cfg_path.clone(),
drafter_safetensors_path.clone(),
] {
if !std::path::Path::new(&p).exists() {
panic!("required artifact missing: {p}");
}
}
let mut gpu = GpuContext::new().expect("Metal device");
let gguf = mlx_native::gguf::GgufFile::open(&target_gguf).expect("open gguf");
let target_cfg = Gemma4Config::from_gguf(&gguf).expect("gemma4 cfg");
let mut progress = LoadProgress::new(false, 0, 0);
let mut target =
MlxModelWeights::load_from_gguf(&gguf, &target_cfg, &mut gpu, &mut progress)
.expect("load target");
let drafter_cfg = DFlashConfig::from_json_path(&drafter_cfg_path).expect("drafter cfg");
let drafter_file =
DFlashWeightsFile::open(&drafter_safetensors_path).expect("drafter file");
let drafter_weights =
DFlashWeights::load(drafter_file.bytes(), &drafter_cfg).expect("drafter weights");
let drafter_tensors = {
let (exec, _reg) = gpu.split();
DFlashModelTensors::upload(exec.device(), &drafter_cfg, &drafter_weights)
.expect("drafter upload")
};
let drafter_cache_cap: u32 = 4096;
let mut drafter_cache = {
let (exec, _reg) = gpu.split();
DFlashKvCache::new(exec.device(), &drafter_cfg, drafter_cache_cap)
.expect("drafter cache")
};
let prompt_tokens: Vec<u32> = vec![
2, 105, 2364, 107, 236935, 236787, 2900, 563, 236743, 236778, 236862, 236778, 105470,
169631, 236787, 106, 107, 105, 4368, 107, 100, 45518, 107, 101,
];
let prompt_len = prompt_tokens.len();
let max_new_tokens = 8usize;
let block_size = 8u32;
eprintln!("[e2e-tmpl] prompt_tokens.len()={prompt_len}");
let max_decode_for_alloc = max_new_tokens + block_size as usize - 1;
let first_token_baseline = target
.forward_prefill_batched(&prompt_tokens, max_decode_for_alloc, 0, &mut gpu)
.expect("baseline prefill");
let mut baseline_new: Vec<u32> = vec![first_token_baseline];
let mut last_tok = first_token_baseline;
for step in 0..(max_new_tokens - 1) {
let mut prof: Option<crate::inference::models::gemma4::TokenProfile> = None;
let next = target
.forward_decode(last_tok, prompt_len + step, &mut gpu, &mut prof)
.expect("baseline forward_decode");
baseline_new.push(next);
last_tok = next;
}
eprintln!("[e2e-tmpl] BASELINE = {baseline_new:?}");
target.rollback_kv(prompt_len + max_new_tokens - 1);
eprintln!("[e2e-tmpl] DIAG: forward_prefill_batched at incremental prefix lengths");
for i in 0..max_new_tokens {
let mut prefix: Vec<u32> = prompt_tokens.clone();
prefix.extend(&baseline_new[..i]);
let argmax = target
.forward_prefill_batched(&prefix, max_decode_for_alloc, 0, &mut gpu)
.expect("diag forward_prefill_batched");
let mark = if argmax == baseline_new[i] {
"✓"
} else {
"✗"
};
eprintln!(
"[e2e-tmpl] DIAG L={} argmax={} baseline_new[{}]={} {}",
prefix.len(),
argmax,
i,
baseline_new[i],
mark,
);
target.rollback_kv(prefix.len());
}
let eos_token_ids: Vec<u32> = vec![]; let spec_output = dispatch_dflash_generate(
&mut target,
&drafter_tensors,
&mut drafter_cache,
&drafter_cfg,
&prompt_tokens,
max_new_tokens,
block_size,
&eos_token_ids,
&mut gpu,
)
.expect("dispatch_dflash_generate");
let spec_new = &spec_output[prompt_len..];
eprintln!("[e2e-tmpl] SPEC = {spec_new:?}");
let n_compare = baseline_new.len().min(spec_new.len());
for i in 0..n_compare {
let mark = if baseline_new[i] == spec_new[i] {
"✓"
} else {
"✗"
};
eprintln!(
"[e2e-tmpl] pos {i}: baseline={} spec={} {mark}",
baseline_new[i], spec_new[i]
);
}
eprintln!("[e2e-tmpl] DIAG: forward_prefill_batched on SPEC's own chain");
let mut spec_self_consistent = true;
target.rollback_kv(prompt_len + max_new_tokens - 1);
for i in 0..max_new_tokens {
let mut prefix: Vec<u32> = prompt_tokens.clone();
prefix.extend(&spec_new[..i]);
let argmax = target
.forward_prefill_batched(&prefix, max_decode_for_alloc, 0, &mut gpu)
.expect("self-consistency prefill");
let mark = if argmax == spec_new[i] { "✓" } else { "✗" };
if argmax != spec_new[i] {
spec_self_consistent = false;
}
eprintln!(
"[e2e-tmpl] SELF L={} argmax={} spec_new[{}]={} {}",
prefix.len(),
argmax,
i,
spec_new[i],
mark,
);
target.rollback_kv(prefix.len());
}
assert!(
spec_self_consistent,
"ORCHESTRATOR self-consistency FAILED: spec_new should match \
forward_prefill_batched(prompt + spec_new[..i]).first_token at \
every i (see SELF rows above for first mismatch).",
);
eprintln!(
"[e2e-tmpl] ORCHESTRATOR self-consistency PASS (spec-decode faithful to batched_prefill)"
);
assert_eq!(spec_new.len(), baseline_new.len(), "length mismatch");
for (i, (b, s)) in baseline_new.iter().zip(spec_new.iter()).enumerate() {
assert_eq!(
b, s,
"coherence gate FAILED on chat-templated prompt at new-token position {i}: \
baseline={b} spec={s} — root cause is forward_prefill_batched coherence \
(see DIAG output for axis-2 failures)."
);
}
eprintln!("[e2e-tmpl] FULL COHERENCE PASS for chat-templated prompt");
}
}