1use crate::Engine;
41use crate::cache::{Cache, KvLayer};
42use crate::forward::argmax;
43use crate::hybrid::HybridModel;
44use crate::model::GpuTensor;
45use cudarc::driver::CudaSlice;
46use memra_gguf::dequant;
47use memra_gguf::safetensors::StModel;
48use std::path::Path;
49
50pub struct Eagle3Draft {
55 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,
63 pub gate_proj: GpuTensor,
64 pub up_proj: GpuTensor,
65 pub down_proj: GpuTensor,
66 pub norm: GpuTensor, pub lm_head: GpuTensor, pub d2t: Vec<i64>, pub n_embd: usize,
72 pub n_head: usize,
73 pub n_head_kv: usize,
74 pub head_dim: usize,
75 pub n_ff: usize,
76 pub draft_vocab: usize,
77 pub rope_dim_count: usize, pub rope_theta: f32, pub eps: f32,
80 pub aux_layers: Vec<usize>, }
82
83fn validate_eagle_tensor(
84 name: &str,
85 info: &memra_gguf::safetensors::StInfo,
86 expected_ne: &[u64],
87) -> Result<Vec<u64>, String> {
88 let ne = info.ne();
89 if !matches!(info.dtype.as_str(), "BF16" | "F32") {
90 return Err(format!(
91 "EAGLE3 tensor {name} has dtype {}, expected BF16 or F32",
92 info.dtype
93 ));
94 }
95 if ne != expected_ne {
96 return Err(format!(
97 "EAGLE3 tensor {name} has shape {ne:?}, expected {expected_ne:?}"
98 ));
99 }
100 Ok(ne)
101}
102
103fn validate_aux_layers(aux_layers: &[usize]) -> Result<(), String> {
104 if aux_layers.len() != 3 {
105 return Err(format!(
106 "EAGLE3 fc.weight consumes exactly three auxiliary hidden states; config declares {}",
107 aux_layers.len()
108 ));
109 }
110 if aux_layers.windows(2).any(|pair| pair[0] >= pair[1]) {
111 return Err("EAGLE3 auxiliary layer ids must be strictly increasing".into());
112 }
113 Ok(())
114}
115
116fn validate_eagle_attention_geometry(
117 n_head: usize,
118 n_head_kv: usize,
119 head_dim: usize,
120) -> Result<(), String> {
121 if n_head == 0
122 || n_head_kv == 0
123 || head_dim == 0
124 || !n_head.is_multiple_of(n_head_kv)
125 || !head_dim.is_multiple_of(32)
126 {
127 return Err(format!(
128 "EAGLE3 attention geometry requires nonzero n_head divisible by n_head_kv and head_dim divisible by 32; got n_head={n_head}, n_head_kv={n_head_kv}, head_dim={head_dim}"
129 ));
130 }
131 n_head
132 .checked_mul(head_dim)
133 .ok_or("EAGLE3 query-head geometry overflow")?;
134 n_head_kv
135 .checked_mul(head_dim)
136 .ok_or("EAGLE3 key/value-head geometry overflow")?;
137 Ok(())
138}
139
140fn validate_d2t_map(
141 d2t: &[i64],
142 draft_vocab: usize,
143 target_vocab: Option<usize>,
144) -> Result<(), String> {
145 if d2t.len() != draft_vocab {
146 return Err(format!(
147 "EAGLE3 d2t has {} entries, expected draft_vocab_size {draft_vocab}",
148 d2t.len()
149 ));
150 }
151 for (draft_id, delta) in d2t.iter().copied().enumerate() {
152 let target = i64::try_from(draft_id)
153 .ok()
154 .and_then(|id| id.checked_add(delta))
155 .filter(|target| (0..=i64::from(u32::MAX)).contains(target))
156 .filter(|target| {
157 target_vocab.is_none_or(|vocab| usize::try_from(*target).is_ok_and(|id| id < vocab))
158 });
159 if target.is_none() {
160 let limit = target_vocab
161 .map(|vocab| format!(" target vocabulary of {vocab} entries"))
162 .unwrap_or_else(|| " target u32 vocabulary".to_string());
163 return Err(format!(
164 "EAGLE3 d2t[{draft_id}]={delta} maps outside the{limit}"
165 ));
166 }
167 }
168 Ok(())
169}
170
171fn load_float(
174 e: &Engine,
175 m: &StModel,
176 name: &str,
177 expected_ne: &[u64],
178) -> Result<GpuTensor, Box<dyn std::error::Error>> {
179 let (info, bytes) = m
180 .raw(name)
181 .ok_or_else(|| format!("EAGLE3 draft missing tensor {name}"))?;
182 let ne = validate_eagle_tensor(name, info, expected_ne)?;
183 let n = ne.iter().try_fold(1u64, |total, extent| {
184 total
185 .checked_mul(*extent)
186 .ok_or_else(|| format!("EAGLE3 tensor {name} element count overflow"))
187 })?;
188 let f32v = dequant::dequantize(info.ggml_type()?, bytes, n as usize);
189 Ok(GpuTensor::Float {
190 data: e.htod(&f32v)?,
191 ne,
192 })
193}
194
195impl Eagle3Draft {
196 pub fn load(e: &Engine, path: &Path) -> Result<Self, Box<dyn std::error::Error>> {
200 let dir = if path.is_file() {
201 path.parent().unwrap_or(Path::new("."))
202 } else {
203 path
204 };
205 let cfg = EagleConfig::from_json(&dir.join("config.json"))?;
206 validate_aux_layers(&cfg.aux_layers)?;
207 validate_eagle_attention_geometry(cfg.n_head, cfg.n_head_kv, cfg.head_dim)?;
208 let m = StModel::open(path)?;
209
210 let d2t = read_i64(&m, "d2t")?;
211 validate_d2t_map(&d2t, cfg.draft_vocab, None)?;
212
213 let n = cfg.hidden_size as u64;
214 let two_n = n.checked_mul(2).ok_or("EAGLE3 hidden geometry overflow")?;
215 let three_n = n.checked_mul(3).ok_or("EAGLE3 hidden geometry overflow")?;
216 let ff = cfg.intermediate_size as u64;
217 let q = cfg
218 .n_head
219 .checked_mul(cfg.head_dim)
220 .ok_or("EAGLE3 q projection geometry overflow")? as u64;
221 let kv = cfg
222 .n_head_kv
223 .checked_mul(cfg.head_dim)
224 .ok_or("EAGLE3 kv projection geometry overflow")? as u64;
225 let vocab = cfg.draft_vocab as u64;
226
227 let draft = Eagle3Draft {
228 fc: load_float(e, &m, "fc.weight", &[three_n, n])?,
229 input_layernorm: load_float(e, &m, "midlayer.input_layernorm.weight", &[n])?,
230 hidden_norm: load_float(e, &m, "midlayer.hidden_norm.weight", &[n])?,
231 q_proj: load_float(e, &m, "midlayer.self_attn.q_proj.weight", &[two_n, q])?,
232 k_proj: load_float(e, &m, "midlayer.self_attn.k_proj.weight", &[two_n, kv])?,
233 v_proj: load_float(e, &m, "midlayer.self_attn.v_proj.weight", &[two_n, kv])?,
234 o_proj: load_float(e, &m, "midlayer.self_attn.o_proj.weight", &[q, n])?,
235 post_attention_layernorm: load_float(
236 e,
237 &m,
238 "midlayer.post_attention_layernorm.weight",
239 &[n],
240 )?,
241 gate_proj: load_float(e, &m, "midlayer.mlp.gate_proj.weight", &[n, ff])?,
242 up_proj: load_float(e, &m, "midlayer.mlp.up_proj.weight", &[n, ff])?,
243 down_proj: load_float(e, &m, "midlayer.mlp.down_proj.weight", &[ff, n])?,
244 norm: load_float(e, &m, "norm.weight", &[n])?,
245 lm_head: load_float(e, &m, "lm_head.weight", &[n, vocab])?,
246 d2t,
247 n_embd: cfg.hidden_size,
248 n_head: cfg.n_head,
249 n_head_kv: cfg.n_head_kv,
250 head_dim: cfg.head_dim,
251 n_ff: cfg.intermediate_size,
252 draft_vocab: cfg.draft_vocab,
253 rope_dim_count: cfg.rope_dim_count(),
254 rope_theta: cfg.rope_theta,
255 eps: cfg.rms_eps,
256 aux_layers: cfg.aux_layers,
257 };
258 Ok(draft)
259 }
260
261 #[inline]
263 pub fn d2t_map(&self, draft_id: u32) -> u32 {
264 (draft_id as i64 + self.d2t[draft_id as usize]) as u32
265 }
266
267 pub fn encode(
271 &self,
272 e: &Engine,
273 aux: &[CudaSlice<f32>],
274 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
275 if aux.len() != self.aux_layers.len() {
276 return Err(format!(
277 "EAGLE3 encode received {} auxiliary states, expected {}",
278 aux.len(),
279 self.aux_layers.len()
280 )
281 .into());
282 }
283 let n = self.n_embd;
284 let mut cat = e.zeros(self.aux_layers.len() * n)?;
285 for (i, a) in aux.iter().enumerate() {
286 e.copy_into(&mut cat, i * n, a, n)?;
287 }
288 e.matmul(&self.fc, &cat, 1) }
290
291 pub fn draft_token(
296 &self,
297 e: &Engine,
298 target: &HybridModel,
299 prev_tok: u32,
300 g: &CudaSlice<f32>,
301 scratch: &mut Eagle3Scratch,
302 pos: usize,
303 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
304 let n = self.n_embd;
305 let eps = self.eps;
306 let pos_d = e.htod_i32(&[pos as i32])?;
307
308 let e_emb = e.htod(&target.embd.gather(n, &[prev_tok]))?;
311 let mut e_norm = e.zeros(n)?;
312 e.rms_norm(
313 &e_emb,
314 self.input_layernorm.float_data(),
315 &mut e_norm,
316 n,
317 1,
318 eps,
319 )?;
320 let res = e.clone_dtod(g)?;
321 let mut g_norm = e.zeros(n)?;
322 e.rms_norm(g, self.hidden_norm.float_data(), &mut g_norm, n, 1, eps)?;
323 let mut cat = e.zeros(2 * n)?;
325 e.copy_into(&mut cat, 0, &e_norm, n)?;
326 e.copy_into(&mut cat, n, &g_norm, n)?;
327
328 let attn = self.attn(e, &cat, &pos_d, scratch)?;
330 let mut x1 = e.zeros(n)?;
332 e.add(&attn, &res, &mut x1, n)?;
333 let mut z = e.zeros(n)?;
335 e.rms_norm(
336 &x1,
337 self.post_attention_layernorm.float_data(),
338 &mut z,
339 n,
340 1,
341 eps,
342 )?;
343 let gate = e.matmul(&self.gate_proj, &z, 1)?;
345 let up = e.matmul(&self.up_proj, &z, 1)?;
346 let mut act = e.zeros(self.n_ff)?;
347 e.silu_mul(&gate, &up, &mut act, self.n_ff)?;
348 let mlp = e.matmul(&self.down_proj, &act, 1)?;
349 let mut g_next = e.zeros(n)?;
351 e.add(&mlp, &x1, &mut g_next, n)?;
352 let mut hn = e.zeros(n)?;
354 e.rms_norm(&g_next, self.norm.float_data(), &mut hn, n, 1, eps)?;
355 let logits = e.matmul(&self.lm_head, &hn, 1)?;
356 let host = e.dtoh(&logits)?;
357 Ok((host, g_next))
358 }
359
360 fn attn(
364 &self,
365 e: &Engine,
366 cat: &CudaSlice<f32>,
367 pos_d: &CudaSlice<i32>,
368 scratch: &mut Eagle3Scratch,
369 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
370 let (nh, nhkv, hd) = (self.n_head, self.n_head_kv, self.head_dim);
371 let scale = 1.0 / (hd as f32).sqrt();
372 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(
378 &mut q,
379 pos_d,
380 hd,
381 self.rope_dim_count,
382 nh,
383 1,
384 self.rope_theta,
385 1.0,
386 )?;
387 e.rope_neox(
388 &mut k,
389 pos_d,
390 hd,
391 self.rope_dim_count,
392 nhkv,
393 1,
394 self.rope_theta,
395 1.0,
396 )?;
397
398 let kv = &mut scratch.kv;
399 e.append_kv_quantized(
400 &k,
401 &v,
402 &mut kv.k,
403 &mut kv.v,
404 kv.len,
405 kv.kv_dim_k,
406 kv.kv_dim_v,
407 kv.k_tok_bytes,
408 kv.v_tok_bytes,
409 false,
410 )?;
411 kv.len += 1;
412 let t_kv = kv.len;
413 let (ktb, vtb) = (kv.k_tok_bytes, kv.v_tok_bytes);
414 let k_view = e.view_u8(&kv.k, t_kv * ktb);
415 let v_view = e.view_u8(&kv.v, t_kv * vtb);
416 let mut attn = e.zeros(nh * hd)?;
417 e.fa_decode(
418 &q, &k_view, &v_view, &mut attn, hd, nh, nhkv, t_kv, scale, ktb, vtb,
419 )?;
420 e.matmul(&self.o_proj, &attn, 1)
421 }
422}
423
424pub struct Eagle3Scratch {
427 pub kv: KvLayer,
428}
429impl Eagle3Scratch {
430 pub fn new(
431 e: &Engine,
432 draft: &Eagle3Draft,
433 cap: usize,
434 ) -> Result<Self, Box<dyn std::error::Error>> {
435 let (nhkv, hd) = (draft.n_head_kv, draft.head_dim);
436 assert!(
437 hd % 32 == 0,
438 "KVQUANT requires head_dim%32==0 (EAGLE3 scratch)"
439 );
440 let kv_dim_k = hd * nhkv;
441 let kv_dim_v = hd * nhkv;
442 let (kbb, vbb) = crate::kv_blk_bytes(); let k_tok_bytes = (kv_dim_k / 32) * kbb;
444 let v_tok_bytes = (kv_dim_v / 32) * vbb;
445 Ok(Eagle3Scratch {
446 kv: KvLayer {
447 k: e.alloc_u8(cap * k_tok_bytes)?,
448 v: e.alloc_u8(cap * v_tok_bytes)?,
449 kv_dim_k,
450 kv_dim_v,
451 k_tok_bytes,
452 v_tok_bytes,
453 len: 0,
454 ring: None,
455 len_d: e.htod_i32(&[0])?,
456 base_d: None,
457 },
458 })
459 }
460 pub fn reset(&mut self) {
461 self.kv.len = 0;
462 }
463}
464
465impl HybridModel {
466 pub fn generate_spec_eagle(
471 &self,
472 e: &Engine,
473 draft: &Eagle3Draft,
474 prompt: &[u32],
475 max_new: usize,
476 k: usize,
477 ) -> Result<(Vec<u32>, usize, usize), Box<dyn std::error::Error>> {
478 self.refuse_hyper("generate_spec_eagle")?;
479 assert!(k >= 1, "k must be >= 1");
480 assert!(!prompt.is_empty(), "prompt must be non-empty");
481 let n_vocab = self.output.out_features();
482 validate_d2t_map(&draft.d2t, draft.draft_vocab, Some(n_vocab))?;
483 let n_embd = self.cfg.n_embd as usize;
484 assert_eq!(n_embd, draft.n_embd, "draft n_embd != target n_embd");
485 let aux = &draft.aux_layers;
486 let max_ctx = prompt.len() + max_new + k + 8;
487 let mut cache = Cache::new(e, &self.cfg, max_ctx)?;
488
489 let mut prime_logits = Vec::new();
491 let mut prime_aux: Vec<CudaSlice<f32>> = Vec::new();
492 for &tok in prompt {
493 let (l, a) = self.decode_step_aux(e, tok, &mut cache, aux)?;
494 prime_logits = l;
495 prime_aux = a;
496 }
497
498 let mut scratch = Eagle3Scratch::new(e, draft, k + 1)?;
499 let mut out: Vec<u32> = Vec::with_capacity(max_new);
500 let mut total_drafted = 0usize;
501 let mut total_accepted = 0usize;
502
503 let shift = std::env::var("MEMRA_EAGLE_ALIGN")
512 .ok()
513 .map(|s| s != "0")
514 .unwrap_or(true);
515 let mut last_token = argmax(&prime_logits) as u32;
516 out.push(last_token);
517 let mut prev_aux = prime_aux;
520 let (mut last_logits, mut g_aux) = self.decode_step_aux(e, last_token, &mut cache, aux)?;
521
522 while out.len() < max_new {
523 let pos = cache.pos;
524 let snap = cache.snapshot(e)?;
525
526 let seed_aux = if shift { &prev_aux } else { &g_aux };
529 let g0 = draft.encode(e, seed_aux)?;
530
531 scratch.reset();
533 let mut draft_toks: Vec<u32> = Vec::with_capacity(k);
534 let mut prev = last_token;
535 let mut g = g0;
536 for j in 0..k {
537 let (dl, g_next) = draft.draft_token(e, self, prev, &g, &mut scratch, pos + j)?;
538 let d_draft = argmax(&dl) as u32;
539 let d_target = draft.d2t_map(d_draft); draft_toks.push(d_target);
541 prev = d_target;
542 g = g_next;
543 }
544
545 let tlogits = self.decode_step_t(e, &draft_toks, pos, &mut cache)?;
547
548 let t_pred = |j: usize| -> u32 {
550 if j == 0 {
551 argmax(&last_logits) as u32
552 } else {
553 argmax(&tlogits[(j - 1) * n_vocab..j * n_vocab]) as u32
554 }
555 };
556 let mut n_acc = 0usize;
557 #[allow(clippy::needless_range_loop)]
558 for j in 0..k {
560 if t_pred(j) == draft_toks[j] {
561 n_acc += 1;
562 } else {
563 break;
564 }
565 }
566 let bonus = t_pred(n_acc);
567 total_drafted += k;
568 total_accepted += n_acc;
569
570 #[allow(clippy::needless_range_loop)]
572 for j in 0..n_acc {
574 if out.len() >= max_new {
575 break;
576 }
577 out.push(draft_toks[j]);
578 }
579 let bonus_emitted = out.len() < max_new;
580 if bonus_emitted {
581 out.push(bonus);
582 }
583 last_token = bonus;
584
585 let pred_is_prev_round = n_acc == 0; let old_g_aux = std::mem::take(&mut g_aux); cache.rollback(e, &snap, 0)?;
599 let mut replay: Vec<u32> = draft_toks[0..n_acc].to_vec();
600 replay.push(bonus);
601 let pred_col = if pred_is_prev_round {
602 None
603 } else {
604 Some(replay.len() - 2)
605 };
606 let (rl, mut a_last, a_pred) =
607 self.decode_step_t_aux2(e, &replay, pos, &mut cache, aux, pred_col)?;
608 last_logits = rl[(replay.len() - 1) * n_vocab..replay.len() * n_vocab].to_vec();
609 prev_aux = if pred_is_prev_round {
610 old_g_aux
611 } else {
612 a_pred.unwrap()
613 };
614 g_aux = std::mem::take(&mut a_last);
615 }
616 out.truncate(max_new);
617 Ok((out, total_drafted, total_accepted))
618 }
619}
620
621struct EagleConfig {
624 hidden_size: usize,
625 n_head: usize,
626 n_head_kv: usize,
627 head_dim: usize,
628 intermediate_size: usize,
629 draft_vocab: usize,
630 rotary_dim: Option<u32>,
634 partial_rotary_factor: Option<f32>,
640 rope_theta: f32,
641 rms_eps: f32,
642 aux_layers: Vec<usize>,
643}
644
645impl EagleConfig {
646 fn rope_dim_count(&self) -> usize {
654 memra_gguf::config::resolve_rope_dim_count(
655 self.rotary_dim,
656 self.partial_rotary_factor,
657 self.head_dim as u32,
658 ) as usize
659 }
660
661 fn from_json(path: &Path) -> Result<Self, Box<dyn std::error::Error>> {
662 Self::from_json_str(&std::fs::read_to_string(path)?)
663 }
664
665 fn from_json_str(txt: &str) -> Result<Self, Box<dyn std::error::Error>> {
666 let num = |key: &str| -> Option<f64> {
668 let pat = format!("\"{key}\"");
669 let i = txt.find(&pat)? + pat.len();
670 let rest = &txt[i..];
671 let c = rest.find(':')? + 1;
672 let tail = rest[c..].trim_start();
673 let end = tail.find([',', '}', '\n']).unwrap_or(tail.len());
674 tail[..end].trim().parse::<f64>().ok()
675 };
676 let aux_layers: Vec<usize> = {
677 let pat = "\"eagle_aux_hidden_state_layer_ids\"";
679 match txt.find(pat) {
680 Some(i) => {
681 let rest = &txt[i + pat.len()..];
682 let lb = rest.find('[').ok_or("no [ after aux ids")?;
683 let rb = rest.find(']').ok_or("no ] after aux ids")?;
684 rest[lb + 1..rb]
685 .split(',')
686 .filter_map(|s| s.trim().parse::<usize>().ok())
687 .collect()
688 }
689 None => vec![1, 15, 28], }
691 };
692 Ok(EagleConfig {
693 hidden_size: num("hidden_size").ok_or("hidden_size")? as usize,
694 n_head: num("num_attention_heads").ok_or("num_attention_heads")? as usize,
695 n_head_kv: num("num_key_value_heads").ok_or("num_key_value_heads")? as usize,
696 head_dim: num("head_dim").ok_or("head_dim")? as usize,
697 intermediate_size: num("intermediate_size").ok_or("intermediate_size")? as usize,
698 draft_vocab: num("draft_vocab_size").ok_or("draft_vocab_size")? as usize,
699 rotary_dim: num("rotary_dim").map(|v| v as u32),
700 partial_rotary_factor: num("partial_rotary_factor").map(|v| v as f32),
701 rope_theta: num("rope_theta").unwrap_or(10000.0) as f32,
702 rms_eps: num("rms_norm_eps").unwrap_or(1e-6) as f32,
703 aux_layers,
704 })
705 }
706}
707
708fn read_i64(m: &StModel, name: &str) -> Result<Vec<i64>, Box<dyn std::error::Error>> {
710 let (info, bytes) = m
711 .raw(name)
712 .ok_or_else(|| format!("EAGLE3 draft missing {name}"))?;
713 if info.dtype != "I64" {
714 return Err(format!("{name} dtype must be I64, found {}", info.dtype).into());
715 }
716 let ne = info.ne();
717 if ne.len() != 1 {
718 return Err(format!("{name} must be rank-1, found shape {ne:?}").into());
719 }
720 let n = usize::try_from(ne[0]).map_err(|_| format!("{name} length does not fit usize"))?;
721 let expected = n
722 .checked_mul(8)
723 .ok_or_else(|| format!("{name} byte length overflow"))?;
724 if bytes.len() != expected {
725 return Err(format!("{name} has {} bytes, expected {expected}", bytes.len()).into());
726 }
727 let mut v = Vec::with_capacity(n);
728 for i in 0..n {
729 v.push(i64::from_le_bytes(
730 bytes[i * 8..i * 8 + 8].try_into().unwrap(),
731 ));
732 }
733 Ok(v)
734}
735
736#[cfg(test)]
747mod tensor_contract_tests {
748 use super::{
749 validate_aux_layers, validate_d2t_map, validate_eagle_attention_geometry,
750 validate_eagle_tensor,
751 };
752 use memra_gguf::safetensors::StInfo;
753
754 #[test]
755 fn every_eagle_weight_is_shape_and_dtype_checked_before_cuda() {
756 let valid = StInfo {
757 dtype: "BF16".into(),
758 shape: vec![8, 4],
759 data_offsets: [0, 64],
760 };
761 assert_eq!(validate_eagle_tensor("q", &valid, &[4, 8]).unwrap(), [4, 8]);
762 let mut bad = valid.clone();
763 bad.dtype = "I8".into();
764 assert!(
765 validate_eagle_tensor("q", &bad, &[4, 8])
766 .unwrap_err()
767 .contains("dtype")
768 );
769 bad = valid.clone();
770 bad.shape = vec![32];
771 assert!(
772 validate_eagle_tensor("q", &bad, &[4, 8])
773 .unwrap_err()
774 .contains("shape")
775 );
776 }
777
778 #[test]
779 fn auxiliary_and_vocabulary_maps_are_closed_before_cuda() {
780 assert!(validate_aux_layers(&[1, 15, 28]).is_ok());
781 for invalid in [&[1, 2][..], &[1, 1, 2], &[2, 1, 3], &[1, 2, 3, 4]] {
782 assert!(
783 validate_aux_layers(invalid).is_err(),
784 "accepted {invalid:?}"
785 );
786 }
787 assert!(validate_d2t_map(&[0, 1, -1], 3, Some(4)).is_ok());
788 assert!(validate_d2t_map(&[0], 2, None).is_err());
789 assert!(validate_d2t_map(&[-1], 1, None).is_err());
790 assert!(validate_d2t_map(&[i64::MAX], 1, None).is_err());
791 let err = validate_d2t_map(&[4], 1, Some(4)).unwrap_err();
792 assert!(err.contains("target vocabulary of 4 entries"), "{err}");
793 assert!(validate_eagle_attention_geometry(32, 8, 128).is_ok());
794 assert!(validate_eagle_attention_geometry(7, 8, 128).is_err());
795 assert!(validate_eagle_attention_geometry(8, 0, 128).is_err());
796 assert!(validate_eagle_attention_geometry(8, 8, 33).is_err());
797 assert!(validate_eagle_attention_geometry(usize::MAX, 1, 2).is_err());
798 }
799}
800
801#[cfg(test)]
802mod draft_rope_width_tests {
803 use super::EagleConfig;
804 use memra_gguf::config::{HfConfig, resolve_rope_dim_count};
805
806 const EAGLE3_QWEN35_9B_CONFIG: &str = r#"{
811 "architectures": [
812 "LlamaForCausalLMEagle3"
813 ],
814 "attention_bias": false,
815 "attention_dropout": 0.0,
816 "bos_token_id": 248040,
817 "draft_vocab_size": 32000,
818 "dtype": "bfloat16",
819 "eos_token_id": 248044,
820 "head_dim": 256,
821 "hidden_act": "silu",
822 "hidden_size": 4096,
823 "initializer_range": 0.02,
824 "intermediate_size": 12288,
825 "max_position_embeddings": 262144,
826 "mlp_bias": false,
827 "model_type": "llama",
828 "num_attention_heads": 16,
829 "num_hidden_layers": 1,
830 "num_key_value_heads": 4,
831 "pad_token_id": null,
832 "partial_rotary_factor": 0.25,
833 "pretraining_tp": 1,
834 "rms_norm_eps": 1e-06,
835 "rope_parameters": {
836 "partial_rotary_factor": 0.25,
837 "rope_theta": 10000000,
838 "rope_type": "default"
839 },
840 "tie_word_embeddings": false,
841 "transformers_version": "5.3.0",
842 "use_cache": true,
843 "vocab_size": 248320,
844 "eagle_config": {
845 "use_aux_hidden_state": true,
846 "eagle_aux_hidden_state_layer_ids": [1, 15, 28]
847 }
848}"#;
849
850 fn edited(from: &str, to: &str) -> String {
853 assert!(
854 EAGLE3_QWEN35_9B_CONFIG.contains(from),
855 "fixture drifted: {from:?} not found — the variant below would test the wrong shape"
856 );
857 EAGLE3_QWEN35_9B_CONFIG.replace(from, to)
858 }
859
860 fn assert_reader_parity(json: &str) -> usize {
865 let draft = EagleConfig::from_json_str(json).expect("draft reader must parse the fixture");
866 let hf = HfConfig::parse(json);
867 assert_eq!(
868 draft.rotary_dim, hf.rotary_dim,
869 "draft scanner and HfConfig::parse disagree on rotary_dim for the same config"
870 );
871 assert_eq!(
872 draft.partial_rotary_factor, hf.partial_rotary_factor,
873 "draft scanner and HfConfig::parse disagree on partial_rotary_factor for the same config"
874 );
875 let expected = resolve_rope_dim_count(
876 hf.rotary_dim,
877 hf.partial_rotary_factor,
878 hf.head_dim.expect("fixture declares head_dim"),
879 ) as usize;
880 assert_eq!(
881 draft.rope_dim_count(),
882 expected,
883 "draft rope width diverged from the shared derivation on the same facts"
884 );
885 draft.rope_dim_count()
886 }
887
888 #[test]
892 fn real_eagle3_qwen35_9b_config_derives_partial_rope_64_of_256() {
893 let cfg = EagleConfig::from_json_str(EAGLE3_QWEN35_9B_CONFIG).expect("real config parses");
894 assert_eq!(cfg.head_dim, 256);
895 assert_eq!(
896 cfg.rotary_dim, None,
897 "no published EAGLE3 draft declares rotary_dim"
898 );
899 assert_eq!(
900 cfg.partial_rotary_factor,
901 Some(0.25),
902 "the declared factor must be READ, not defaulted — unwrap_or(1.0) is the bug class"
903 );
904 assert_eq!(
905 cfg.rope_dim_count(),
906 64,
907 "eagle3-qwen35-9b rotates 64 of 256 head dims; full rope silently corrupts the \
908 pass-through band 64..256 — no shape error, fluent output, wrecked long context"
909 );
910 assert_eq!(assert_reader_parity(EAGLE3_QWEN35_9B_CONFIG), 64);
911 }
912
913 #[test]
916 fn nested_only_partial_rotary_spelling_is_still_partial_rope() {
917 let json = edited("\n \"partial_rotary_factor\": 0.25,", "");
918 let cfg = EagleConfig::from_json_str(&json).expect("nested-only config parses");
919 assert_eq!(
920 cfg.partial_rotary_factor,
921 Some(0.25),
922 "rope_parameters spelling must be read"
923 );
924 assert_eq!(cfg.rope_dim_count(), 64);
925 assert_reader_parity(&json);
926 }
927
928 #[test]
932 fn no_rope_declaration_is_full_rope() {
933 let json = edited("\n \"partial_rotary_factor\": 0.25,", "")
934 .replace("\n \"partial_rotary_factor\": 0.25,", "");
935 assert!(
936 !json.contains("partial_rotary_factor"),
937 "variant edit failed: a factor spelling survived"
938 );
939 let cfg = EagleConfig::from_json_str(&json).expect("undeclared-rope config parses");
940 assert_eq!(cfg.partial_rotary_factor, None);
941 assert_eq!(
942 cfg.rope_dim_count(),
943 256,
944 "absent declaration = every head dim rotates"
945 );
946 assert_reader_parity(&json);
947 }
948
949 #[test]
954 fn explicit_rotary_dim_wins_over_the_fraction() {
955 let json = edited(
956 "\n \"partial_rotary_factor\": 0.25,",
957 "\n \"partial_rotary_factor\": 0.25,\n \"rotary_dim\": 32,",
958 );
959 let cfg = EagleConfig::from_json_str(&json).expect("explicit-dims config parses");
960 assert_eq!(cfg.rotary_dim, Some(32));
961 assert_eq!(
962 cfg.rope_dim_count(),
963 32,
964 "explicit rotary_dim is the more specific declaration and must win over the fraction"
965 );
966 assert_reader_parity(&json);
967 }
968
969 #[test]
974 fn malformed_factor_takes_full_width_not_a_wider_than_head_rotation() {
975 let over = EAGLE3_QWEN35_9B_CONFIG.replace(
976 "\"partial_rotary_factor\": 0.25",
977 "\"partial_rotary_factor\": 2.0",
978 );
979 let cfg = EagleConfig::from_json_str(&over).expect("factor-2.0 config parses");
980 assert_eq!(cfg.partial_rotary_factor, Some(2.0));
981 assert_eq!(
982 cfg.rope_dim_count(),
983 256,
984 "factor 2.0 must take the FULL head width (256), never 512 — the old \
985 factor*head_dim arithmetic rotated past the head allocation"
986 );
987 assert_reader_parity(&over);
988
989 let zero = EAGLE3_QWEN35_9B_CONFIG.replace(
990 "\"partial_rotary_factor\": 0.25",
991 "\"partial_rotary_factor\": 0.0",
992 );
993 let cfg = EagleConfig::from_json_str(&zero).expect("factor-0.0 config parses");
994 assert_eq!(
995 cfg.rope_dim_count(),
996 256,
997 "factor 0.0 is malformed and takes the full width, not the old max(2) stub rotation"
998 );
999 assert_reader_parity(&zero);
1000 }
1001}