use super::config::Eagle3DrafterConfig;
use super::weights::{Eagle3Weights, Eagle3WeightsError};
use mlx_native::{DType, MlxBuffer, MlxDevice, MlxError};
use safetensors::tensor::TensorView;
#[derive(Debug, thiserror::Error)]
pub enum Eagle3TensorsError {
#[error("eagle3 tensors mlx: {0}")]
Mlx(#[from] MlxError),
#[error("eagle3 tensors weights: {0}")]
Weights(#[from] Eagle3WeightsError),
#[error("eagle3 tensors: missing manifest entry `{0}`")]
MissingEntry(String),
}
pub struct Eagle3DrafterTensors {
pub embed_tokens: Option<MlxBuffer>,
pub fc: MlxBuffer,
pub input_norm: Option<MlxBuffer>,
pub fc_norm: Vec<MlxBuffer>,
pub norm: MlxBuffer,
pub lm_head: Option<MlxBuffer>,
pub draft_id_to_target_id: Option<Vec<i64>>,
pub input_layernorm: MlxBuffer,
pub hidden_norm: MlxBuffer,
pub post_attention_layernorm: MlxBuffer,
pub q_proj: MlxBuffer,
pub k_proj: MlxBuffer,
pub v_proj: MlxBuffer,
pub o_proj: MlxBuffer,
pub q_norm: Option<MlxBuffer>,
pub k_norm: Option<MlxBuffer>,
pub q_bias: Option<MlxBuffer>,
pub k_bias: Option<MlxBuffer>,
pub v_bias: Option<MlxBuffer>,
pub o_bias: Option<MlxBuffer>,
pub mlp_gate: MlxBuffer,
pub mlp_up: MlxBuffer,
pub mlp_down: MlxBuffer,
}
fn decode_bf16_bytes_to_f32(bytes: &[u8]) -> Result<Vec<f32>, Eagle3TensorsError> {
let n_elem = bytes.len() / 2;
if bytes.len() != n_elem * 2 {
return Err(Eagle3TensorsError::Mlx(MlxError::InvalidArgument(format!(
"decode_bf16_bytes_to_f32: data len {} not BF16-aligned (odd byte count)",
bytes.len()
))));
}
let mut out = Vec::with_capacity(n_elem);
for i in 0..n_elem {
let lo = bytes[i * 2] as u32;
let hi = bytes[i * 2 + 1] as u32;
let bf16_bits = lo | (hi << 8);
out.push(f32::from_bits(bf16_bits << 16));
}
Ok(out)
}
fn upload_bf16(device: &MlxDevice, view: &TensorView<'_>) -> Result<MlxBuffer, Eagle3TensorsError> {
let shape: Vec<usize> = view.shape().to_vec();
let byte_len = view.data().len();
let mut buf = device.alloc_buffer(byte_len, DType::BF16, shape)?;
let dst: &mut [u8] = buf.as_mut_slice::<u8>().map_err(|e| {
Eagle3TensorsError::Mlx(MlxError::InvalidArgument(format!("buffer slice: {e}")))
})?;
debug_assert_eq!(dst.len(), byte_len);
dst.copy_from_slice(view.data());
Ok(buf)
}
fn upload_bf16_as_f32(
device: &MlxDevice,
view: &TensorView<'_>,
) -> Result<MlxBuffer, Eagle3TensorsError> {
let f32_values = decode_bf16_bytes_to_f32(view.data())?;
let n_elem = f32_values.len();
let shape: Vec<usize> = view.shape().to_vec();
let mut buf = device.alloc_buffer(n_elem * 4, DType::F32, shape)?;
let dst: &mut [f32] = buf.as_mut_slice::<f32>().map_err(|e| {
Eagle3TensorsError::Mlx(MlxError::InvalidArgument(format!("f32 slice: {e}")))
})?;
debug_assert_eq!(dst.len(), n_elem);
dst.copy_from_slice(&f32_values);
Ok(buf)
}
fn decode_i64_le(view: &TensorView<'_>) -> Result<Vec<i64>, Eagle3TensorsError> {
let bytes = view.data();
if bytes.len() % 8 != 0 {
return Err(Eagle3TensorsError::Mlx(MlxError::InvalidArgument(format!(
"decode_i64_le: data len {} not I64-aligned (not multiple of 8)",
bytes.len()
))));
}
let n = bytes.len() / 8;
let mut out = Vec::with_capacity(n);
for i in 0..n {
let mut arr = [0u8; 8];
arr.copy_from_slice(&bytes[i * 8..i * 8 + 8]);
out.push(i64::from_le_bytes(arr));
}
Ok(out)
}
fn fetch<'a, 'b>(
weights: &'a Eagle3Weights<'b>,
name: &str,
) -> Result<&'a TensorView<'b>, Eagle3TensorsError> {
weights
.tensor(name)
.ok_or_else(|| Eagle3TensorsError::MissingEntry(name.to_string()))
}
impl Eagle3DrafterTensors {
pub fn upload(
device: &MlxDevice,
cfg: &Eagle3DrafterConfig,
weights: &Eagle3Weights<'_>,
) -> Result<Self, Eagle3TensorsError> {
let embed_tokens = if cfg.has_own_embed_tokens {
Some(upload_bf16(device, fetch(weights, "embed_tokens.weight")?)?)
} else {
None
};
let fc = upload_bf16(device, fetch(weights, "fc.weight")?)?;
let input_norm = if cfg.norm_before_fc {
Some(upload_bf16_as_f32(
device,
fetch(weights, "input_norm.weight")?,
)?)
} else {
None
};
let fc_norm = if cfg.fc_norm {
let mut v = Vec::with_capacity(cfg.num_aux_hidden_states);
for i in 0..cfg.num_aux_hidden_states {
v.push(upload_bf16_as_f32(
device,
fetch(weights, &format!("fc_norm.{i}.weight"))?,
)?);
}
v
} else {
Vec::new()
};
let norm = upload_bf16_as_f32(device, fetch(weights, "norm.weight")?)?;
let lm_head = if cfg.tie_lm_head {
None
} else {
Some(upload_bf16(device, fetch(weights, "lm_head.weight")?)?)
};
let draft_id_to_target_id = if cfg.include_draft_id_mapping {
Some(decode_i64_le(fetch(weights, "draft_id_to_target_id")?)?)
} else {
None
};
let input_layernorm =
upload_bf16_as_f32(device, fetch(weights, "layers.0.input_layernorm.weight")?)?;
let hidden_norm =
upload_bf16_as_f32(device, fetch(weights, "layers.0.hidden_norm.weight")?)?;
let post_attention_layernorm = upload_bf16_as_f32(
device,
fetch(weights, "layers.0.post_attention_layernorm.weight")?,
)?;
let q_proj = upload_bf16(device, fetch(weights, "layers.0.self_attn.q_proj.weight")?)?;
let k_proj = upload_bf16(device, fetch(weights, "layers.0.self_attn.k_proj.weight")?)?;
let v_proj = upload_bf16(device, fetch(weights, "layers.0.self_attn.v_proj.weight")?)?;
let o_proj = upload_bf16(device, fetch(weights, "layers.0.self_attn.o_proj.weight")?)?;
let (q_norm, k_norm) = if cfg.use_qk_norm {
(
Some(upload_bf16_as_f32(
device,
fetch(weights, "layers.0.self_attn.q_norm.weight")?,
)?),
Some(upload_bf16_as_f32(
device,
fetch(weights, "layers.0.self_attn.k_norm.weight")?,
)?),
)
} else {
(None, None)
};
let (q_bias, k_bias, v_bias, o_bias) = if cfg.attention_bias {
(
Some(upload_bf16_as_f32(
device,
fetch(weights, "layers.0.self_attn.q_proj.bias")?,
)?),
Some(upload_bf16_as_f32(
device,
fetch(weights, "layers.0.self_attn.k_proj.bias")?,
)?),
Some(upload_bf16_as_f32(
device,
fetch(weights, "layers.0.self_attn.v_proj.bias")?,
)?),
Some(upload_bf16_as_f32(
device,
fetch(weights, "layers.0.self_attn.o_proj.bias")?,
)?),
)
} else {
(None, None, None, None)
};
let mlp_gate = upload_bf16(device, fetch(weights, "layers.0.mlp.gate_proj.weight")?)?;
let mlp_up = upload_bf16(device, fetch(weights, "layers.0.mlp.up_proj.weight")?)?;
let mlp_down = upload_bf16(device, fetch(weights, "layers.0.mlp.down_proj.weight")?)?;
let tensors = Self {
embed_tokens,
fc,
input_norm,
fc_norm,
norm,
lm_head,
draft_id_to_target_id,
input_layernorm,
hidden_norm,
post_attention_layernorm,
q_proj,
k_proj,
v_proj,
o_proj,
q_norm,
k_norm,
q_bias,
k_bias,
v_bias,
o_bias,
mlp_gate,
mlp_up,
mlp_down,
};
debug_assert_eq!(tensors.fc.dtype(), DType::BF16);
debug_assert_eq!(tensors.norm.dtype(), DType::F32);
debug_assert_eq!(tensors.input_layernorm.dtype(), DType::F32);
debug_assert_eq!(tensors.hidden_norm.dtype(), DType::F32);
debug_assert_eq!(tensors.post_attention_layernorm.dtype(), DType::F32);
debug_assert_eq!(tensors.q_proj.dtype(), DType::BF16);
debug_assert_eq!(tensors.k_proj.dtype(), DType::BF16);
debug_assert_eq!(tensors.v_proj.dtype(), DType::BF16);
debug_assert_eq!(tensors.o_proj.dtype(), DType::BF16);
debug_assert_eq!(tensors.mlp_gate.dtype(), DType::BF16);
debug_assert_eq!(tensors.mlp_up.dtype(), DType::BF16);
debug_assert_eq!(tensors.mlp_down.dtype(), DType::BF16);
if let Some(b) = &tensors.embed_tokens {
debug_assert_eq!(b.dtype(), DType::BF16);
}
if let Some(b) = &tensors.input_norm {
debug_assert_eq!(b.dtype(), DType::F32);
}
for b in &tensors.fc_norm {
debug_assert_eq!(b.dtype(), DType::F32);
}
if let Some(b) = &tensors.lm_head {
debug_assert_eq!(b.dtype(), DType::BF16);
}
if let Some(b) = &tensors.q_norm {
debug_assert_eq!(b.dtype(), DType::F32);
}
if let Some(b) = &tensors.k_norm {
debug_assert_eq!(b.dtype(), DType::F32);
}
for opt in [
&tensors.q_bias,
&tensors.k_bias,
&tensors.v_bias,
&tensors.o_bias,
] {
if let Some(b) = opt {
debug_assert_eq!(b.dtype(), DType::F32);
}
}
Ok(tensors)
}
pub fn gpu_resident_bytes(&self) -> usize {
let mut total = 0;
total += self.fc.byte_len();
total += self.norm.byte_len();
total += self.input_layernorm.byte_len();
total += self.hidden_norm.byte_len();
total += self.post_attention_layernorm.byte_len();
total += self.q_proj.byte_len();
total += self.k_proj.byte_len();
total += self.v_proj.byte_len();
total += self.o_proj.byte_len();
total += self.mlp_gate.byte_len();
total += self.mlp_up.byte_len();
total += self.mlp_down.byte_len();
if let Some(b) = &self.embed_tokens {
total += b.byte_len();
}
if let Some(b) = &self.input_norm {
total += b.byte_len();
}
for b in &self.fc_norm {
total += b.byte_len();
}
if let Some(b) = &self.lm_head {
total += b.byte_len();
}
if let Some(b) = &self.q_norm {
total += b.byte_len();
}
if let Some(b) = &self.k_norm {
total += b.byte_len();
}
if let Some(b) = &self.q_bias {
total += b.byte_len();
}
if let Some(b) = &self.k_bias {
total += b.byte_len();
}
if let Some(b) = &self.v_bias {
total += b.byte_len();
}
if let Some(b) = &self.o_bias {
total += b.byte_len();
}
total
}
pub fn cpu_resident_bytes(&self) -> usize {
self.draft_id_to_target_id
.as_ref()
.map_or(0, |v| v.len() * std::mem::size_of::<i64>())
}
pub fn total_resident_bytes(&self) -> usize {
self.gpu_resident_bytes() + self.cpu_resident_bytes()
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::*;
use crate::inference::spec_decode::eagle3::weights::{
expected_manifest, Eagle3Weights, ExpectedTensor,
};
use safetensors::tensor::Dtype as SafeDtype;
use std::collections::BTreeMap;
fn bf16_bytes_from_f32(values: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(values.len() * 2);
for v in values {
let bf16_bits = (v.to_bits() >> 16) as u16;
out.push((bf16_bits & 0xff) as u8);
out.push(((bf16_bits >> 8) & 0xff) as u8);
}
out
}
#[test]
fn adr_037_e4b_decode_bf16_round_trips_canonical_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let canonical = [0.0f32, 1.0, -1.0, 2.0, 0.5, -0.5];
let bytes = bf16_bytes_from_f32(&canonical);
let decoded = decode_bf16_bytes_to_f32(&bytes).expect("decode ok");
for (i, (got, want)) in decoded.iter().zip(canonical.iter()).enumerate() {
assert_eq!(
got.to_bits(),
want.to_bits(),
"canonical[{i}] = {want} round-trip got {got}"
);
}
}
#[test]
fn adr_037_e4b_decode_bf16_rejects_odd_length_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err = decode_bf16_bytes_to_f32(&[0x80, 0x3f, 0x00]).unwrap_err();
assert!(format!("{err}").contains("not BF16-aligned"), "got: {err}");
}
fn build_synthetic_safetensors(manifest: &[ExpectedTensor]) -> Vec<u8> {
let mut storage: Vec<Vec<u8>> = Vec::with_capacity(manifest.len());
let mut tensors: BTreeMap<String, TensorView> = BTreeMap::new();
for exp in manifest {
let elem_bytes = match exp.dtype {
SafeDtype::BF16 => 2,
SafeDtype::I64 => 8,
_ => panic!("unexpected dtype in test"),
};
let nelem: usize = exp.shape.iter().product();
storage.push(vec![0u8; nelem * elem_bytes]);
}
for (i, exp) in manifest.iter().enumerate() {
let view = TensorView::new(exp.dtype, exp.shape.clone(), storage[i].as_slice())
.expect("synthetic tensor view");
tensors.insert(exp.name.clone(), view);
}
safetensors::serialize(&tensors, None::<std::collections::HashMap<String, String>>)
.expect("serialize synthetic")
}
fn tiny_cfg() -> Eagle3DrafterConfig {
Eagle3DrafterConfig {
hidden_size: 256,
intermediate_size: 512,
head_dim: 32,
num_q_heads: 8,
num_kv_heads: 4,
vocab_size: 1000,
draft_vocab_size: 1000,
target_hidden_size: 256,
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: 32,
norm_before_residual: false,
}
}
#[test]
fn adr_037_e4b1_upload_default_qwen35_config_succeeds_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return, };
let cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_synthetic_safetensors(&manifest);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload to GPU");
assert_eq!(tensors.fc.dtype(), DType::BF16);
assert_eq!(
tensors.norm.dtype(),
DType::F32,
"norm cast to F32 on upload"
);
assert_eq!(tensors.input_layernorm.dtype(), DType::F32);
assert_eq!(tensors.hidden_norm.dtype(), DType::F32);
assert_eq!(tensors.post_attention_layernorm.dtype(), DType::F32);
assert_eq!(tensors.q_proj.dtype(), DType::BF16);
assert!(
tensors.embed_tokens.is_some(),
"has_own_embed_tokens = true"
);
assert!(tensors.input_norm.is_none(), "norm_before_fc = false");
assert_eq!(tensors.fc_norm.len(), 3, "fc_norm = true, num_aux = 3");
assert!(tensors.q_norm.is_some(), "use_qk_norm = true");
assert!(tensors.k_norm.is_some(), "use_qk_norm = true");
assert!(tensors.q_bias.is_none(), "attention_bias = false");
assert!(tensors.lm_head.is_some(), "tie_lm_head = false");
assert!(
tensors.draft_id_to_target_id.is_some(),
"include_draft_id_mapping = true"
);
}
#[test]
fn adr_037_e4b1_upload_all_gates_off_minimum_tensors_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut cfg = tiny_cfg();
cfg.norm_before_fc = false;
cfg.fc_norm = false;
cfg.use_qk_norm = false;
cfg.attention_bias = false;
cfg.tie_lm_head = true;
cfg.include_draft_id_mapping = false;
cfg.has_own_embed_tokens = false;
let manifest = expected_manifest(&cfg);
let blob = build_synthetic_safetensors(&manifest);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors =
Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("minimum-config upload");
assert!(tensors.embed_tokens.is_none());
assert!(tensors.input_norm.is_none());
assert!(tensors.fc_norm.is_empty());
assert!(tensors.lm_head.is_none());
assert!(tensors.draft_id_to_target_id.is_none());
assert!(tensors.q_norm.is_none() && tensors.k_norm.is_none());
assert!(tensors.q_bias.is_none());
}
#[test]
fn adr_037_e4b1_upload_all_gates_on_maximum_tensors_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut cfg = tiny_cfg();
cfg.norm_before_fc = true;
cfg.fc_norm = true;
cfg.use_qk_norm = true;
cfg.attention_bias = true;
cfg.tie_lm_head = false;
cfg.include_draft_id_mapping = true;
cfg.has_own_embed_tokens = true;
let manifest = expected_manifest(&cfg);
let blob = build_synthetic_safetensors(&manifest);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors =
Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("maximum-config upload");
assert!(tensors.embed_tokens.is_some());
assert!(tensors.input_norm.is_some());
assert_eq!(tensors.fc_norm.len(), cfg.num_aux_hidden_states);
assert!(tensors.lm_head.is_some());
assert!(tensors.draft_id_to_target_id.is_some());
assert!(tensors.q_norm.is_some() && tensors.k_norm.is_some());
assert!(tensors.q_bias.is_some());
assert!(tensors.k_bias.is_some());
assert!(tensors.v_bias.is_some());
assert!(tensors.o_bias.is_some());
}
#[test]
fn adr_037_e4b1_gpu_resident_bytes_includes_f32_cast_expansion_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_synthetic_safetensors(&manifest);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let safetensors_data_bytes: usize = weights.tensors.iter().map(|t| t.data().len()).sum();
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload to GPU");
let total_bytes = tensors.total_resident_bytes();
let cast_elems = cfg.hidden_size + cfg.hidden_size + cfg.hidden_size + cfg.head_dim + cfg.head_dim + cfg.num_aux_hidden_states * cfg.target_hidden_size + cfg.hidden_size; let cast_expansion_bytes = cast_elems * 2; assert_eq!(
total_bytes,
safetensors_data_bytes + cast_expansion_bytes,
"total bytes (gpu+cpu) = safetensors data ({safetensors_data_bytes}) + F32 cast expansion ({cast_expansion_bytes})"
);
assert_eq!(
tensors.gpu_resident_bytes() + tensors.cpu_resident_bytes(),
total_bytes
);
assert_eq!(
tensors.cpu_resident_bytes(),
cfg.draft_vocab_size * 8,
"CPU bytes should be draft_vocab_size * 8 (i64 mapping)"
);
}
}