use anyhow::{anyhow, ensure, Result};
#[derive(Debug, Clone)]
pub struct Eagle3DrafterConfig {
pub hidden_size: usize,
pub intermediate_size: usize,
pub head_dim: usize,
pub num_q_heads: usize,
pub num_kv_heads: usize,
pub vocab_size: usize,
pub draft_vocab_size: usize,
pub target_hidden_size: usize,
pub num_aux_hidden_states: usize,
pub rms_norm_eps: f32,
pub norm_before_fc: bool,
pub fc_norm: bool,
pub use_qk_norm: bool,
pub attention_bias: bool,
pub tie_lm_head: bool,
pub include_draft_id_mapping: bool,
pub has_own_embed_tokens: bool,
pub rope_theta: f32,
pub rope_dim: usize,
pub norm_before_residual: bool,
}
impl Eagle3DrafterConfig {
#[inline]
pub fn fc_input_size(&self) -> usize {
self.target_hidden_size * self.num_aux_hidden_states
}
#[inline]
pub fn q_proj_out(&self) -> usize {
self.num_q_heads * self.head_dim
}
#[inline]
pub fn kv_proj_out(&self) -> usize {
self.num_kv_heads * self.head_dim
}
#[inline]
pub fn qkv_input_width(&self) -> usize {
2 * self.hidden_size
}
pub fn validate(&self) -> Result<()> {
ensure!(self.hidden_size > 0, "hidden_size must be > 0");
ensure!(self.intermediate_size > 0, "intermediate_size must be > 0");
ensure!(self.head_dim > 0, "head_dim must be > 0");
ensure!(self.num_q_heads > 0, "num_q_heads must be > 0");
ensure!(self.num_kv_heads > 0, "num_kv_heads must be > 0");
ensure!(
self.num_q_heads % self.num_kv_heads == 0,
"num_q_heads ({}) must be divisible by num_kv_heads ({}) — GQA invariant",
self.num_q_heads,
self.num_kv_heads,
);
let _q_out = self.num_q_heads.checked_mul(self.head_dim).ok_or_else(|| {
anyhow!(
"num_q_heads * head_dim overflows usize (num_q_heads={}, head_dim={})",
self.num_q_heads,
self.head_dim
)
})?;
let _kv_out = self
.num_kv_heads
.checked_mul(self.head_dim)
.ok_or_else(|| {
anyhow!(
"num_kv_heads * head_dim overflows usize (num_kv_heads={}, head_dim={})",
self.num_kv_heads,
self.head_dim
)
})?;
ensure!(self.vocab_size > 0, "vocab_size must be > 0");
ensure!(self.draft_vocab_size > 0, "draft_vocab_size must be > 0");
ensure!(
self.draft_vocab_size <= self.vocab_size,
"draft_vocab_size ({}) must be <= vocab_size ({})",
self.draft_vocab_size,
self.vocab_size,
);
ensure!(
self.target_hidden_size > 0,
"target_hidden_size must be > 0"
);
ensure!(
self.num_aux_hidden_states > 0,
"num_aux_hidden_states must be > 0"
);
ensure!(
self.num_aux_hidden_states <= 64,
"num_aux_hidden_states must be <= 64 (matches Eagle3HiddenCollector)"
);
ensure!(self.rms_norm_eps > 0.0, "rms_norm_eps must be > 0");
ensure!(
self.rope_theta.is_finite() && self.rope_theta > 0.0,
"rope_theta ({}) must be finite and > 0",
self.rope_theta
);
ensure!(self.rope_dim > 0, "rope_dim must be > 0");
ensure!(
self.head_dim % 2 == 0,
"head_dim ({}) must be even (NeoX RoPE pairing requires head_dim/2)",
self.head_dim
);
ensure!(
self.rope_dim == self.head_dim,
"rope_dim ({}) must equal head_dim ({}) — partial rotation not supported by apply_imrope NeoX pairing",
self.rope_dim,
self.head_dim
);
ensure!(
self.rope_dim % 2 == 0,
"rope_dim ({}) must be even (RoPE rotates pairs)",
self.rope_dim
);
self.target_hidden_size
.checked_mul(self.num_aux_hidden_states)
.ok_or_else(|| anyhow!("fc_input_size overflow"))?;
self.hidden_size
.checked_mul(2)
.ok_or_else(|| anyhow!("qkv_input_width overflow"))?;
Ok(())
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
pub(crate) mod tests {
use super::*;
pub fn qwen35_default() -> Eagle3DrafterConfig {
Eagle3DrafterConfig {
hidden_size: 5120,
intermediate_size: 13824,
head_dim: 128,
num_q_heads: 40,
num_kv_heads: 8,
vocab_size: 152064,
draft_vocab_size: 152064,
target_hidden_size: 5120,
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,
}
}
#[test]
fn adr_037_e3b_qwen35_default_validates_2026_05_22() {
qwen35_default()
.validate()
.expect("default should validate");
}
#[test]
fn adr_037_e3b_fc_input_size_formula_2026_05_22() {
let cfg = qwen35_default();
assert_eq!(cfg.fc_input_size(), 5120 * 3);
}
#[test]
fn adr_037_e3b_qkv_input_width_is_2x_hidden_2026_05_22() {
let cfg = qwen35_default();
assert_eq!(cfg.qkv_input_width(), 2 * 5120);
}
#[test]
fn adr_037_e3b_gqa_invariant_enforced_2026_05_22() {
let mut cfg = qwen35_default();
cfg.num_kv_heads = 7; let err = cfg.validate().unwrap_err().to_string();
assert!(err.contains("GQA invariant"), "got: {err}");
}
#[test]
fn adr_038_g4_cfa5_llama_style_q_proj_out_validates_2026_05_23() {
let mut cfg = qwen35_default();
cfg.num_q_heads = 32; cfg.validate()
.expect("Llama-style q_proj_out != hidden_size must validate (ADR-038 G4-CFA-5)");
}
#[test]
fn adr_037_e3b_draft_vocab_size_at_most_target_vocab_2026_05_22() {
let mut cfg = qwen35_default();
cfg.draft_vocab_size = cfg.vocab_size + 1;
let err = cfg.validate().unwrap_err().to_string();
assert!(err.contains("must be <="), "got: {err}");
}
#[test]
fn adr_037_e3b_num_aux_at_most_64_2026_05_22() {
let mut cfg = qwen35_default();
cfg.num_aux_hidden_states = 65;
let err = cfg.validate().unwrap_err().to_string();
assert!(err.contains("<= 64"), "got: {err}");
}
#[test]
fn g4_cfa4_norm_before_residual_false_validates_2026_05_22() {
let mut cfg = qwen35_default();
cfg.norm_before_residual = false;
cfg.validate()
.expect("norm_before_residual=false must validate");
}
#[test]
fn g4_cfa4_norm_before_residual_true_validates_2026_05_22() {
let mut cfg = qwen35_default();
cfg.norm_before_residual = true;
cfg.validate()
.expect("norm_before_residual=true must validate");
}
#[test]
fn g4_cfa4_default_gemma4_eagle3_config_shape_2026_05_22() {
let cfg = Eagle3DrafterConfig {
hidden_size: 5376,
intermediate_size: 21504,
head_dim: 256,
num_q_heads: 32, num_kv_heads: 16, vocab_size: 262144,
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,
};
cfg.validate()
.expect("Gemma4 RedHatAI config shape must validate");
assert_eq!(
cfg.fc_input_size(),
5376 * 3,
"fc_input_size = 3 aux * 5376"
);
assert_eq!(cfg.norm_before_residual, true);
assert_eq!(
cfg.q_proj_out(),
8192,
"q_proj_out = num_q_heads(32) * head_dim(256)"
);
assert_eq!(
cfg.kv_proj_out(),
4096,
"kv_proj_out = num_kv_heads(16) * head_dim(256)"
);
}
}