use anyhow::Result;
#[derive(Debug, Clone)]
pub struct LayerActivations {
pub layer_inputs: Vec<Vec<f32>>,
pub layer_outputs: Vec<Vec<f32>>,
pub num_layers: u32,
pub seq_len: u32,
pub hidden_size: u32,
pub target_layer_filter: Option<Vec<usize>>,
}
impl LayerActivations {
pub fn element_count(&self) -> usize {
let per_layer = (self.seq_len as usize) * (self.hidden_size as usize) * 2;
match self.target_layer_filter.as_ref() {
Some(filter) => {
filter
.iter()
.filter(|&&i| i < (self.num_layers as usize))
.count()
* per_layer
}
None => (self.num_layers as usize) * per_layer,
}
}
pub fn is_target_layer(&self, layer_idx: usize) -> bool {
self.target_layer_filter
.as_ref()
.map_or(true, |f| f.contains(&layer_idx))
}
pub fn validate(&self) -> Result<()> {
if self.layer_inputs.len() != self.num_layers as usize {
anyhow::bail!(
"LayerActivations: layer_inputs.len() = {} != num_layers = {}",
self.layer_inputs.len(),
self.num_layers
);
}
if self.layer_outputs.len() != self.num_layers as usize {
anyhow::bail!(
"LayerActivations: layer_outputs.len() = {} != num_layers = {}",
self.layer_outputs.len(),
self.num_layers
);
}
let expected_per_layer = (self.seq_len as usize) * (self.hidden_size as usize);
if let Some(filter) = self.target_layer_filter.as_ref() {
for &f in filter {
if f >= (self.num_layers as usize) {
anyhow::bail!(
"LayerActivations: target_layer_filter contains out-of-range index {} (num_layers={})",
f,
self.num_layers,
);
}
}
}
let is_target = |i: usize| -> bool { self.is_target_layer(i) };
for (i, v) in self.layer_inputs.iter().enumerate() {
if !is_target(i) {
if !v.is_empty() {
anyhow::bail!(
"LayerActivations: layer_inputs[{}] non-target index expected empty, got len={}",
i,
v.len(),
);
}
continue;
}
if v.len() != expected_per_layer {
anyhow::bail!(
"LayerActivations: layer_inputs[{}].len() = {} != seq_len({}) * hidden_size({}) = {}",
i,
v.len(),
self.seq_len,
self.hidden_size,
expected_per_layer
);
}
}
for (i, v) in self.layer_outputs.iter().enumerate() {
if !is_target(i) {
if !v.is_empty() {
anyhow::bail!(
"LayerActivations: layer_outputs[{}] non-target index expected empty, got len={}",
i,
v.len(),
);
}
continue;
}
if v.len() != expected_per_layer {
anyhow::bail!(
"LayerActivations: layer_outputs[{}].len() = {} != {}",
i,
v.len(),
expected_per_layer
);
}
}
Ok(())
}
}
pub trait ActivationCapture {
fn run_calibration_prompt(&mut self, tokens: &[u32]) -> Result<LayerActivations>;
}
pub struct MockActivationCapture {
pub num_layers: u32,
pub hidden_size: u32,
}
impl MockActivationCapture {
pub fn new(num_layers: u32, hidden_size: u32) -> Self {
Self {
num_layers,
hidden_size,
}
}
}
impl ActivationCapture for MockActivationCapture {
fn run_calibration_prompt(&mut self, tokens: &[u32]) -> Result<LayerActivations> {
if tokens.is_empty() {
anyhow::bail!("MockActivationCapture: tokens must be non-empty");
}
let seq = tokens.len();
let h = self.hidden_size as usize;
let mut layer_inputs = Vec::with_capacity(self.num_layers as usize);
let mut layer_outputs = Vec::with_capacity(self.num_layers as usize);
for l in 0..self.num_layers as usize {
let mut inp = vec![0.0f32; seq * h];
let mut out = vec![0.0f32; seq * h];
for t in 0..seq {
for j in 0..h {
let base = (tokens[t] as f32) * 0.001 + (l as f32) * 0.01 + (j as f32) * 0.0001;
inp[t * h + j] = base;
out[t * h + j] = base + 1.0;
}
}
layer_inputs.push(inp);
layer_outputs.push(out);
}
Ok(LayerActivations {
layer_inputs,
layer_outputs,
num_layers: self.num_layers,
seq_len: seq as u32,
hidden_size: self.hidden_size,
target_layer_filter: None,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn activations_validate_accepts_matching_shapes() {
let act = LayerActivations {
layer_inputs: vec![vec![0.0; 8]; 3],
layer_outputs: vec![vec![0.0; 8]; 3],
num_layers: 3,
seq_len: 2,
hidden_size: 4,
target_layer_filter: None,
};
act.validate().expect("should validate");
}
#[test]
fn activations_validate_rejects_wrong_layer_count() {
let act = LayerActivations {
layer_inputs: vec![vec![0.0; 8]; 2], layer_outputs: vec![vec![0.0; 8]; 3],
num_layers: 3,
seq_len: 2,
hidden_size: 4,
target_layer_filter: None,
};
assert!(act.validate().is_err());
}
#[test]
fn activations_validate_rejects_wrong_element_count() {
let act = LayerActivations {
layer_inputs: vec![vec![0.0; 9]; 3], layer_outputs: vec![vec![0.0; 8]; 3],
num_layers: 3,
seq_len: 2,
hidden_size: 4,
target_layer_filter: None,
};
assert!(act.validate().is_err());
}
#[test]
fn activations_element_count_matches_expected() {
let act = LayerActivations {
layer_inputs: vec![vec![0.0; 8]; 3],
layer_outputs: vec![vec![0.0; 8]; 3],
num_layers: 3,
seq_len: 2,
hidden_size: 4,
target_layer_filter: None,
};
assert_eq!(act.element_count(), 3 * 2 * 4 * 2);
}
#[test]
fn activations_validate_accepts_filtered_layers_2026_05_21() {
let act = LayerActivations {
layer_inputs: vec![
vec![0.0; 8],
vec![], vec![0.0; 8],
],
layer_outputs: vec![vec![0.0; 8], vec![], vec![0.0; 8]],
num_layers: 3,
seq_len: 2,
hidden_size: 4,
target_layer_filter: Some(vec![0, 2]),
};
act.validate().expect("filtered validate should pass");
}
#[test]
fn activations_validate_rejects_filtered_target_short_2026_05_21() {
let act = LayerActivations {
layer_inputs: vec![vec![0.0; 4] , vec![], vec![0.0; 8]],
layer_outputs: vec![vec![0.0; 8], vec![], vec![0.0; 8]],
num_layers: 3,
seq_len: 2,
hidden_size: 4,
target_layer_filter: Some(vec![0, 2]),
};
assert!(act.validate().is_err());
}
#[test]
fn activations_validate_rejects_non_target_nonempty_2026_05_21() {
let act = LayerActivations {
layer_inputs: vec![
vec![0.0; 8],
vec![0.0; 8],
vec![0.0; 8],
],
layer_outputs: vec![vec![0.0; 8], vec![], vec![0.0; 8]],
num_layers: 3,
seq_len: 2,
hidden_size: 4,
target_layer_filter: Some(vec![0, 2]),
};
assert!(act.validate().is_err());
}
#[test]
fn mock_produces_correct_shapes() {
let mut mock = MockActivationCapture::new(5, 16);
let tokens = vec![10u32, 20, 30];
let act = mock.run_calibration_prompt(&tokens).expect("run");
act.validate().expect("shapes valid");
assert_eq!(act.num_layers, 5);
assert_eq!(act.seq_len, 3);
assert_eq!(act.hidden_size, 16);
assert_eq!(act.layer_inputs.len(), 5);
assert_eq!(act.layer_outputs.len(), 5);
for v in &act.layer_inputs {
assert_eq!(v.len(), 3 * 16);
}
for v in &act.layer_outputs {
assert_eq!(v.len(), 3 * 16);
}
}
#[test]
fn mock_is_deterministic() {
let mut mock1 = MockActivationCapture::new(2, 4);
let mut mock2 = MockActivationCapture::new(2, 4);
let tokens = vec![1u32, 2, 3];
let a1 = mock1.run_calibration_prompt(&tokens).unwrap();
let a2 = mock2.run_calibration_prompt(&tokens).unwrap();
for l in 0..2 {
for i in 0..12 {
assert_eq!(
a1.layer_inputs[l][i].to_bits(),
a2.layer_inputs[l][i].to_bits()
);
assert_eq!(
a1.layer_outputs[l][i].to_bits(),
a2.layer_outputs[l][i].to_bits()
);
}
}
}
#[test]
fn mock_differs_between_layers() {
let mut mock = MockActivationCapture::new(3, 4);
let tokens = vec![5u32];
let act = mock.run_calibration_prompt(&tokens).unwrap();
let l0 = &act.layer_inputs[0];
let l1 = &act.layer_inputs[1];
let l2 = &act.layer_inputs[2];
let mut any_differ_01 = false;
let mut any_differ_12 = false;
for i in 0..4 {
if (l0[i] - l1[i]).abs() > 1e-6 {
any_differ_01 = true;
}
if (l1[i] - l2[i]).abs() > 1e-6 {
any_differ_12 = true;
}
}
assert!(any_differ_01, "layer 0 == layer 1 (mock broken)");
assert!(any_differ_12, "layer 1 == layer 2 (mock broken)");
}
#[test]
fn mock_differs_between_tokens() {
let mut mock = MockActivationCapture::new(1, 4);
let a1 = mock.run_calibration_prompt(&[1u32, 2, 3]).unwrap();
let a2 = mock.run_calibration_prompt(&[100u32, 200, 300]).unwrap();
let mut any_differ = false;
for i in 0..12 {
if (a1.layer_inputs[0][i] - a2.layer_inputs[0][i]).abs() > 1e-6 {
any_differ = true;
break;
}
}
assert!(
any_differ,
"mock token content is not encoded in activations"
);
}
#[test]
fn mock_rejects_empty_tokens() {
let mut mock = MockActivationCapture::new(2, 4);
let result = mock.run_calibration_prompt(&[]);
assert!(result.is_err(), "empty tokens should error");
}
#[test]
fn trait_object_usable_for_cross_adr_consumption() {
fn accumulate_mean_input_per_layer(
capture: &mut dyn ActivationCapture,
tokens: &[u32],
) -> Result<Vec<f32>> {
let act = capture.run_calibration_prompt(tokens)?;
let mut means = Vec::with_capacity(act.num_layers as usize);
for l in 0..act.num_layers as usize {
let sum: f32 = act.layer_inputs[l].iter().sum();
let n = act.layer_inputs[l].len() as f32;
means.push(sum / n);
}
Ok(means)
}
let mut mock = MockActivationCapture::new(4, 8);
let tokens = vec![42u32, 100, 7];
let means = accumulate_mean_input_per_layer(&mut mock, &tokens).expect("accum");
assert_eq!(means.len(), 4);
for i in 1..means.len() {
assert!(
means[i] > means[i - 1],
"layer {} mean should exceed layer {}",
i,
i - 1
);
}
}
}