use crate::Engine;
use crate::cache::{Cache, KvLayer};
use crate::forward::argmax;
use crate::hybrid::HybridModel;
use crate::model::GpuTensor;
use cudarc::driver::CudaSlice;
use memra_gguf::dequant;
use memra_gguf::safetensors::StModel;
use std::path::Path;
pub struct Eagle3Draft {
pub fc: GpuTensor, pub input_layernorm: GpuTensor, pub hidden_norm: GpuTensor, pub q_proj: GpuTensor, pub k_proj: GpuTensor, pub v_proj: GpuTensor, pub o_proj: GpuTensor, pub post_attention_layernorm: GpuTensor,
pub gate_proj: GpuTensor,
pub up_proj: GpuTensor,
pub down_proj: GpuTensor,
pub norm: GpuTensor, pub lm_head: GpuTensor, pub d2t: Vec<i64>,
pub n_embd: usize,
pub n_head: usize,
pub n_head_kv: usize,
pub head_dim: usize,
pub n_ff: usize,
pub draft_vocab: usize,
pub rope_dim_count: usize, pub rope_theta: f32, pub eps: f32,
pub aux_layers: Vec<usize>, }
fn load_float(
e: &Engine,
m: &StModel,
name: &str,
) -> Result<GpuTensor, Box<dyn std::error::Error>> {
let (info, bytes) = m
.raw(name)
.ok_or_else(|| format!("EAGLE3 draft missing tensor {name}"))?;
let ne = info.ne(); let n: u64 = ne.iter().product();
let f32v = dequant::dequantize(info.ggml_type(), bytes, n as usize);
Ok(GpuTensor::Float {
data: e.htod(&f32v)?,
ne,
})
}
impl Eagle3Draft {
pub fn load(e: &Engine, path: &Path) -> Result<Self, Box<dyn std::error::Error>> {
let dir = if path.is_file() {
path.parent().unwrap_or(Path::new("."))
} else {
path
};
let cfg = EagleConfig::from_json(&dir.join("config.json"))?;
let m = StModel::open(path)?;
let d2t = read_i64(&m, "d2t")?;
assert_eq!(d2t.len(), cfg.draft_vocab, "d2t len != draft_vocab_size");
let draft = Eagle3Draft {
fc: load_float(e, &m, "fc.weight")?,
input_layernorm: load_float(e, &m, "midlayer.input_layernorm.weight")?,
hidden_norm: load_float(e, &m, "midlayer.hidden_norm.weight")?,
q_proj: load_float(e, &m, "midlayer.self_attn.q_proj.weight")?,
k_proj: load_float(e, &m, "midlayer.self_attn.k_proj.weight")?,
v_proj: load_float(e, &m, "midlayer.self_attn.v_proj.weight")?,
o_proj: load_float(e, &m, "midlayer.self_attn.o_proj.weight")?,
post_attention_layernorm: load_float(
e,
&m,
"midlayer.post_attention_layernorm.weight",
)?,
gate_proj: load_float(e, &m, "midlayer.mlp.gate_proj.weight")?,
up_proj: load_float(e, &m, "midlayer.mlp.up_proj.weight")?,
down_proj: load_float(e, &m, "midlayer.mlp.down_proj.weight")?,
norm: load_float(e, &m, "norm.weight")?,
lm_head: load_float(e, &m, "lm_head.weight")?,
d2t,
n_embd: cfg.hidden_size,
n_head: cfg.n_head,
n_head_kv: cfg.n_head_kv,
head_dim: cfg.head_dim,
n_ff: cfg.intermediate_size,
draft_vocab: cfg.draft_vocab,
rope_dim_count: ((cfg.partial_rotary_factor * cfg.head_dim as f32).round() as usize)
.max(2),
rope_theta: cfg.rope_theta,
eps: cfg.rms_eps,
aux_layers: cfg.aux_layers,
};
assert_eq!(
draft.fc.in_features(),
3 * draft.n_embd,
"fc in != 3*n_embd"
);
assert_eq!(draft.fc.out_features(), draft.n_embd, "fc out != n_embd");
assert_eq!(
draft.q_proj.in_features(),
2 * draft.n_embd,
"q_proj in != 2*n_embd"
);
assert_eq!(
draft.q_proj.out_features(),
draft.n_head * draft.head_dim,
"q_proj out"
);
assert_eq!(
draft.lm_head.out_features(),
draft.draft_vocab,
"lm_head out != draft_vocab"
);
Ok(draft)
}
#[inline]
pub fn d2t_map(&self, draft_id: u32) -> u32 {
(draft_id as i64 + self.d2t[draft_id as usize]) as u32
}
pub fn encode(
&self,
e: &Engine,
aux: &[CudaSlice<f32>],
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
assert_eq!(aux.len(), self.aux_layers.len(), "aux count != #aux layers");
let n = self.n_embd;
let mut cat = e.zeros(self.aux_layers.len() * n)?;
for (i, a) in aux.iter().enumerate() {
e.copy_into(&mut cat, i * n, a, n)?;
}
e.matmul(&self.fc, &cat, 1) }
pub fn draft_token(
&self,
e: &Engine,
target: &HybridModel,
prev_tok: u32,
g: &CudaSlice<f32>,
scratch: &mut Eagle3Scratch,
pos: usize,
) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let n = self.n_embd;
let eps = self.eps;
let pos_d = e.htod_i32(&[pos as i32])?;
let e_emb = e.htod(&target.embd.gather(n, &[prev_tok]))?;
let mut e_norm = e.zeros(n)?;
e.rms_norm(
&e_emb,
self.input_layernorm.float_data(),
&mut e_norm,
n,
1,
eps,
)?;
let res = e.clone_dtod(g)?;
let mut g_norm = e.zeros(n)?;
e.rms_norm(g, self.hidden_norm.float_data(), &mut g_norm, n, 1, eps)?;
let mut cat = e.zeros(2 * n)?;
e.copy_into(&mut cat, 0, &e_norm, n)?;
e.copy_into(&mut cat, n, &g_norm, n)?;
let attn = self.attn(e, &cat, &pos_d, scratch)?;
let mut x1 = e.zeros(n)?;
e.add(&attn, &res, &mut x1, n)?;
let mut z = e.zeros(n)?;
e.rms_norm(
&x1,
self.post_attention_layernorm.float_data(),
&mut z,
n,
1,
eps,
)?;
let gate = e.matmul(&self.gate_proj, &z, 1)?;
let up = e.matmul(&self.up_proj, &z, 1)?;
let mut act = e.zeros(self.n_ff)?;
e.silu_mul(&gate, &up, &mut act, self.n_ff)?;
let mlp = e.matmul(&self.down_proj, &act, 1)?;
let mut g_next = e.zeros(n)?;
e.add(&mlp, &x1, &mut g_next, n)?;
let mut hn = e.zeros(n)?;
e.rms_norm(&g_next, self.norm.float_data(), &mut hn, n, 1, eps)?;
let logits = e.matmul(&self.lm_head, &hn, 1)?;
let host = e.dtoh(&logits)?;
Ok((host, g_next))
}
fn attn(
&self,
e: &Engine,
cat: &CudaSlice<f32>,
pos_d: &CudaSlice<i32>,
scratch: &mut Eagle3Scratch,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (nh, nhkv, hd) = (self.n_head, self.n_head_kv, self.head_dim);
let scale = 1.0 / (hd as f32).sqrt();
let mut q = e.matmul(&self.q_proj, cat, 1)?; let mut k = e.matmul(&self.k_proj, cat, 1)?; let v = e.matmul(&self.v_proj, cat, 1)?;
e.rope_neox(
&mut q,
pos_d,
hd,
self.rope_dim_count,
nh,
1,
self.rope_theta,
1.0,
)?;
e.rope_neox(
&mut k,
pos_d,
hd,
self.rope_dim_count,
nhkv,
1,
self.rope_theta,
1.0,
)?;
let kv = &mut scratch.kv;
e.append_kv_quantized(
&k,
&v,
&mut kv.k,
&mut kv.v,
kv.len,
kv.kv_dim_k,
kv.kv_dim_v,
kv.k_tok_bytes,
kv.v_tok_bytes,
false,
)?;
kv.len += 1;
let t_kv = kv.len;
let (ktb, vtb) = (kv.k_tok_bytes, kv.v_tok_bytes);
let k_view = e.view_u8(&kv.k, t_kv * ktb);
let v_view = e.view_u8(&kv.v, t_kv * vtb);
let mut attn = e.zeros(nh * hd)?;
e.fa_decode(
&q, &k_view, &v_view, &mut attn, hd, nh, nhkv, t_kv, scale, ktb, vtb,
)?;
e.matmul(&self.o_proj, &attn, 1)
}
}
pub struct Eagle3Scratch {
pub kv: KvLayer,
}
impl Eagle3Scratch {
pub fn new(
e: &Engine,
draft: &Eagle3Draft,
cap: usize,
) -> Result<Self, Box<dyn std::error::Error>> {
let (nhkv, hd) = (draft.n_head_kv, draft.head_dim);
assert!(
hd % 32 == 0,
"KVQUANT requires head_dim%32==0 (EAGLE3 scratch)"
);
let kv_dim_k = hd * nhkv;
let kv_dim_v = hd * nhkv;
let (kbb, vbb) = crate::kv_blk_bytes(); let k_tok_bytes = (kv_dim_k / 32) * kbb;
let v_tok_bytes = (kv_dim_v / 32) * vbb;
Ok(Eagle3Scratch {
kv: KvLayer {
k: e.alloc_u8(cap * k_tok_bytes)?,
v: e.alloc_u8(cap * v_tok_bytes)?,
kv_dim_k,
kv_dim_v,
k_tok_bytes,
v_tok_bytes,
len: 0,
ring: None,
len_d: e.htod_i32(&[0])?,
},
})
}
pub fn reset(&mut self) {
self.kv.len = 0;
}
}
impl HybridModel {
pub fn generate_spec_eagle(
&self,
e: &Engine,
draft: &Eagle3Draft,
prompt: &[u32],
max_new: usize,
k: usize,
) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
assert!(k >= 1, "k must be >= 1");
assert!(!prompt.is_empty(), "prompt must be non-empty");
let n_vocab = self.output.out_features();
let n_embd = self.cfg.n_embd as usize;
assert_eq!(n_embd, draft.n_embd, "draft n_embd != target n_embd");
let aux = &draft.aux_layers;
let max_ctx = prompt.len() + max_new + k + 8;
let mut cache = Cache::new(e, &self.cfg, max_ctx)?;
let mut prime_logits = Vec::new();
let mut prime_aux: Vec<CudaSlice<f32>> = Vec::new();
for &tok in prompt {
let (l, a) = self.decode_step_aux(e, tok, &mut cache, aux)?;
prime_logits = l;
prime_aux = a;
}
let mut scratch = Eagle3Scratch::new(e, draft, k + 1)?;
let mut out: Vec<u32> = Vec::with_capacity(max_new);
let mut total_drafted = 0usize;
let mut total_accepted = 0usize;
let shift = std::env::var("MEMRA_EAGLE_ALIGN")
.ok()
.map(|s| s != "0")
.unwrap_or(true);
let mut last_token = argmax(&prime_logits) as u32;
out.push(last_token);
let mut prev_aux = prime_aux;
let (mut last_logits, mut g_aux) = self.decode_step_aux(e, last_token, &mut cache, aux)?;
while out.len() < max_new {
let pos = cache.pos;
let snap = cache.snapshot(e)?;
let seed_aux = if shift { &prev_aux } else { &g_aux };
let g0 = draft.encode(e, seed_aux)?;
scratch.reset();
let mut draft_toks: Vec<u32> = Vec::with_capacity(k);
let mut prev = last_token;
let mut g = g0;
for j in 0..k {
let (dl, g_next) = draft.draft_token(e, self, prev, &g, &mut scratch, pos + j)?;
let d_draft = argmax(&dl) as u32;
let d_target = draft.d2t_map(d_draft); draft_toks.push(d_target);
prev = d_target;
g = g_next;
}
let tlogits = self.decode_step_t(e, &draft_toks, pos, &mut cache)?;
let t_pred = |j: usize| -> u32 {
if j == 0 {
argmax(&last_logits) as u32
} else {
argmax(&tlogits[(j - 1) * n_vocab..j * n_vocab]) as u32
}
};
let mut n_acc = 0usize;
for j in 0..k {
if t_pred(j) == draft_toks[j] {
n_acc += 1;
} else {
break;
}
}
let bonus = t_pred(n_acc);
total_drafted += k;
total_accepted += n_acc;
for j in 0..n_acc {
if out.len() >= max_new {
break;
}
out.push(draft_toks[j]);
}
let bonus_emitted = out.len() < max_new;
if bonus_emitted {
out.push(bonus);
}
last_token = bonus;
let pred_is_prev_round = n_acc == 0; let old_g_aux = std::mem::take(&mut g_aux); cache.rollback(e, &snap, 0)?;
let mut replay: Vec<u32> = draft_toks[0..n_acc].to_vec();
replay.push(bonus);
let pred_col = if pred_is_prev_round {
None
} else {
Some(replay.len() - 2)
};
let (rl, mut a_last, a_pred) =
self.decode_step_t_aux2(e, &replay, pos, &mut cache, aux, pred_col)?;
last_logits = rl[(replay.len() - 1) * n_vocab..replay.len() * n_vocab].to_vec();
prev_aux = if pred_is_prev_round {
old_g_aux
} else {
a_pred.unwrap()
};
g_aux = std::mem::take(&mut a_last);
}
out.truncate(max_new);
Ok((out, total_drafted, total_accepted))
}
}
struct EagleConfig {
hidden_size: usize,
n_head: usize,
n_head_kv: usize,
head_dim: usize,
intermediate_size: usize,
draft_vocab: usize,
partial_rotary_factor: f32,
rope_theta: f32,
rms_eps: f32,
aux_layers: Vec<usize>,
}
impl EagleConfig {
fn from_json(path: &Path) -> Result<Self, Box<dyn std::error::Error>> {
let txt = std::fs::read_to_string(path)?;
let num = |key: &str| -> Option<f64> {
let pat = format!("\"{key}\"");
let i = txt.find(&pat)? + pat.len();
let rest = &txt[i..];
let c = rest.find(':')? + 1;
let tail = rest[c..].trim_start();
let end = tail
.find(|ch: char| ch == ',' || ch == '}' || ch == '\n')
.unwrap_or(tail.len());
tail[..end].trim().parse::<f64>().ok()
};
let aux_layers: Vec<usize> = {
let pat = "\"eagle_aux_hidden_state_layer_ids\"";
match txt.find(pat) {
Some(i) => {
let rest = &txt[i + pat.len()..];
let lb = rest.find('[').ok_or("no [ after aux ids")?;
let rb = rest.find(']').ok_or("no ] after aux ids")?;
rest[lb + 1..rb]
.split(',')
.filter_map(|s| s.trim().parse::<usize>().ok())
.collect()
}
None => vec![1, 15, 28], }
};
Ok(EagleConfig {
hidden_size: num("hidden_size").ok_or("hidden_size")? as usize,
n_head: num("num_attention_heads").ok_or("num_attention_heads")? as usize,
n_head_kv: num("num_key_value_heads").ok_or("num_key_value_heads")? as usize,
head_dim: num("head_dim").ok_or("head_dim")? as usize,
intermediate_size: num("intermediate_size").ok_or("intermediate_size")? as usize,
draft_vocab: num("draft_vocab_size").ok_or("draft_vocab_size")? as usize,
partial_rotary_factor: num("partial_rotary_factor").unwrap_or(1.0) as f32,
rope_theta: num("rope_theta").unwrap_or(10000.0) as f32,
rms_eps: num("rms_norm_eps").unwrap_or(1e-6) as f32,
aux_layers,
})
}
}
fn read_i64(m: &StModel, name: &str) -> Result<Vec<i64>, Box<dyn std::error::Error>> {
let (info, bytes) = m
.raw(name)
.ok_or_else(|| format!("EAGLE3 draft missing {name}"))?;
assert_eq!(info.dtype, "I64", "{name} dtype != I64");
let n = bytes.len() / 8;
let mut v = Vec::with_capacity(n);
for i in 0..n {
v.push(i64::from_le_bytes(
bytes[i * 8..i * 8 + 8].try_into().unwrap(),
));
}
Ok(v)
}