use anyhow::{anyhow, ensure, Result};
#[derive(Debug, Clone)]
pub struct Eagle3HiddenCollector {
target_layer_ids: Vec<usize>,
seq_len: usize,
hidden_size: usize,
buffer: Vec<f32>,
written_mask: u64,
}
impl Eagle3HiddenCollector {
pub fn new(target_layer_ids: Vec<usize>, seq_len: usize, hidden_size: usize) -> Result<Self> {
ensure!(
!target_layer_ids.is_empty(),
"Eagle3HiddenCollector: target_layer_ids must be non-empty"
);
ensure!(
target_layer_ids.len() <= 64,
"Eagle3HiddenCollector: at most 64 aux layers supported (written_mask is u64); got {}",
target_layer_ids.len()
);
let mut sorted = target_layer_ids.clone();
sorted.sort_unstable();
for w in sorted.windows(2) {
ensure!(
w[0] != w[1],
"Eagle3HiddenCollector: target_layer_ids has duplicate entry {}",
w[0]
);
}
ensure!(seq_len > 0, "Eagle3HiddenCollector: seq_len must be > 0");
ensure!(
hidden_size > 0,
"Eagle3HiddenCollector: hidden_size must be > 0"
);
let total = seq_len
.checked_mul(target_layer_ids.len())
.and_then(|v| v.checked_mul(hidden_size))
.ok_or_else(|| {
anyhow!(
"Eagle3HiddenCollector: seq_len({}) * num_aux({}) * hidden_size({}) overflows usize",
seq_len,
target_layer_ids.len(),
hidden_size
)
})?;
Ok(Self {
target_layer_ids,
seq_len,
hidden_size,
buffer: vec![0.0f32; total],
written_mask: 0,
})
}
#[inline]
pub fn num_aux(&self) -> usize {
self.target_layer_ids.len()
}
#[inline]
pub fn seq_len(&self) -> usize {
self.seq_len
}
#[inline]
pub fn hidden_size(&self) -> usize {
self.hidden_size
}
#[inline]
pub fn fc_input_size(&self) -> usize {
self.hidden_size * self.num_aux()
}
#[inline]
pub fn target_layer_ids(&self) -> &[usize] {
&self.target_layer_ids
}
pub fn capture_index_for(&self, target_layer_idx: usize) -> Option<usize> {
self.target_layer_ids
.iter()
.position(|&i| i == target_layer_idx)
}
pub fn write_layer_slab(&mut self, capture_idx: usize, slab: &[f32]) -> Result<()> {
ensure!(
capture_idx < self.num_aux(),
"Eagle3HiddenCollector::write_layer_slab: capture_idx {} >= num_aux {}",
capture_idx,
self.num_aux()
);
let expected = self.seq_len * self.hidden_size;
ensure!(
slab.len() == expected,
"Eagle3HiddenCollector::write_layer_slab: slab len {} != seq_len({}) * hidden_size({}) = {}",
slab.len(),
self.seq_len,
self.hidden_size,
expected
);
let bit = 1u64 << capture_idx;
ensure!(
(self.written_mask & bit) == 0,
"Eagle3HiddenCollector::write_layer_slab: capture_idx {} already written",
capture_idx
);
let num_aux = self.num_aux();
let hs = self.hidden_size;
for token_pos in 0..self.seq_len {
let src = token_pos * hs;
let dst = (token_pos * num_aux + capture_idx) * hs;
self.buffer[dst..dst + hs].copy_from_slice(&slab[src..src + hs]);
}
self.written_mask |= bit;
Ok(())
}
pub fn is_complete(&self) -> bool {
let full_mask = if self.num_aux() == 64 {
!0u64
} else {
(1u64 << self.num_aux()) - 1
};
self.written_mask == full_mask
}
pub fn concatenated_hidden(&self) -> Result<&[f32]> {
ensure!(
self.is_complete(),
"Eagle3HiddenCollector::concatenated_hidden: incomplete capture — written_mask={:#066b}, expected {} bits set",
self.written_mask,
self.num_aux()
);
Ok(&self.buffer)
}
pub fn reset(&mut self) {
self.written_mask = 0;
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::*;
fn make_slab(seed: u64, seq_len: usize, hidden_size: usize) -> Vec<f32> {
(0..seq_len * hidden_size)
.map(|i| ((seed.wrapping_add(i as u64)) % 1000) as f32 * 0.001)
.collect()
}
#[test]
fn adr_037_e3a_constructor_rejects_empty_layer_ids_2026_05_22() {
let err = Eagle3HiddenCollector::new(vec![], 8, 128).unwrap_err();
assert!(err.to_string().contains("non-empty"), "got: {err}");
}
#[test]
fn adr_037_e3a_constructor_rejects_duplicate_layer_ids_2026_05_22() {
let err = Eagle3HiddenCollector::new(vec![4, 16, 4], 8, 128).unwrap_err();
assert!(err.to_string().contains("duplicate"), "got: {err}");
}
#[test]
fn adr_037_e3a_constructor_rejects_excessive_layer_count_2026_05_22() {
let layer_ids: Vec<usize> = (0..65).collect();
let err = Eagle3HiddenCollector::new(layer_ids, 8, 128).unwrap_err();
assert!(err.to_string().contains("at most 64"), "got: {err}");
}
#[test]
fn adr_037_e3a_constructor_rejects_zero_seq_len_2026_05_22() {
let err = Eagle3HiddenCollector::new(vec![4], 0, 128).unwrap_err();
assert!(err.to_string().contains("seq_len"), "got: {err}");
}
#[test]
fn adr_037_e3a_constructor_rejects_zero_hidden_size_2026_05_22() {
let err = Eagle3HiddenCollector::new(vec![4], 8, 0).unwrap_err();
assert!(err.to_string().contains("hidden_size"), "got: {err}");
}
#[test]
fn adr_037_e3a_fc_input_size_matches_vllm_contract_2026_05_22() {
let c = Eagle3HiddenCollector::new(vec![4, 16, 31], 1, 5120).unwrap();
assert_eq!(c.num_aux(), 3);
assert_eq!(c.hidden_size(), 5120);
assert_eq!(c.fc_input_size(), 15360);
}
#[test]
fn adr_037_e3a_capture_index_for_returns_position_or_none_2026_05_22() {
let c = Eagle3HiddenCollector::new(vec![4, 16, 31], 1, 128).unwrap();
assert_eq!(c.capture_index_for(4), Some(0));
assert_eq!(c.capture_index_for(16), Some(1));
assert_eq!(c.capture_index_for(31), Some(2));
assert_eq!(c.capture_index_for(5), None);
assert_eq!(c.capture_index_for(63), None);
}
#[test]
fn adr_037_e3a_write_layer_slab_rejects_wrong_size_2026_05_22() {
let mut c = Eagle3HiddenCollector::new(vec![4, 16, 31], 8, 128).unwrap();
let bad = vec![0.0f32; 100]; let err = c.write_layer_slab(0, &bad).unwrap_err();
assert!(err.to_string().contains("slab len"), "got: {err}");
}
#[test]
fn adr_037_e3a_write_layer_slab_rejects_out_of_range_2026_05_22() {
let mut c = Eagle3HiddenCollector::new(vec![4, 16, 31], 8, 128).unwrap();
let slab = make_slab(0, 8, 128);
let err = c.write_layer_slab(3, &slab).unwrap_err();
assert!(err.to_string().contains("capture_idx 3"), "got: {err}");
}
#[test]
fn adr_037_e3a_write_layer_slab_rejects_double_write_2026_05_22() {
let mut c = Eagle3HiddenCollector::new(vec![4, 16, 31], 4, 16).unwrap();
let slab = make_slab(0, 4, 16);
c.write_layer_slab(1, &slab).unwrap();
let err = c.write_layer_slab(1, &slab).unwrap_err();
assert!(err.to_string().contains("already written"), "got: {err}");
}
#[test]
fn adr_037_e3a_concatenated_hidden_rejects_incomplete_capture_2026_05_22() {
let mut c = Eagle3HiddenCollector::new(vec![4, 16, 31], 4, 16).unwrap();
let slab = make_slab(0, 4, 16);
c.write_layer_slab(0, &slab).unwrap();
c.write_layer_slab(1, &slab).unwrap();
let err = c.concatenated_hidden().unwrap_err();
assert!(err.to_string().contains("incomplete"), "got: {err}");
}
#[test]
fn adr_037_e3a_concatenated_hidden_layout_matches_vllm_2026_05_22() {
let seq_len = 3;
let hidden_size = 4;
let mut c = Eagle3HiddenCollector::new(vec![10, 20], seq_len, hidden_size).unwrap();
let mut layer0 = Vec::with_capacity(seq_len * hidden_size);
for token in 0..seq_len {
for _ in 0..hidden_size {
layer0.push((token + 1) as f32);
}
}
let mut layer1 = Vec::with_capacity(seq_len * hidden_size);
for token in 0..seq_len {
for _ in 0..hidden_size {
layer1.push(((token + 1) * 10) as f32);
}
}
c.write_layer_slab(0, &layer0).unwrap();
c.write_layer_slab(1, &layer1).unwrap();
let cat = c.concatenated_hidden().unwrap();
assert_eq!(cat.len(), seq_len * 2 * hidden_size);
for token in 0..seq_len {
let base = token * 2 * hidden_size;
for d in 0..hidden_size {
assert_eq!(
cat[base + d],
(token + 1) as f32,
"token {token} layer0 dim {d}: got {}",
cat[base + d]
);
}
for d in 0..hidden_size {
assert_eq!(
cat[base + hidden_size + d],
((token + 1) * 10) as f32,
"token {token} layer1 dim {d}: got {}",
cat[base + hidden_size + d]
);
}
}
}
#[test]
fn adr_037_e3a_layer_id_order_preserved_in_concat_2026_05_22() {
let layer_ids = vec![31, 4, 16]; let c = Eagle3HiddenCollector::new(layer_ids, 4, 16).unwrap();
assert_eq!(c.target_layer_ids(), &[31, 4, 16]);
assert_eq!(c.capture_index_for(31), Some(0));
assert_eq!(c.capture_index_for(4), Some(1));
assert_eq!(c.capture_index_for(16), Some(2));
}
#[test]
fn adr_037_e3a_reset_clears_written_mask_2026_05_22() {
let mut c = Eagle3HiddenCollector::new(vec![4, 16], 4, 16).unwrap();
let slab = make_slab(0, 4, 16);
c.write_layer_slab(0, &slab).unwrap();
c.write_layer_slab(1, &slab).unwrap();
assert!(c.is_complete());
c.reset();
assert!(!c.is_complete());
let err = c.concatenated_hidden().unwrap_err();
assert!(err.to_string().contains("incomplete"), "got: {err}");
c.write_layer_slab(0, &slab).unwrap();
c.write_layer_slab(1, &slab).unwrap();
assert!(c.is_complete());
}
#[test]
fn adr_037_e3a_realistic_qwen35_shape_2026_05_22() {
let layer_ids = vec![8, 16, 32, 48];
let seq_len = 200;
let hidden_size = 5120;
let mut c = Eagle3HiddenCollector::new(layer_ids.clone(), seq_len, hidden_size).unwrap();
assert_eq!(c.num_aux(), 4);
assert_eq!(c.fc_input_size(), 4 * 5120);
assert_eq!(c.buffer.len(), seq_len * 4 * hidden_size);
for (capture_idx, &layer_id) in layer_ids.iter().enumerate() {
let slab = make_slab(layer_id as u64 * 1000, seq_len, hidden_size);
c.write_layer_slab(capture_idx, &slab).unwrap();
}
let cat = c.concatenated_hidden().unwrap();
assert_eq!(cat.len(), seq_len * 4 * hidden_size);
}
}