use anyhow::{anyhow, Context};
use crate::inference::models::gemma4::MlxModelWeights;
use crate::inference::spec_decode::dflash::hidden_capture::{
extract_final_layer_slab, DFlashCaptureSession,
};
use crate::inference::spec_decode::ngram_proposer::{propose as ngram_propose, NgramConfig};
use crate::inference::spec_decode::verifier::accept_prefix_argmax;
use crate::serve::gpu::GpuContext;
#[derive(Debug, Default, Clone, Copy)]
pub struct NgramRoundProfile {
pub propose_ms: f64,
pub verify_prefill_ms: f64,
pub argmax_ms: f64,
pub total_ms: f64,
pub accept_count: usize,
pub draft_len: usize,
}
pub fn dispatch_ngram_generate(
target: &mut MlxModelWeights,
prompt_tokens: &[u32],
max_new_tokens: usize,
k: u32,
min_ngram: u32,
max_ngram: u32,
eos_token_ids: &[u32],
gpu: &mut GpuContext,
) -> anyhow::Result<Vec<u32>> {
if prompt_tokens.is_empty() {
anyhow::bail!("dispatch_ngram_generate: empty prompt");
}
if k < 1 {
anyhow::bail!("dispatch_ngram_generate: k must be >= 1; got {k}");
}
if min_ngram < 1 || max_ngram < min_ngram {
anyhow::bail!(
"dispatch_ngram_generate: invalid ngram range min={min_ngram} max={max_ngram}"
);
}
let profile_on = std::env::var("HF2Q_SPEC_NGRAM_PROFILE").as_deref() == Ok("1");
let hs = target.hidden_size;
let final_layer_idx = target.layers.len() - 1;
let combined_capture_ids: Vec<usize> = vec![final_layer_idx];
let max_decode_for_alloc = max_new_tokens;
let first_token = target
.forward_prefill_batched(prompt_tokens, max_decode_for_alloc, 0, gpu)
.map_err(|e| anyhow!("ngram: initial prefill: {e}"))?;
let mut output: Vec<u32> = prompt_tokens.to_vec();
output.push(first_token);
if eos_token_ids.contains(&first_token) || max_new_tokens <= 1 {
return Ok(output);
}
let mut total_propose_ms = 0.0;
let mut total_verify_ms = 0.0;
let mut total_argmax_ms = 0.0;
let mut total_round_ms = 0.0;
let mut total_drafts_proposed: usize = 0;
let mut total_drafts_accepted: usize = 0;
let mut total_rounds: usize = 0;
let mut total_empty_propose: usize = 0;
let cfg = NgramConfig {
min_ngram: min_ngram as usize,
max_ngram: max_ngram as usize,
k: k as usize,
max_model_len: prompt_tokens.len() + max_new_tokens,
};
while output.len() < prompt_tokens.len() + max_new_tokens {
let t_round = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let t_propose = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let drafts = ngram_propose(&output, &cfg);
if let Some(t) = t_propose {
total_propose_ms += t.elapsed().as_secs_f64() * 1000.0;
}
if drafts.is_empty() {
total_empty_propose += 1;
let last_tok = *output.last().expect("non-empty output");
let seq_pos = output.len() - 1;
let mut prof = None;
let next = target
.forward_decode(last_tok, seq_pos, gpu, &mut prof)
.map_err(|e| anyhow!("ngram: fallback decode: {e}"))?;
output.push(next);
if eos_token_ids.contains(&next) {
break;
}
if let Some(t) = t_round {
total_round_ms += t.elapsed().as_secs_f64() * 1000.0;
}
total_rounds += 1;
continue;
}
let last_tok = *output.last().expect("non-empty output");
let mut verify_input: Vec<u32> = Vec::with_capacity(1 + drafts.len());
verify_input.push(last_tok);
verify_input.extend(drafts.iter().copied());
let verify_seq_len = verify_input.len();
let xlen_sdpa = std::env::var("HF2Q_DFLASH_XLEN_SDPA").as_deref() == Ok("1");
let (verify_prefix, start_pos, xlen_max_decode) = if xlen_sdpa {
(verify_input, output.len() - 1, output.len() - 1)
} else {
let mut v = output.clone();
v.extend(drafts.iter().copied());
let len = v.len();
(v, 0usize, max_decode_for_alloc.min(len.saturating_sub(1)))
};
let verify_prefix_len = verify_prefix.len();
let session =
DFlashCaptureSession::new(combined_capture_ids.clone(), verify_prefix_len, hs, false);
target.install_dflash_capture(session);
let t_verify = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
target
.forward_prefill_batched(&verify_prefix, xlen_max_decode, start_pos, gpu)
.map_err(|e| anyhow!("ngram: verify forward: {e}"))?;
if let Some(t) = t_verify {
total_verify_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let captured = target
.take_dflash_capture()
.ok_or_else(|| anyhow!("ngram: verify capture vanished"))?;
let final_slab = extract_final_layer_slab(
&captured.hidden_output,
&combined_capture_ids,
final_layer_idx,
verify_prefix_len,
hs,
)
.context("ngram: extract final layer slab")?;
let t_argmax = if profile_on {
Some(std::time::Instant::now())
} else {
None
};
let argmaxes = if xlen_sdpa {
target
.per_position_argmax_from_hidden_batched_impl(
&final_slab,
verify_seq_len as u32,
true,
gpu,
)
.map_err(|e| anyhow!("ngram: argmax (xlen): {e}"))?
} else {
let tail_start = output.len() - 1;
let tail_slab: &[f32] = &final_slab[tail_start * hs..verify_prefix_len * hs];
target
.per_position_argmax_from_hidden_batched_impl(
tail_slab,
verify_seq_len as u32,
true,
gpu,
)
.map_err(|e| anyhow!("ngram: argmax (re-prefill): {e}"))?
};
if let Some(t) = t_argmax {
total_argmax_ms += t.elapsed().as_secs_f64() * 1000.0;
}
let (accept_count, fallback) = accept_prefix_argmax(&drafts, &argmaxes);
let drafts_len = drafts.len();
output.extend_from_slice(&drafts[..accept_count]);
output.push(fallback);
let n_rejected = drafts_len - accept_count;
if n_rejected > 0 {
target.rollback_kv(n_rejected);
}
if let Some(t) = t_round {
total_round_ms += t.elapsed().as_secs_f64() * 1000.0;
}
total_rounds += 1;
total_drafts_proposed += drafts_len;
total_drafts_accepted += accept_count;
for &tok in &drafts[..accept_count] {
if eos_token_ids.contains(&tok) {
return Ok(output);
}
}
if eos_token_ids.contains(&fallback) {
return Ok(output);
}
}
if profile_on {
let mean_accept_per_proposing_round = if total_rounds > total_empty_propose {
total_drafts_accepted as f64 / (total_rounds - total_empty_propose) as f64
} else {
0.0
};
let mean_accept_rate = if total_drafts_proposed > 0 {
total_drafts_accepted as f64 / total_drafts_proposed as f64
} else {
0.0
};
eprintln!(
"[HF2Q_SPEC_NGRAM_PROFILE] rounds={total_rounds} empty_propose={total_empty_propose} \
drafts_proposed={total_drafts_proposed} drafts_accepted={total_drafts_accepted} \
mean_accept_rate={:.3} mean_accept_per_proposing_round={:.2} \
cumulative_ms: propose={:.2} verify_prefill={:.2} argmax={:.2} total={:.2}",
mean_accept_rate,
mean_accept_per_proposing_round,
total_propose_ms,
total_verify_ms,
total_argmax_ms,
total_round_ms,
);
}
Ok(output)
}
#[cfg(test)]
mod tests {
}