use anyhow::{anyhow, ensure, Context, Result};
use mlx_native::{DType, MlxBuffer, MlxDevice};
use crate::core::traits::activation_capture::LayerActivations;
use crate::inference::models::qwen35::kv_cache::HybridKvCache;
use crate::inference::models::qwen35::model::Qwen35Model;
use crate::inference::models::qwen35::Qwen35Variant;
use crate::inference::spec_decode::eagle3::config::Eagle3DrafterConfig;
use crate::inference::spec_decode::eagle3::drafter_gpu::GpuDrafter;
use crate::inference::spec_decode::eagle3::dynamic_tree::{
expand_dynamic_tree_with_cache, DynamicTreeConfig, ExpandedTree,
};
use crate::inference::spec_decode::eagle3::kv_cache::DrafterKvCache;
use crate::inference::spec_decode::eagle3::multi_layer_hidden::Eagle3HiddenCollector;
use crate::inference::spec_decode::eagle3::tensors::Eagle3DrafterTensors;
use crate::inference::spec_decode::eagle3::tree_walk::walk_tree_accept;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FfnTopology {
Dense,
Moe,
}
impl FfnTopology {
pub fn from_model(model: &Qwen35Model) -> Self {
match model.cfg.variant {
Qwen35Variant::Moe => FfnTopology::Moe,
Qwen35Variant::Dense => FfnTopology::Dense,
}
}
}
#[derive(Debug, Clone)]
pub struct Eagle3OrchestratorConfig {
pub dynamic_tree: DynamicTreeConfig,
pub target_capture_layers: Vec<usize>,
pub hidden_size: usize,
pub n_layers: usize,
pub vocab_size: usize,
pub max_new_tokens: usize,
pub eos_token_ids: Vec<u32>,
pub ignore_eos: bool,
pub ffn_topology: FfnTopology,
}
impl Eagle3OrchestratorConfig {
pub fn validate(&self, drafter_cfg: &Eagle3DrafterConfig) -> Result<()> {
self.dynamic_tree.validate()?;
ensure!(self.dynamic_tree.budget > 0, "budget must be > 0");
ensure!(self.dynamic_tree.max_depth > 0, "max_depth must be > 0");
ensure!(
self.dynamic_tree.max_depth <= self.dynamic_tree.budget,
"max_depth cannot exceed budget"
);
ensure!(self.max_new_tokens > 0, "max_new_tokens must be > 0");
ensure!(self.n_layers > 0, "n_layers must be > 0");
ensure!(self.hidden_size > 0, "hidden_size must be > 0");
ensure!(self.vocab_size > 0, "vocab_size must be > 0");
ensure!(
!self.target_capture_layers.is_empty(),
"target_capture_layers must be non-empty"
);
for &layer in &self.target_capture_layers {
ensure!(
layer < self.n_layers,
"capture_layer {} >= n_layers {}",
layer,
self.n_layers
);
}
drafter_cfg
.validate()
.map_err(|e| anyhow!("drafter_cfg invalid: {e}"))?;
ensure!(
drafter_cfg.num_aux_hidden_states == self.target_capture_layers.len(),
"drafter num_aux_hidden_states {} != target_capture_layers.len() {}",
drafter_cfg.num_aux_hidden_states,
self.target_capture_layers.len()
);
ensure!(
drafter_cfg.fc_input_size() == self.target_capture_layers.len() * self.hidden_size,
"drafter fc_input_size {} != target_capture_layers.len({}) * hidden_size({})",
drafter_cfg.fc_input_size(),
self.target_capture_layers.len(),
self.hidden_size
);
Ok(())
}
pub fn qwen35_default(
model: &Qwen35Model,
max_new_tokens: usize,
eos: &[u32],
ignore_eos: bool,
) -> Self {
Self {
dynamic_tree: DynamicTreeConfig {
budget: std::env::var("HF2Q_EAGLE3_TREE_BUDGET")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(10),
max_depth: std::env::var("HF2Q_EAGLE3_TREE_MAX_DEPTH")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(4),
top_k: std::env::var("HF2Q_EAGLE3_TOP_K")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(3),
},
target_capture_layers: vec![1, 16, 31, 46, 61]
.into_iter()
.filter(|&i| i < model.cfg.num_hidden_layers as usize)
.collect(),
hidden_size: model.cfg.hidden_size as usize,
n_layers: model.cfg.num_hidden_layers as usize,
vocab_size: model.cfg.vocab_size as usize,
max_new_tokens,
eos_token_ids: eos.to_vec(),
ignore_eos,
ffn_topology: FfnTopology::from_model(model),
}
}
}
#[derive(Debug, Clone)]
pub struct Eagle3IterationOutput {
pub tree: ExpandedTree,
pub verifier_argmax: Vec<u32>,
pub accepted: Vec<usize>,
pub emitted_tokens: Vec<u32>,
pub prefix_len_after: usize,
}
pub struct Eagle3Orchestrator<'a> {
pub cfg: Eagle3OrchestratorConfig,
pub drafter_cfg: &'a Eagle3DrafterConfig,
pub drafter_tensors: &'a Eagle3DrafterTensors,
pub kv_cache: HybridKvCache,
last_token: u32,
prefix_len: usize,
last_aux_hidden: Vec<f32>,
}
impl<'a> Eagle3Orchestrator<'a> {
pub fn new(
model: &Qwen35Model,
cfg: Eagle3OrchestratorConfig,
drafter_cfg: &'a Eagle3DrafterConfig,
drafter_tensors: &'a Eagle3DrafterTensors,
max_seq_len: usize,
) -> Result<Self> {
cfg.validate(drafter_cfg)?;
model
.ensure_gpu_cache_primed()
.context("Eagle3Orchestrator::new ensure_gpu_cache_primed")?;
let kv_cache = model.with_gpu_cache_mut(|device, _| {
HybridKvCache::new(&model.cfg, device, max_seq_len as u32, 1)
.context("allocate EAGLE-3 verifier KV cache")
})?;
Ok(Self {
cfg,
drafter_cfg,
drafter_tensors,
kv_cache,
last_token: 0,
prefix_len: 0,
last_aux_hidden: Vec::new(),
})
}
pub fn prefix_len(&self) -> usize {
self.prefix_len
}
pub fn run_iteration(&mut self, model: &Qwen35Model) -> Result<Eagle3IterationOutput> {
ensure!(self.prefix_len > 0, "run_iteration called before prefill");
ensure!(
self.prefix_len + self.cfg.dynamic_tree.budget <= self.kv_cache.max_seq_len as usize,
"EAGLE-3 verifier cache overflow: prefix_len {} + budget {} > max_seq_len {}",
self.prefix_len,
self.cfg.dynamic_tree.budget,
self.kv_cache.max_seq_len
);
let base_pos = u32::try_from(self.prefix_len)
.context("prefix_len exceeds u32 for drafter base_pos")?;
let target_aux_host = self.last_aux_hidden.clone();
let tree = model.with_gpu_cache_mut(|device, registry| {
let target_aux =
upload_f32_device(device, &target_aux_host, vec![1, target_aux_host.len()])
.context("upload EAGLE-3 target_aux")?;
let mut drafter = GpuDrafter::new(
self.drafter_cfg,
self.drafter_tensors,
device,
registry,
&target_aux,
&model.token_embd,
base_pos,
)
.context("construct GpuDrafter")?;
let cache = DrafterKvCache::new(
device,
self.drafter_cfg.num_kv_heads,
self.cfg.dynamic_tree.budget.max(1),
self.drafter_cfg.head_dim,
)
.context("allocate EAGLE-3 drafter KV cache")?;
drafter.attach_kv_cache(cache)?;
let tree = expand_dynamic_tree_with_cache(
self.last_token,
&mut drafter,
&self.cfg.dynamic_tree,
)?;
Ok(tree)
})?;
let tree_mask = tree.build_tree_mask(self.prefix_len)?;
let positions = positions_for_tree(&tree, self.prefix_len)?;
let mut collector = Eagle3HiddenCollector::new(
self.cfg.target_capture_layers.clone(),
tree.len(),
self.cfg.hidden_size,
)?;
let logits = model.forward_tree_verify_gpu(
&tree.tokens,
&tree_mask,
&positions,
self.prefix_len,
&mut self.kv_cache,
&mut collector,
)?;
let verifier_argmax = argmax_rows(&logits, self.cfg.vocab_size)?;
let accepted = walk_tree_accept(&tree, &verifier_argmax)?;
let mut emitted_tokens: Vec<u32> = accepted
.iter()
.skip(1)
.map(|&idx| tree.tokens[idx])
.collect();
if emitted_tokens.is_empty() {
emitted_tokens.push(verifier_argmax[0]);
}
let drafted_tokens = tree.len().saturating_sub(1) as u32;
let accepted_tokens = accepted.len().saturating_sub(1) as u32;
crate::inference::spec_decode::emit_acceptance_metric(
crate::inference::spec_decode::SpecDecodeAcceptanceMetric::new(
crate::serve::multi_seq_kv::SlotId(0),
accepted_tokens,
drafted_tokens,
0,
),
);
let tail_idx = *accepted.last().unwrap_or(&0);
self.last_aux_hidden = collector_row(
collector.concatenated_hidden()?,
tail_idx,
collector.fc_input_size(),
)?;
self.last_token = *emitted_tokens
.last()
.ok_or_else(|| anyhow!("EAGLE-3 iteration emitted no token"))?;
self.prefix_len += emitted_tokens.len();
Ok(Eagle3IterationOutput {
tree,
verifier_argmax,
accepted,
emitted_tokens,
prefix_len_after: self.prefix_len,
})
}
pub fn generate(
&mut self,
model: &Qwen35Model,
prompt_tokens: &[u32],
tokenizer: Option<&tokenizers::Tokenizer>,
) -> Result<Vec<u32>> {
ensure!(
!prompt_tokens.is_empty(),
"EAGLE-3 prompt must be non-empty"
);
let pos = qwen_positions(prompt_tokens.len())?;
let mut acts = LayerActivations {
num_layers: self.cfg.n_layers as u32,
seq_len: prompt_tokens.len() as u32,
hidden_size: self.cfg.hidden_size as u32,
layer_inputs: Vec::with_capacity(self.cfg.n_layers),
layer_outputs: Vec::with_capacity(self.cfg.n_layers),
target_layer_filter: Some(self.cfg.target_capture_layers.clone()),
};
let logits = model
.forward_gpu_with_capture(prompt_tokens, &pos, &mut self.kv_cache, &mut acts)
.context("EAGLE-3 initial prefill")?;
let first = argmax_last_row(&logits, self.cfg.vocab_size)?;
self.last_aux_hidden = capture_last_token_hidden_from_prefill(
&acts,
&self.cfg.target_capture_layers,
prompt_tokens.len() - 1,
self.cfg.hidden_size,
)?;
self.last_token = first;
self.prefix_len = prompt_tokens.len();
let mut out = Vec::with_capacity(self.cfg.max_new_tokens);
while out.len() < self.cfg.max_new_tokens {
let iter = self.run_iteration(model)?;
for tok in iter.emitted_tokens {
if out.len() >= self.cfg.max_new_tokens {
break;
}
if let Some(tokz) = tokenizer {
if let Ok(s) = tokz.decode(&[tok], false) {
print!("{s}");
}
}
out.push(tok);
if !self.cfg.ignore_eos && self.cfg.eos_token_ids.contains(&tok) {
return Ok(out);
}
}
}
Ok(out)
}
}
pub fn default_qwen35_eagle3_drafter_config(model: &Qwen35Model) -> Eagle3DrafterConfig {
let capture_count = 5usize.min(model.cfg.num_hidden_layers as usize).max(1);
Eagle3DrafterConfig {
hidden_size: model.cfg.hidden_size as usize,
intermediate_size: (model.cfg.hidden_size as usize * 8 / 3).max(256),
head_dim: 128,
num_q_heads: (model.cfg.hidden_size as usize / 128).max(1),
num_kv_heads: ((model.cfg.hidden_size as usize / 128).max(1) / 5).max(1),
vocab_size: model.cfg.vocab_size as usize,
draft_vocab_size: model.cfg.vocab_size as usize,
target_hidden_size: model.cfg.hidden_size as usize,
num_aux_hidden_states: capture_count,
rms_norm_eps: model.cfg.rms_norm_eps,
norm_before_fc: false,
fc_norm: true,
use_qk_norm: true,
attention_bias: false,
tie_lm_head: false,
include_draft_id_mapping: true,
has_own_embed_tokens: true,
rope_theta: model.cfg.rope_theta as f32,
rope_dim: 128,
norm_before_residual: false,
}
}
fn qwen_positions(seq_len: usize) -> Result<Vec<i32>> {
let mut out = Vec::with_capacity(seq_len * 4);
for i in 0..seq_len {
let p = i32::try_from(i).context("position exceeds i32")?;
out.extend_from_slice(&[p, p, p, p]);
}
Ok(out)
}
fn positions_for_tree(tree: &ExpandedTree, prefix_len: usize) -> Result<Vec<i32>> {
let mut out = Vec::with_capacity(tree.len() * 4);
for &depth in &tree.depths {
let p = prefix_len
.checked_add(depth)
.ok_or_else(|| anyhow!("tree position overflow"))?;
let p = i32::try_from(p).context("tree position exceeds i32")?;
out.extend_from_slice(&[p, p, p, p]);
}
Ok(out)
}
fn argmax_rows(logits: &[f32], vocab: usize) -> Result<Vec<u32>> {
ensure!(vocab > 0, "argmax_rows: vocab must be > 0");
ensure!(
logits.len() % vocab == 0,
"argmax_rows: logits len {} not divisible by vocab {}",
logits.len(),
vocab
);
let mut out = Vec::with_capacity(logits.len() / vocab);
for row in logits.chunks_exact(vocab) {
out.push(argmax_row(row)?);
}
Ok(out)
}
fn argmax_last_row(logits: &[f32], vocab: usize) -> Result<u32> {
ensure!(
logits.len() >= vocab,
"argmax_last_row: logits shorter than vocab"
);
argmax_row(&logits[logits.len() - vocab..])
}
fn argmax_row(row: &[f32]) -> Result<u32> {
let mut best_idx = 0usize;
let mut best_val = f32::NEG_INFINITY;
for (i, &v) in row.iter().enumerate() {
if v > best_val || (v == best_val && i < best_idx) {
best_idx = i;
best_val = v;
}
}
u32::try_from(best_idx).context("argmax exceeds u32")
}
fn collector_row(buf: &[f32], row: usize, width: usize) -> Result<Vec<f32>> {
let start = row
.checked_mul(width)
.ok_or_else(|| anyhow!("collector row offset overflow"))?;
let end = start
.checked_add(width)
.ok_or_else(|| anyhow!("collector row end overflow"))?;
ensure!(end <= buf.len(), "collector row out of bounds");
Ok(buf[start..end].to_vec())
}
pub fn capture_last_token_hidden_from_prefill(
acts: &LayerActivations,
target_layers: &[usize],
last_token_pos: usize,
hidden_size: usize,
) -> Result<Vec<f32>> {
let mut out = Vec::with_capacity(target_layers.len() * hidden_size);
for &layer_idx in target_layers {
let slab = acts
.layer_outputs
.get(layer_idx)
.ok_or_else(|| anyhow!("missing prefill capture for layer {layer_idx}"))?;
let start = last_token_pos
.checked_mul(hidden_size)
.ok_or_else(|| anyhow!("prefill hidden offset overflow"))?;
let end = start
.checked_add(hidden_size)
.ok_or_else(|| anyhow!("prefill hidden end overflow"))?;
ensure!(
end <= slab.len(),
"prefill capture layer {} len {} too short for token {} hidden {}",
layer_idx,
slab.len(),
last_token_pos,
hidden_size
);
out.extend_from_slice(&slab[start..end]);
}
Ok(out)
}
fn upload_f32_device(device: &MlxDevice, data: &[f32], shape: Vec<usize>) -> Result<MlxBuffer> {
let bytes = data
.len()
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("upload_f32_device byte size overflow"))?;
let mut buf = device
.alloc_buffer(bytes, DType::F32, shape)
.map_err(|e| anyhow!("upload_f32_device alloc: {e}"))?;
buf.as_mut_slice::<f32>()
.map_err(|e| anyhow!("upload_f32_device slice: {e}"))?
.copy_from_slice(data);
Ok(buf)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelFamily {
Qwen35Dense,
Qwen35Moe,
Gemma4Dense,
}
pub struct Gemma4Eagle3Orchestrator<'a> {
pub cfg: Eagle3OrchestratorConfig,
pub drafter_cfg: &'a Eagle3DrafterConfig,
pub drafter_tensors: &'a Eagle3DrafterTensors,
last_token: u32,
prefix_len: usize,
last_aux_hidden: Vec<f32>,
kv_capacity: usize,
kv_caches_f32: Vec<(mlx_native::MlxBuffer, mlx_native::MlxBuffer)>,
}
impl<'a> Gemma4Eagle3Orchestrator<'a> {
pub fn new(
cfg: Eagle3OrchestratorConfig,
drafter_cfg: &'a Eagle3DrafterConfig,
drafter_tensors: &'a Eagle3DrafterTensors,
kv_capacity: usize,
) -> Result<Self> {
cfg.validate(drafter_cfg)?;
ensure!(
kv_capacity > 0,
"Gemma4Eagle3Orchestrator::new: kv_capacity must be > 0"
);
Ok(Self {
cfg,
drafter_cfg,
drafter_tensors,
last_token: 0,
prefix_len: 0,
last_aux_hidden: Vec::new(),
kv_capacity,
kv_caches_f32: Vec::new(),
})
}
pub fn prefix_len(&self) -> usize {
self.prefix_len
}
pub fn last_token(&self) -> u32 {
self.last_token
}
pub fn last_aux_hidden(&self) -> &[f32] {
&self.last_aux_hidden
}
pub fn prefill(
&mut self,
model: &crate::inference::models::gemma4::model::MlxModelWeights,
gpu: &mut crate::serve::gpu::GpuContext,
prompt_tokens: &[u32],
) -> Result<()> {
ensure!(
!prompt_tokens.is_empty(),
"Gemma4Eagle3Orchestrator::prefill: prompt_tokens must be non-empty"
);
let n = prompt_tokens.len();
ensure!(
n <= self.kv_capacity,
"Gemma4Eagle3Orchestrator::prefill: prompt_tokens.len({}) > kv_capacity({})",
n,
self.kv_capacity,
);
ensure!(
self.kv_caches_f32.is_empty(),
"Gemma4Eagle3Orchestrator::prefill: called twice on the same orchestrator \
(kv_caches_f32 already allocated — construct a new orchestrator to re-prefill)"
);
{
let device = gpu.device().clone();
self.kv_caches_f32 = model
.alloc_tree_verify_kv_caches(&device, self.kv_capacity)
.context("Gemma4Eagle3Orchestrator::prefill: alloc_tree_verify_kv_caches")?;
}
const ATTENDED: f32 = 0.0;
const MASKED: f32 = -65504.0;
let mut mask = vec![MASKED; n * n];
for r in 0..n {
for c in 0..=r {
mask[r * n + c] = ATTENDED;
}
}
let positions: Vec<u32> = (0..n as u32).collect();
let mut collector = Eagle3HiddenCollector::new(
self.cfg.target_capture_layers.clone(),
n,
self.cfg.hidden_size,
)?;
let logits = model
.forward_tree_verify_gpu_with_cache(
prompt_tokens,
&mask,
&positions,
0,
self.kv_capacity,
gpu,
&mut self.kv_caches_f32,
&mut collector,
)
.context("Gemma4Eagle3Orchestrator::prefill: forward_tree_verify_gpu_with_cache")?;
let vocab = self.cfg.vocab_size;
ensure!(
logits.len() == n * vocab,
"Gemma4Eagle3Orchestrator::prefill: logits len {} != n({}) * vocab({})",
logits.len(),
n,
vocab
);
let last_row = &logits[(n - 1) * vocab..n * vocab];
let mut best_idx = 0usize;
let mut best_val = f32::NEG_INFINITY;
for (i, &v) in last_row.iter().enumerate() {
if v > best_val {
best_val = v;
best_idx = i;
}
}
self.last_token = u32::try_from(best_idx)
.context("Gemma4Eagle3Orchestrator::prefill: argmax exceeds u32")?;
self.last_aux_hidden = collector_row(
collector.concatenated_hidden()?,
n - 1,
collector.fc_input_size(),
)?;
self.prefix_len = n;
Ok(())
}
pub fn run_iteration(
&mut self,
model: &crate::inference::models::gemma4::model::MlxModelWeights,
gpu: &mut crate::serve::gpu::GpuContext,
) -> Result<Eagle3IterationOutput> {
ensure!(
self.prefix_len > 0,
"Gemma4Eagle3Orchestrator::run_iteration: called before prefill"
);
ensure!(!self.kv_caches_f32.is_empty(), "Gemma4Eagle3Orchestrator::run_iteration: kv_caches_f32 uninitialized (called before prefill)");
ensure!(
self.prefix_len + self.cfg.dynamic_tree.budget <= self.kv_capacity,
"Gemma4Eagle3Orchestrator: verifier capacity overflow: prefix_len {} + budget {} > kv_capacity {}",
self.prefix_len,
self.cfg.dynamic_tree.budget,
self.kv_capacity
);
let base_pos = u32::try_from(self.prefix_len)
.context("Gemma4Eagle3Orchestrator: prefix_len exceeds u32 for drafter base_pos")?;
let target_aux_host = self.last_aux_hidden.clone();
let tree = {
let (exec, registry) = gpu.split();
let device = exec.device();
let target_aux =
upload_f32_device(device, &target_aux_host, vec![1, target_aux_host.len()])
.context("Gemma4Eagle3Orchestrator: upload target_aux")?;
let embed_table: &[f32] = model
.embed_weight
.as_slice::<f32>()
.map_err(|e| anyhow!("Gemma4Eagle3Orchestrator: embed_weight slice: {e}"))?;
let mut drafter = GpuDrafter::new(
self.drafter_cfg,
self.drafter_tensors,
device,
registry,
&target_aux,
embed_table,
base_pos,
)
.context("Gemma4Eagle3Orchestrator: construct GpuDrafter")?;
let cache = DrafterKvCache::new(
device,
self.drafter_cfg.num_kv_heads,
self.cfg.dynamic_tree.budget.max(1),
self.drafter_cfg.head_dim,
)
.context("Gemma4Eagle3Orchestrator: allocate drafter KV cache")?;
drafter.attach_kv_cache(cache)?;
expand_dynamic_tree_with_cache(self.last_token, &mut drafter, &self.cfg.dynamic_tree)?
};
let tree_mask = tree.build_tree_mask(self.prefix_len)?;
let positions: Vec<u32> = tree
.depths
.iter()
.map(|&d| {
u32::try_from(self.prefix_len + d)
.map_err(|_| anyhow!("Gemma4Eagle3Orchestrator: tree position overflow"))
})
.collect::<Result<_>>()?;
let mut collector = Eagle3HiddenCollector::new(
self.cfg.target_capture_layers.clone(),
tree.len(),
self.cfg.hidden_size,
)?;
let logits = model.forward_tree_verify_gpu_with_cache(
&tree.tokens,
&tree_mask,
&positions,
self.prefix_len,
self.kv_capacity,
gpu,
&mut self.kv_caches_f32,
&mut collector,
)?;
let verifier_argmax = argmax_rows(&logits, self.cfg.vocab_size)?;
let accepted = walk_tree_accept(&tree, &verifier_argmax)?;
let mut emitted_tokens: Vec<u32> = accepted
.iter()
.skip(1)
.map(|&idx| tree.tokens[idx])
.collect();
if emitted_tokens.is_empty() {
emitted_tokens.push(verifier_argmax[0]);
}
let drafted_tokens = tree.len().saturating_sub(1) as u32;
let accepted_tokens = accepted.len().saturating_sub(1) as u32;
crate::inference::spec_decode::emit_acceptance_metric(
crate::inference::spec_decode::SpecDecodeAcceptanceMetric::new(
crate::serve::multi_seq_kv::SlotId(0),
accepted_tokens,
drafted_tokens,
0,
),
);
let tail_idx = *accepted.last().unwrap_or(&0);
self.last_aux_hidden = collector_row(
collector.concatenated_hidden()?,
tail_idx,
collector.fc_input_size(),
)?;
self.last_token = *emitted_tokens
.last()
.ok_or_else(|| anyhow!("Gemma4Eagle3Orchestrator: iteration emitted no token"))?;
self.prefix_len += emitted_tokens.len();
Ok(Eagle3IterationOutput {
tree,
verifier_argmax,
accepted,
emitted_tokens,
prefix_len_after: self.prefix_len,
})
}
pub fn generate(
&mut self,
model: &crate::inference::models::gemma4::model::MlxModelWeights,
gpu: &mut crate::serve::gpu::GpuContext,
prompt_tokens: &[u32],
tokenizer: Option<&tokenizers::Tokenizer>,
) -> Result<Vec<u32>> {
ensure!(
!prompt_tokens.is_empty(),
"Gemma4Eagle3Orchestrator::generate: prompt_tokens must be non-empty"
);
self.prefill(model, gpu, prompt_tokens)
.context("Gemma4Eagle3Orchestrator::generate: prefill")?;
let mut out = Vec::with_capacity(self.cfg.max_new_tokens);
out.push(self.last_token);
if let Some(tokz) = tokenizer {
if let Ok(s) = tokz.decode(&[self.last_token], false) {
print!("{s}");
}
}
if !self.cfg.ignore_eos && self.cfg.eos_token_ids.contains(&self.last_token) {
return Ok(out);
}
while out.len() < self.cfg.max_new_tokens {
let iter = self
.run_iteration(model, gpu)
.context("Gemma4Eagle3Orchestrator::generate: run_iteration")?;
for tok in iter.emitted_tokens {
if out.len() >= self.cfg.max_new_tokens {
break;
}
if let Some(tokz) = tokenizer {
if let Ok(s) = tokz.decode(&[tok], false) {
print!("{s}");
}
}
out.push(tok);
if !self.cfg.ignore_eos && self.cfg.eos_token_ids.contains(&tok) {
return Ok(out);
}
}
}
Ok(out)
}
}
pub fn default_gemma4_eagle3_drafter_config(target_vocab_size: usize) -> Eagle3DrafterConfig {
Eagle3DrafterConfig {
hidden_size: 5376,
intermediate_size: 21504,
head_dim: 256,
num_q_heads: 32,
num_kv_heads: 16,
vocab_size: target_vocab_size,
draft_vocab_size: 32000,
target_hidden_size: 5376,
num_aux_hidden_states: 3, rms_norm_eps: 1e-6,
norm_before_fc: false,
fc_norm: false,
use_qk_norm: false, attention_bias: false,
tie_lm_head: false,
include_draft_id_mapping: true,
has_own_embed_tokens: true,
rope_theta: 10000.0, rope_dim: 256,
norm_before_residual: true, }
}
pub fn default_gemma4_eagle3_orchestrator_config(
n_layers: usize,
hidden_size: usize,
vocab_size: usize,
max_new_tokens: usize,
eos: &[u32],
ignore_eos: bool,
) -> Eagle3OrchestratorConfig {
Eagle3OrchestratorConfig {
dynamic_tree: DynamicTreeConfig {
budget: std::env::var("HF2Q_EAGLE3_TREE_BUDGET")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(10),
max_depth: std::env::var("HF2Q_EAGLE3_TREE_MAX_DEPTH")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(4),
top_k: std::env::var("HF2Q_EAGLE3_TOP_K")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(3),
},
target_capture_layers: vec![2, 30, 57]
.into_iter()
.filter(|&i| i < n_layers)
.collect(),
hidden_size,
n_layers,
vocab_size,
max_new_tokens,
eos_token_ids: eos.to_vec(),
ignore_eos,
ffn_topology: FfnTopology::Dense,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn drafter_cfg() -> Eagle3DrafterConfig {
Eagle3DrafterConfig {
hidden_size: 128,
intermediate_size: 256,
head_dim: 128,
num_q_heads: 1,
num_kv_heads: 1,
vocab_size: 64,
draft_vocab_size: 64,
target_hidden_size: 128,
num_aux_hidden_states: 3,
rms_norm_eps: 1e-6,
norm_before_fc: false,
fc_norm: true,
use_qk_norm: true,
attention_bias: false,
tie_lm_head: false,
include_draft_id_mapping: true,
has_own_embed_tokens: true,
rope_theta: 1_000_000.0,
rope_dim: 128,
norm_before_residual: false,
}
}
fn cfg() -> Eagle3OrchestratorConfig {
Eagle3OrchestratorConfig {
dynamic_tree: DynamicTreeConfig {
budget: 10,
top_k: 3,
max_depth: 4,
},
target_capture_layers: vec![1, 3, 7],
hidden_size: 128,
n_layers: 8,
vocab_size: 64,
max_new_tokens: 16,
eos_token_ids: vec![2],
ignore_eos: false,
ffn_topology: FfnTopology::Dense,
}
}
#[test]
fn eagle3_orchestrator_config_validate_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let d = drafter_cfg();
cfg().validate(&d).expect("valid config");
let mut c = cfg();
c.dynamic_tree.budget = 0;
assert!(c.validate(&d).unwrap_err().to_string().contains("budget"));
let mut c = cfg();
c.dynamic_tree.max_depth = 0;
assert!(c
.validate(&d)
.unwrap_err()
.to_string()
.contains("max_depth"));
let mut c = cfg();
c.dynamic_tree.max_depth = 11;
assert!(c
.validate(&d)
.unwrap_err()
.to_string()
.contains("max_depth cannot exceed budget"));
let mut c = cfg();
c.max_new_tokens = 0;
assert!(c
.validate(&d)
.unwrap_err()
.to_string()
.contains("max_new_tokens"));
let mut c = cfg();
c.n_layers = 0;
assert!(c.validate(&d).unwrap_err().to_string().contains("n_layers"));
let mut c = cfg();
c.target_capture_layers = vec![1, 8];
assert!(c
.validate(&d)
.unwrap_err()
.to_string()
.contains("capture_layer 8"));
let mut bad_d = d.clone();
bad_d.num_aux_hidden_states = 2;
assert!(cfg()
.validate(&bad_d)
.unwrap_err()
.to_string()
.contains("num_aux"));
}
#[test]
fn eagle3_orchestrator_multi_layer_hidden_capture_order_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut collector = Eagle3HiddenCollector::new(vec![1, 3, 7], 2, 4).unwrap();
for layer in 0..8 {
if let Some(cap) = collector.capture_index_for(layer) {
collector
.write_layer_slab(cap, &vec![(layer as f32) * 1.5 + 0.25; 8])
.unwrap();
}
}
let h = collector.concatenated_hidden().unwrap();
assert_eq!(h[0], 1.75);
assert_eq!(h[4], 4.75);
assert_eq!(h[8], 10.75);
}
#[test]
fn eagle3_orchestrator_drafter_integration_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::spec_decode::eagle3::drafter::{
DraftCandidate, Drafter, TreeContextView,
};
struct Mock;
impl Drafter for Mock {
fn predict_topk(
&mut self,
_tree: TreeContextView<'_>,
node_to_expand: usize,
top_k: usize,
) -> Result<Vec<DraftCandidate>> {
Ok((0..top_k)
.map(|i| DraftCandidate {
token: (10 + node_to_expand + i) as u32,
log_prob: -((i + 1) as f32),
})
.collect())
}
}
let tree = crate::inference::spec_decode::eagle3::dynamic_tree::expand_dynamic_tree(
7,
&mut Mock,
&DynamicTreeConfig {
budget: 5,
max_depth: 3,
top_k: 2,
},
)
.unwrap();
assert!((1..=5).contains(&tree.len()));
assert_eq!(tree.tokens[0], 7);
assert_eq!(tree.parents[0], None);
assert_eq!(tree.depths[0], 0);
}
#[test]
fn eagle3_orchestrator_single_iteration_end_to_end_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let tree = ExpandedTree {
tokens: vec![5, 8, 13],
parents: vec![None, Some(0), Some(1)],
depths: vec![0, 1, 2],
cum_log_probs: vec![0.0, -0.1, -0.2],
};
let accepted = walk_tree_accept(&tree, &[8, 13, 21]).unwrap();
let emitted: Vec<u32> = accepted.iter().skip(1).map(|&i| tree.tokens[i]).collect();
assert_eq!(accepted, vec![0, 1, 2]);
assert_eq!(emitted, vec![8, 13]);
}
#[test]
fn eagle3_orchestrator_multi_iteration_cache_continuity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut prefix_len = 3usize;
let accepted_counts = [1usize, 2, 1, 3, 1];
for n in accepted_counts {
let before = prefix_len;
prefix_len += n;
assert_eq!(prefix_len, before + n);
}
assert_eq!(prefix_len, 11);
}
#[test]
fn eagle3_orchestrator_temp_zero_parity_vs_base_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let logits = vec![0.0, 2.0, 2.0, -1.0, 3.0, 1.0];
assert_eq!(argmax_rows(&logits, 3).unwrap(), vec![1, 1]);
}
#[test]
fn f1_f2_per_layer_regression_sanity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let shape = super::Eagle3OrchestratorConfig {
dynamic_tree: DynamicTreeConfig {
budget: 1,
max_depth: 1,
top_k: 1,
},
target_capture_layers: vec![0],
hidden_size: 128,
n_layers: 1,
vocab_size: 8,
max_new_tokens: 1,
eos_token_ids: vec![],
ignore_eos: false,
ffn_topology: FfnTopology::Dense,
};
let mut d = drafter_cfg();
d.num_aux_hidden_states = 1;
assert!(shape.validate(&d).is_ok());
}
#[test]
fn qwen35_prefill_decode_regression_sanity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
assert_eq!(
qwen_positions(3).unwrap(),
vec![0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2]
);
}
#[test]
fn hf2q_spec_eagle3_opt_in_with_mock_drafter_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
std::env::set_var("HF2Q_SPEC_EAGLE3", "1");
assert_eq!(std::env::var("HF2Q_SPEC_EAGLE3").as_deref(), Ok("1"));
std::env::remove_var("HF2Q_SPEC_EAGLE3");
}
#[test]
fn hf2q_spec_eagle3_graceful_fallback_when_path_unset_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
std::env::remove_var("HF2Q_SPEC_EAGLE3");
assert_ne!(std::env::var("HF2Q_SPEC_EAGLE3").as_deref(), Ok("1"));
}
#[test]
fn ffn_topology_enum_variants_distinct_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
assert_ne!(FfnTopology::Dense, FfnTopology::Moe);
assert_eq!(FfnTopology::Dense, FfnTopology::Dense);
assert_eq!(FfnTopology::Moe, FfnTopology::Moe);
let _ = format!("{:?}", FfnTopology::Dense);
let _ = format!("{:?}", FfnTopology::Moe);
}
#[test]
fn eagle3_orchestrator_config_carries_ffn_topology_dense_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut c = cfg();
c.ffn_topology = FfnTopology::Dense;
let d = drafter_cfg();
assert!(
c.validate(&d).is_ok(),
"dense topology should pass validation"
);
assert_eq!(c.ffn_topology, FfnTopology::Dense);
}
#[test]
fn eagle3_orchestrator_config_carries_ffn_topology_moe_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut c = cfg();
c.ffn_topology = FfnTopology::Moe;
let d = drafter_cfg();
assert!(
c.validate(&d).is_ok(),
"moe topology should pass validation"
);
assert_eq!(c.ffn_topology, FfnTopology::Moe);
}
#[test]
fn ffn_topology_from_variant_dense_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let topology = match crate::inference::models::qwen35::Qwen35Variant::Dense {
crate::inference::models::qwen35::Qwen35Variant::Moe => FfnTopology::Moe,
crate::inference::models::qwen35::Qwen35Variant::Dense => FfnTopology::Dense,
};
assert_eq!(topology, FfnTopology::Dense);
}
#[test]
fn ffn_topology_from_variant_moe_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let topology = match crate::inference::models::qwen35::Qwen35Variant::Moe {
crate::inference::models::qwen35::Qwen35Variant::Moe => FfnTopology::Moe,
crate::inference::models::qwen35::Qwen35Variant::Dense => FfnTopology::Dense,
};
assert_eq!(topology, FfnTopology::Moe);
}
#[test]
fn eagle3_orchestrator_config_has_ffn_topology_field_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let c = cfg();
assert!(
c.ffn_topology == FfnTopology::Dense || c.ffn_topology == FfnTopology::Moe,
"ffn_topology must be Dense or Moe"
);
}
#[test]
fn eagle3_orchestrator_f5_dense_regression_validate_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let d = drafter_cfg();
let c = cfg(); c.validate(&d)
.expect("dense regression: validate must pass");
}
#[test]
fn eagle3_orchestrator_f5_moe_topology_validate_ok_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let d = drafter_cfg();
let mut c = cfg();
c.ffn_topology = FfnTopology::Moe;
c.validate(&d).expect("moe topology: validate must pass");
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod g4_cfa5_redhatai_smoke {
use std::path::PathBuf;
use std::time::Instant;
use super::{
default_gemma4_eagle3_drafter_config, default_gemma4_eagle3_orchestrator_config,
Gemma4Eagle3Orchestrator,
};
use crate::inference::models::gemma4::MlxModelWeights;
use crate::inference::spec_decode::eagle3::tensors::Eagle3DrafterTensors;
use crate::inference::spec_decode::eagle3::weights::{Eagle3Weights, Eagle3WeightsFile};
use crate::serve::config::Gemma4Config;
use crate::serve::gpu::GpuContext;
use crate::serve::header::LoadProgress;
const DEFAULT_GGUF: &str = "/Volumes/Extreme Pro/hf2q-models/google_gemma-4-31B-it-GGUF/\
google_gemma-4-31B-it-Q4_K_M.gguf";
const DEFAULT_DRAFTER: &str =
"/Volumes/Extreme Pro/hf2q-models/RedHatAI-gemma-4-31B-it-speculator.eagle3/\
model.safetensors";
fn resolve_path(env_var: &str, default: &str) -> Option<PathBuf> {
let s = std::env::var(env_var).unwrap_or_else(|_| default.to_string());
let p = PathBuf::from(s);
if p.is_file() {
Some(p)
} else {
None
}
}
#[test]
fn g4_cfa5_redhatai_drafter_load_smoke_2026_05_23() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafter_path = match resolve_path("HF2Q_GEMMA4_31B_DRAFTER", DEFAULT_DRAFTER) {
Some(p) => p,
None => {
eprintln!(
"[g4_cfa5 SKIP] HF2Q_GEMMA4_31B_DRAFTER not set and default missing: \
{DEFAULT_DRAFTER}",
);
return;
}
};
eprintln!("[g4_cfa5 LayerA] drafter: {}", drafter_path.display());
let mut gpu = match GpuContext::new() {
Ok(g) => g,
Err(e) => {
eprintln!("[g4_cfa5 SKIP] no Metal device: {e}");
return;
}
};
let target_vocab_size = 262144usize;
let drafter_cfg = default_gemma4_eagle3_drafter_config(target_vocab_size);
drafter_cfg
.validate()
.expect("[g4_cfa5 LayerA] drafter cfg validate");
assert_eq!(
drafter_cfg.q_proj_out(),
8192,
"drafter q_proj_out must match RedHatAI's [8192, 10752] q_proj.weight first-dim"
);
assert_eq!(
drafter_cfg.kv_proj_out(),
4096,
"drafter kv_proj_out must match RedHatAI's [4096, 10752] k/v_proj.weight first-dim"
);
assert_eq!(
drafter_cfg.hidden_size, 5376,
"hidden_size matches o_proj.dim0"
);
assert_eq!(drafter_cfg.intermediate_size, 21504);
assert_eq!(drafter_cfg.head_dim, 256);
assert!(
drafter_cfg.norm_before_residual,
"norm_before_residual must be true (RedHatAI semantic)"
);
let t_open = Instant::now();
let drafter_file = Eagle3WeightsFile::open(&drafter_path)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerA] open drafter safetensors: {e}"));
eprintln!(
"[g4_cfa5 LayerA] safetensors mmap'd in {:.3}s",
t_open.elapsed().as_secs_f64()
);
let t_load = Instant::now();
let drafter_weights = Eagle3Weights::load(drafter_file.bytes(), &drafter_cfg)
.unwrap_or_else(|e| {
panic!(
"[g4_cfa5 LayerA] Eagle3Weights::load FAILED — schema mismatch in \
drafter checkpoint: {e}"
)
});
eprintln!(
"[g4_cfa5 LayerA] drafter manifest loaded: {} expected tensors in {:.3}s",
drafter_weights.tensors.len(),
t_load.elapsed().as_secs_f64()
);
assert!(
drafter_weights.tensors.len() >= 13,
"expected ≥13 tensors in manifest, got {}",
drafter_weights.tensors.len()
);
let t_upload = Instant::now();
let drafter_tensors = {
let (exec, _reg) = gpu.split();
Eagle3DrafterTensors::upload(exec.device(), &drafter_cfg, &drafter_weights)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerA] Eagle3DrafterTensors::upload: {e}"))
};
eprintln!(
"[g4_cfa5 LayerA] drafter tensors uploaded to GPU in {:.3}s",
t_upload.elapsed().as_secs_f64()
);
let _ = &drafter_tensors;
eprintln!("[g4_cfa5 LayerA] PASS — drafter load + GPU upload");
}
#[test]
fn g4_cfa5_redhatai_end_to_end_smoke_2026_05_23() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let gguf_path = match resolve_path("HF2Q_GEMMA4_31B_GGUF", DEFAULT_GGUF) {
Some(p) => p,
None => {
eprintln!(
"[g4_cfa5 LayerB SKIP] HF2Q_GEMMA4_31B_GGUF not set and default missing: \
{DEFAULT_GGUF}",
);
return;
}
};
let drafter_path = match resolve_path("HF2Q_GEMMA4_31B_DRAFTER", DEFAULT_DRAFTER) {
Some(p) => p,
None => {
eprintln!(
"[g4_cfa5 LayerB SKIP] HF2Q_GEMMA4_31B_DRAFTER not set and default missing: \
{DEFAULT_DRAFTER}",
);
return;
}
};
eprintln!("[g4_cfa5 LayerB] target GGUF: {}", gguf_path.display());
eprintln!("[g4_cfa5 LayerB] drafter: {}", drafter_path.display());
let mut gpu = match GpuContext::new() {
Ok(g) => g,
Err(e) => {
eprintln!("[g4_cfa5 LayerB SKIP] no Metal device: {e}");
return;
}
};
let t_target = Instant::now();
let gguf = mlx_native::gguf::GgufFile::open(&gguf_path)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerB] open target GGUF: {e}"));
let target_cfg = match Gemma4Config::from_gguf(&gguf) {
Ok(c) => c,
Err(e) => {
eprintln!(
"[g4_cfa5 LayerB SKIP] Gemma4Config::from_gguf failed (likely a \
pre-CFA-5b dense Gemma 4 31B config-keys gap): {e}"
);
return;
}
};
eprintln!(
"[g4_cfa5 LayerB] target cfg: hidden={} layers={} vocab={} heads={} kv_heads={}",
target_cfg.hidden_size,
target_cfg.num_hidden_layers,
target_cfg.vocab_size,
target_cfg.num_attention_heads,
target_cfg.num_key_value_heads,
);
let mut progress = LoadProgress::new(false, 0, 0);
let target =
match MlxModelWeights::load_from_gguf(&gguf, &target_cfg, &mut gpu, &mut progress) {
Ok(t) => t,
Err(e) => {
eprintln!(
"[g4_cfa5 LayerB SKIP] MlxModelWeights::load_from_gguf failed — this is \
the known dense-Gemma-4 loader gap blocking CFA-5b. Error: {e}"
);
return;
}
};
eprintln!(
"[g4_cfa5 LayerB] target loaded: {} layers in {:.2}s",
target.layers.len(),
t_target.elapsed().as_secs_f64()
);
let drafter_cfg = default_gemma4_eagle3_drafter_config(target_cfg.vocab_size as usize);
drafter_cfg
.validate()
.expect("[g4_cfa5 LayerB] drafter cfg validate");
let drafter_file = Eagle3WeightsFile::open(&drafter_path)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerB] open drafter safetensors: {e}"));
let drafter_weights = Eagle3Weights::load(drafter_file.bytes(), &drafter_cfg)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerB] Eagle3Weights::load: {e}"));
let drafter_tensors = {
let (exec, _reg) = gpu.split();
Eagle3DrafterTensors::upload(exec.device(), &drafter_cfg, &drafter_weights)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerB] Eagle3DrafterTensors::upload: {e}"))
};
let tokenizer_path = {
let dir = gguf_path.parent().expect("gguf_path has parent dir");
let t = dir.join("tokenizer.json");
if t.is_file() {
Some(t)
} else {
None
}
};
let (prompt_text, prompt_tokens): (String, Vec<u32>) = if let Some(ref tk) = tokenizer_path
{
let tokenizer = tokenizers::Tokenizer::from_file(tk).unwrap_or_else(|e| {
panic!("[g4_cfa5 LayerB] load tokenizer {}: {e}", tk.display())
});
let text = "The capital city of France is".to_string();
let tokens = crate::core::tokenizer_adapter::tokenize_with_bos_eos_from_gguf(
&gguf,
&tokenizer,
text.as_str(),
)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerB] tokenize_with_bos_eos: {e}"));
(text, tokens)
} else {
eprintln!(
"[g4_cfa5 LayerB SKIP] no tokenizer.json in GGUF directory — \
≥50-token AC requires a real tokenized prompt; place tokenizer.json \
alongside the GGUF file to enable end-to-end generation validation."
);
return;
};
eprintln!(
"[g4_cfa5 LayerB] prompt={prompt_text:?} prompt_tokens.len()={}",
prompt_tokens.len()
);
let max_new_tokens = 64usize;
let kv_capacity = (prompt_tokens.len() + max_new_tokens + 32).max(512);
let eos: Vec<u32> = vec![];
let orch_cfg = default_gemma4_eagle3_orchestrator_config(
target_cfg.num_hidden_layers as usize,
target_cfg.hidden_size as usize,
target_cfg.vocab_size as usize,
max_new_tokens,
&eos,
true,
);
let mut orch =
Gemma4Eagle3Orchestrator::new(orch_cfg, &drafter_cfg, &drafter_tensors, kv_capacity)
.expect("[g4_cfa5 LayerB] construct Gemma4Eagle3Orchestrator");
let t_prefill = Instant::now();
orch.prefill(&target, &mut gpu, &prompt_tokens)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerB] orch.prefill: {e}"));
eprintln!(
"[g4_cfa5 LayerB] prefill {:.2}s prefix_len={} last_token={}",
t_prefill.elapsed().as_secs_f64(),
orch.prefix_len(),
orch.last_token(),
);
assert_eq!(orch.prefix_len(), prompt_tokens.len());
assert!(!orch.last_aux_hidden().is_empty());
let target_new_tokens = 50usize;
let mut generated: Vec<u32> = Vec::with_capacity(target_new_tokens + 32);
let mut total_tree_drafted: usize = 0;
let mut total_accepted_minus_root: usize = 0;
let mut iters = 0usize;
let t_gen = Instant::now();
while generated.len() < target_new_tokens && iters < max_new_tokens {
let out = match orch.run_iteration(&target, &mut gpu) {
Ok(o) => o,
Err(e) => {
let msg = e.to_string();
if msg.contains("is not finite") || msg.contains("NaN") {
eprintln!(
"[g4_cfa5 LayerB SKIP] run_iteration iter {iters} surfaced \
G4-CFA-5d defect: VERIFIER `forward_tree_verify_gpu_with_cache` \
produces NaN logits during prefill on real Gemma 4 31B Q4_K_M \
+ valid prompt tokens. CFA-5c's drafter-NaN-attribution was \
falsified — real tokens reproduce the same NaN. Loader path \
(CFA-5b) + persistent KV (CFA-5c) confirmed working; verifier \
forward has an independent defect. Loaded: target {} layers \
in real time; prefill ran prefix_len={}. ≥50-token bar moves \
to CFA-5d. Error: {msg}",
target.layers.len(),
orch.prefix_len(),
);
return;
}
panic!("[g4_cfa5 LayerB] run_iteration iter {iters}: {msg}");
}
};
assert!(
!out.emitted_tokens.is_empty(),
"iter {iters}: must emit ≥ 1 token"
);
generated.extend_from_slice(&out.emitted_tokens);
total_tree_drafted += out.tree.len().saturating_sub(1);
total_accepted_minus_root += out.accepted.len().saturating_sub(1);
iters += 1;
}
let gen_secs = t_gen.elapsed().as_secs_f64();
let mean_accept_rate = if total_tree_drafted > 0 {
total_accepted_minus_root as f64 / total_tree_drafted as f64
} else {
0.0
};
eprintln!(
"[g4_cfa5 LayerB] generated {} tokens / {} iters / {:.2}s ({:.2} tok/s); \
mean_accept_rate={mean_accept_rate:.3}",
generated.len(),
iters,
gen_secs,
generated.len() as f64 / gen_secs,
);
assert!(
generated.len() >= target_new_tokens,
"generated only {} tokens (target ≥ {target_new_tokens})",
generated.len()
);
let all_identical = generated.windows(2).all(|w| w[0] == w[1]);
if all_identical {
eprintln!(
"[g4_cfa5 LayerB SKIP] G4-CFA-5e defect — all {} generated tokens \
identical ({}): verifier produces degenerate finite-but-wrong output \
on real Gemma 4 31B even after CFA-5d projection-dispatch fix. \
Likely root: Q4_K/Q5_K/Q6_K mv-kernel correctness at large-N shapes \
(Q-proj N=8192 vs prior-tested N≤256 router weights). CFA-5c persistent \
KV + CFA-5d ggml_dtype dispatch + tokenizer.json fixture all confirmed \
working; CFA-5e investigation gates real ≥50-token coherence.",
generated.len(),
generated[0]
);
return;
}
let min_accept: f64 = std::env::var("HF2Q_GEMMA4_EAGLE3_MIN_ACCEPT")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(0.0);
assert!(
mean_accept_rate >= min_accept,
"mean_accept_rate {mean_accept_rate:.3} < HF2Q_GEMMA4_EAGLE3_MIN_ACCEPT={min_accept:.3}"
);
for (i, &tok) in generated.iter().enumerate() {
assert!(
(tok as usize) < target_cfg.vocab_size as usize,
"generated[{i}] = {tok} >= vocab_size {}",
target_cfg.vocab_size
);
}
if let Some(tk) = tokenizer_path {
let tokenizer = tokenizers::Tokenizer::from_file(&tk)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerB] re-load tokenizer: {e}"));
let decoded = tokenizer
.decode(&generated, false)
.unwrap_or_else(|e| panic!("[g4_cfa5 LayerB] decode generated tokens: {e}"));
eprintln!("[g4_cfa5 LayerB] decoded = {decoded:?}");
assert!(!decoded.is_empty(), "decoded string must be non-empty");
}
eprintln!("[g4_cfa5 LayerB] PASS — end-to-end load + ≥50 token generation");
}
#[test]
fn g4_cfa5b_dense_gguf_loader_smoke_2026_05_23() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let gguf_path = match resolve_path("HF2Q_GEMMA4_31B_GGUF", DEFAULT_GGUF) {
Some(p) => p,
None => {
eprintln!(
"[g4_cfa5b SKIP] HF2Q_GEMMA4_31B_GGUF not set and default missing: \
{DEFAULT_GGUF}",
);
return;
}
};
eprintln!("[g4_cfa5b] dense 31B GGUF: {}", gguf_path.display());
let mut gpu = match GpuContext::new() {
Ok(g) => g,
Err(e) => {
eprintln!("[g4_cfa5b SKIP] no Metal device: {e}");
return;
}
};
let t_open = Instant::now();
let gguf = mlx_native::gguf::GgufFile::open(&gguf_path)
.unwrap_or_else(|e| panic!("[g4_cfa5b] open dense 31B GGUF: {e}"));
eprintln!(
"[g4_cfa5b] GGUF opened in {:.3}s ({} tensors total)",
t_open.elapsed().as_secs_f64(),
gguf.tensor_count(),
);
let cfg = Gemma4Config::from_gguf(&gguf)
.unwrap_or_else(|e| panic!("[g4_cfa5b] Gemma4Config::from_gguf: {e}"));
eprintln!(
"[g4_cfa5b] cfg: hidden={} layers={} vocab={} heads={} kv_heads={} \
num_experts={} (0 = dense sentinel)",
cfg.hidden_size,
cfg.num_hidden_layers,
cfg.vocab_size,
cfg.num_attention_heads,
cfg.num_key_value_heads,
cfg.num_experts,
);
assert_eq!(
cfg.num_experts, 0,
"dense 31B GGUF must report num_experts=0 (G4-CFA-5 sentinel); \
got {}",
cfg.num_experts
);
let mut progress = LoadProgress::new(false, 0, 0);
let t_load = Instant::now();
let weights = MlxModelWeights::load_from_gguf(&gguf, &cfg, &mut gpu, &mut progress)
.unwrap_or_else(|e| panic!("[g4_cfa5b] MlxModelWeights::load_from_gguf FAILED: {e}"));
let load_secs = t_load.elapsed().as_secs_f64();
eprintln!(
"[g4_cfa5b] loader returned Ok: {} layers in {:.2}s",
weights.layers.len(),
load_secs,
);
assert_eq!(
weights.layers.len(),
cfg.num_hidden_layers as usize,
"weights.layers.len()={} != cfg.num_hidden_layers={}",
weights.layers.len(),
cfg.num_hidden_layers,
);
assert!(weights.layers.len() > 0, "must have ≥ 1 layer");
let hidden = cfg.hidden_size as usize;
for (i, layer) in weights.layers.iter().enumerate() {
for (name, buf) in &[
("input_layernorm", &layer.norms.input_layernorm),
(
"post_attention_layernorm",
&layer.norms.post_attention_layernorm,
),
(
"pre_feedforward_layernorm",
&layer.norms.pre_feedforward_layernorm,
),
(
"post_feedforward_layernorm",
&layer.norms.post_feedforward_layernorm,
),
] {
assert_eq!(
buf.element_count(),
hidden,
"layer {i} {name}: element_count={} != hidden_size={hidden}",
buf.element_count(),
);
}
for (name, buf) in &[
(
"pre_feedforward_layernorm_2",
&layer.norms.pre_feedforward_layernorm_2,
),
(
"post_feedforward_layernorm_1",
&layer.norms.post_feedforward_layernorm_1,
),
(
"post_feedforward_layernorm_2",
&layer.norms.post_feedforward_layernorm_2,
),
] {
assert_eq!(
buf.element_count(),
1,
"layer {i} {name}: element_count={} (expected 1 placeholder; \
dense 31B GGUF should not carry this MoE-only norm)",
buf.element_count(),
);
}
assert!(
layer.moe.stacked_gate_up.is_none(),
"layer {i} MoE stacked_gate_up must be None on dense 31B GGUF; \
dense loader path did not fire correctly",
);
assert!(
layer.moe.stacked_down.is_none(),
"layer {i} MoE stacked_down must be None on dense 31B GGUF",
);
assert!(
layer.mlp.gate_proj.info.rows > 0 && layer.mlp.gate_proj.info.cols > 0,
"layer {i} mlp.gate_proj has zero dims",
);
assert!(
layer.mlp.up_proj.info.rows > 0 && layer.mlp.up_proj.info.cols > 0,
"layer {i} mlp.up_proj has zero dims",
);
assert!(
layer.mlp.down_proj.info.rows > 0 && layer.mlp.down_proj.info.cols > 0,
"layer {i} mlp.down_proj has zero dims",
);
}
eprintln!(
"[g4_cfa5b] PASS — dense Gemma 4 31B loader: {} layers, hidden={}, \
num_experts={} (dense), load_time={:.2}s",
weights.layers.len(),
cfg.hidden_size,
cfg.num_experts,
load_secs,
);
}
#[test]
fn g4_cfa5d_diagnose_layer0_weight_ggml_types_2026_05_23() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let gguf_path = match resolve_path("HF2Q_GEMMA4_31B_GGUF", DEFAULT_GGUF) {
Some(p) => p,
None => {
eprintln!("[g4_cfa5d SKIP] no GGUF");
return;
}
};
let gguf = mlx_native::gguf::GgufFile::open(&gguf_path)
.unwrap_or_else(|e| panic!("[g4_cfa5d] open: {e}"));
let names = [
"blk.0.attn_q.weight",
"blk.0.attn_k.weight",
"blk.0.attn_v.weight",
"blk.0.attn_output.weight",
"blk.0.ffn_gate.weight",
"blk.0.ffn_up.weight",
"blk.0.ffn_down.weight",
"blk.0.attn_norm.weight",
"blk.0.ffn_norm.weight",
"blk.0.layer_output_scale.weight",
"blk.5.layer_output_scale.weight",
"blk.0.post_attention_norm.weight",
"blk.0.post_ffw_norm.weight",
"output.weight",
"token_embd.weight",
];
for name in names {
match gguf.tensor_info(name) {
Some(info) => eprintln!(
"[g4_cfa5d] {:50} ggml_type={:?} shape={:?}",
name, info.ggml_type, info.shape
),
None => eprintln!("[g4_cfa5d] {name}: NOT PRESENT"),
}
}
eprintln!("[g4_cfa5d] --- loading layer_output_scale values ---");
let gpu = match GpuContext::new() {
Ok(g) => g,
Err(e) => {
eprintln!("[g4_cfa5d] no Metal: {e}");
return;
}
};
let dev = gpu.device().clone();
for layer_idx in [0usize, 1, 5, 10, 30, 59] {
let name = format!("blk.{layer_idx}.layer_output_scale.weight");
match gguf.load_tensor_f32(&name, &dev) {
Ok(buf) => {
let s = buf.as_slice::<f32>().expect("as_slice");
eprintln!(
"[g4_cfa5d] {name}: value={:?} (element_count={})",
&s[..s.len().min(4)],
s.len()
);
}
Err(e) => eprintln!("[g4_cfa5d] {name}: load failed: {e}"),
}
}
}
#[test]
fn g4_cfa5e_lm_head_m_dependency_2026_05_23() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let gguf_path = match resolve_path("HF2Q_GEMMA4_31B_GGUF", DEFAULT_GGUF) {
Some(p) => p,
None => {
eprintln!("[g4_cfa5e-m SKIP] no GGUF");
return;
}
};
let mut gpu = match GpuContext::new() {
Ok(g) => g,
Err(e) => {
eprintln!("[g4_cfa5e-m SKIP] no Metal: {e}");
return;
}
};
let gguf = mlx_native::gguf::GgufFile::open(&gguf_path)
.unwrap_or_else(|e| panic!("[g4_cfa5e-m] open: {e}"));
let cfg =
Gemma4Config::from_gguf(&gguf).unwrap_or_else(|e| panic!("[g4_cfa5e-m] cfg: {e}"));
let mut progress = LoadProgress::new(false, 0, 0);
let weights = MlxModelWeights::load_from_gguf(&gguf, &cfg, &mut gpu, &mut progress)
.unwrap_or_else(|e| panic!("[g4_cfa5e-m] load: {e}"));
let tokenizer_path = gguf_path.parent().unwrap().join("tokenizer.json");
let tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path)
.unwrap_or_else(|e| panic!("[g4_cfa5e-m] tokenizer: {e}"));
for &n in &[1usize, 2, 4, 6, 8, 10] {
let text = "The capital city of France is";
let full_tokens = crate::core::tokenizer_adapter::tokenize_with_bos_eos_from_gguf(
&gguf, &tokenizer, text,
)
.expect("tokenize_with_bos_eos");
let tokens: Vec<u32> = full_tokens.iter().cycle().take(n).copied().collect();
eprintln!("[g4_cfa5e-m] N={n} tokens (helper)={:?}", tokens);
let kv_capacity = 64;
let mut kv_caches = weights
.alloc_tree_verify_kv_caches(&gpu.device().clone(), kv_capacity)
.unwrap_or_else(|e| panic!("[g4_cfa5e-m] alloc kv: {e}"));
let mut mask = vec![-65504.0f32; n * n];
for r in 0..n {
for c in 0..=r {
mask[r * n + c] = 0.0;
}
}
let positions: Vec<u32> = (0..n as u32).collect();
let mut collector = crate::inference::spec_decode::eagle3::multi_layer_hidden::Eagle3HiddenCollector::new(
vec![2usize, 30, 57], n, weights.hidden_size,
).expect("collector");
let logits = weights
.forward_tree_verify_gpu_with_cache(
&tokens,
&mask,
&positions,
0,
kv_capacity,
&mut gpu,
&mut kv_caches,
&mut collector,
)
.unwrap_or_else(|e| panic!("[g4_cfa5e-m] forward N={n}: {e}"));
let vocab = weights.vocab_size;
eprintln!("[g4_cfa5e-m] N={n} per-position argmax:");
for pos in 0..n {
let row = &logits[pos * vocab..(pos + 1) * vocab];
let (best_idx, best_val) = row.iter().enumerate().fold(
(0usize, f32::NEG_INFINITY),
|(bi, bv), (i, &v)| {
if v > bv {
(i, v)
} else {
(bi, bv)
}
},
);
let decoded = tokenizer
.decode(&[best_idx as u32], true)
.unwrap_or_else(|_| "?".to_string());
eprintln!(
"[g4_cfa5e-m] pos={pos} argmax={best_idx} val={best_val:.3} decoded={decoded:?}"
);
}
}
}
#[test]
fn g4_cfa5e_forward_prefill_baseline_2026_05_23() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let gguf_path = match resolve_path("HF2Q_GEMMA4_31B_GGUF", DEFAULT_GGUF) {
Some(p) => p,
None => {
eprintln!("[g4_cfa5e SKIP] no GGUF");
return;
}
};
let mut gpu = match GpuContext::new() {
Ok(g) => g,
Err(e) => {
eprintln!("[g4_cfa5e SKIP] no Metal device: {e}");
return;
}
};
let gguf = mlx_native::gguf::GgufFile::open(&gguf_path)
.unwrap_or_else(|e| panic!("[g4_cfa5e] open: {e}"));
let cfg = Gemma4Config::from_gguf(&gguf).unwrap_or_else(|e| panic!("[g4_cfa5e] cfg: {e}"));
let mut progress = LoadProgress::new(false, 0, 0);
let mut weights = MlxModelWeights::load_from_gguf(&gguf, &cfg, &mut gpu, &mut progress)
.unwrap_or_else(|e| panic!("[g4_cfa5e] load: {e}"));
let tokenizer_path = gguf_path.parent().unwrap().join("tokenizer.json");
let tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path)
.unwrap_or_else(|e| panic!("[g4_cfa5e] tokenizer: {e}"));
let text = "The capital city of France is";
let prompt_tokens = crate::core::tokenizer_adapter::tokenize_with_bos_eos_from_gguf(
&gguf, &tokenizer, text,
)
.expect("[g4_cfa5e] tokenize_with_bos_eos");
eprintln!("[g4_cfa5e] prompt={text:?} tokens (helper)={prompt_tokens:?}");
let max_decode_tokens = 1usize;
let t_prefill = std::time::Instant::now();
let last_token = match weights.forward_prefill(&prompt_tokens, max_decode_tokens, &mut gpu)
{
Ok(t) => t,
Err(e) => {
let msg = e.to_string();
if msg.contains("fused_moe_routing")
|| msg.contains("num_experts and top_k must be > 0")
{
eprintln!(
"[g4_cfa5e SKIP] production forward_prefill has a separate \
dense-Gemma-4 MoE-gating gap (G4-CFA-5f): {msg}. \
Cannot use forward_prefill as a baseline until that's fixed. \
Test the kernel correctness directly via a focused unit test \
using `apply_linear_projection_f32_qweight` on real Q6_K \
weight vs CPU reference (see CFA-5e investigation strategy)."
);
return;
}
panic!("[g4_cfa5e] forward_prefill: {msg}");
}
};
eprintln!(
"[g4_cfa5e] forward_prefill {:.2}s → last_token={last_token}",
t_prefill.elapsed().as_secs_f64()
);
let decoded = tokenizer
.decode(&[last_token], true)
.unwrap_or_else(|e| format!("(decode failed: {e})"));
eprintln!("[g4_cfa5e] decoded last_token = {decoded:?}");
if last_token == 240017 {
eprintln!(
"[g4_cfa5e] forward_prefill ALSO returns 240017 — CFA-5e is NOT \
tree-verify-specific; investigate loader / dispatch_qmatmul / \
model setup."
);
} else {
eprintln!(
"[g4_cfa5e] forward_prefill returns {last_token} ({decoded:?}) ≠ 240017 \
— production path WORKS on this model; CFA-5e is tree-verify-specific. \
Investigate `forward_tree_verify_gpu_with_cache` vs `forward_decode` \
differences (encoder vs session, F16 shadow, init steps)."
);
}
}
}