#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum MetricsMode {
#[default]
Off,
CheapOnline,
AttentionProfile,
HeavyDiagnostics,
}
#[derive(Debug, Clone, Default)]
pub struct LayerMetrics {
pub mode: MetricsMode,
pub layer_idx: usize,
pub latency_ns: u64,
pub input_norm: f32,
pub output_norm: f32,
pub update_ratio: f32,
pub block_influence: f32,
pub entropy: Option<Vec<f32>>,
pub sparsity: Option<Vec<f32>>,
pub kv_page_mass: Option<Vec<f32>>,
pub pattern_label: Option<Vec<String>>,
}
#[derive(Debug)]
pub struct ForwardMetrics {
pub mode: MetricsMode,
pub layers: Vec<LayerMetrics>,
pub total_ns: u64,
}
#[derive(Debug, Clone)]
pub struct OnlineSoftmaxEntropy {
m: f32, l: f32, r: f32, count: usize,
}
impl Default for OnlineSoftmaxEntropy {
fn default() -> Self {
Self::new()
}
}
impl OnlineSoftmaxEntropy {
pub fn new() -> Self {
Self {
m: 0.0,
l: 0.0,
r: 0.0,
count: 0,
}
}
#[inline]
pub fn update(&mut self, logit: f32) {
assert!(logit.is_finite(), "logit must be finite, got {logit}");
self.update_finite(logit);
}
#[inline]
pub fn try_update(&mut self, logit: f32) -> Result<(), crate::error::InferenceError> {
if !logit.is_finite() {
return Err(crate::error::InferenceError::InvalidInput(format!(
"attention entropy logit must be finite, got {logit}"
)));
}
self.update_finite(logit);
Ok(())
}
#[inline]
fn update_finite(&mut self, logit: f32) {
debug_assert!(logit.is_finite());
if self.count == 0 {
self.m = logit;
self.l = 1.0;
self.r = 0.0; self.count = 1;
return;
}
if logit > self.m {
let diff = self.m - logit; let alpha = diff.exp();
self.r = if alpha == 0.0 {
0.0
} else {
alpha * (self.r + self.l * diff)
};
self.l = alpha * self.l + 1.0;
self.m = logit;
} else {
let shifted = logit - self.m; let w = shifted.exp();
self.l += w;
self.r += if w == 0.0 { 0.0 } else { w * shifted };
}
self.count += 1;
}
pub fn entropy_nats(&self) -> f32 {
if self.count < 2 {
return 0.0;
}
(self.l.ln() - self.r / self.l).max(0.0)
}
pub fn entropy_bits(&self) -> f32 {
self.entropy_nats() / std::f32::consts::LN_2
}
pub fn normalized_entropy(&self) -> f32 {
if self.count < 2 {
return 0.0;
}
self.entropy_nats() / (self.count as f32).ln()
}
pub fn count(&self) -> usize {
self.count
}
pub fn reset(&mut self) {
self.m = 0.0;
self.l = 0.0;
self.r = 0.0;
self.count = 0;
}
}
#[inline]
pub fn l2_norm(data: &[f32]) -> f32 {
let sum_sq: f64 = data
.iter()
.map(|&x| {
let x = f64::from(x);
x * x
})
.sum();
sum_sq.sqrt() as f32
}
#[cfg(test)]
mod tests {
use super::*;
const TOL: f32 = 1e-5;
fn assert_close(a: f32, b: f32, msg: &str) {
let diff = (a - b).abs();
let scale = a.abs().max(b.abs()).max(1e-8);
assert!(
diff / scale < TOL,
"{msg}: got {a}, expected {b}, diff={diff}"
);
}
#[test]
fn test_metrics_mode_default() {
assert_eq!(MetricsMode::default(), MetricsMode::Off);
}
#[test]
fn test_metrics_mode_heavy_diagnostics() {
assert_ne!(MetricsMode::HeavyDiagnostics, MetricsMode::AttentionProfile);
}
#[test]
fn test_entropy_uniform() {
for n in [2usize, 8, 128, 512] {
let mut acc = OnlineSoftmaxEntropy::new();
let logit = 0.5_f32; for _ in 0..n {
acc.update(logit);
}
let expected = (n as f32).ln();
assert_close(
acc.entropy_nats(),
expected,
&format!("uniform entropy N={n}"),
);
}
}
#[test]
fn test_entropy_peaked() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(100.0);
for _ in 0..127 {
acc.update(-100.0);
}
assert!(
acc.entropy_nats() < 0.01,
"peaked entropy should be near 0, got {}",
acc.entropy_nats()
);
}
#[test]
fn test_entropy_two_values() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(1.0);
acc.update(1.0);
assert_close(acc.entropy_nats(), std::f32::consts::LN_2, "50/50 entropy");
}
#[test]
fn test_entropy_matches_naive() {
let logits: Vec<f32> = (0..256)
.map(|i| {
let x = (i as u32)
.wrapping_mul(1_664_525)
.wrapping_add(1_013_904_223);
((x >> 8) as f32) / 16_777_216.0 * 10.0 - 5.0
})
.collect();
let mut acc = OnlineSoftmaxEntropy::new();
for &l in &logits {
acc.update(l);
}
let online_h = acc.entropy_nats();
let max_l = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = logits.iter().map(|&l| (l - max_l).exp()).collect();
let sum_exp: f32 = exps.iter().sum();
let naive_h: f32 = exps
.iter()
.map(|&e| {
let p = e / sum_exp;
if p > 0.0 { -p * p.ln() } else { 0.0 }
})
.sum();
let diff = (online_h - naive_h).abs();
assert!(
diff < 1e-4,
"online={online_h} naive={naive_h} diff={diff} > 1e-4"
);
}
#[test]
fn test_entropy_single_value() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(5.0);
assert_eq!(acc.entropy_nats(), 0.0, "single value → 0");
assert_eq!(acc.count(), 1);
}
#[test]
fn test_entropy_extreme_logits() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(100.0);
acc.update(-100.0);
let h = acc.entropy_nats();
assert!(
h.is_finite(),
"extreme logits produced non-finite entropy: {h}"
);
assert!(h >= 0.0, "entropy should be non-negative, got {h}");
}
#[test]
fn test_entropy_f32_max_uniform() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(f32::MAX);
acc.update(f32::MAX);
assert_close(
acc.entropy_nats(),
std::f32::consts::LN_2,
"f32::MAX uniform pair",
);
let mut acc3 = OnlineSoftmaxEntropy::new();
for _ in 0..3 {
acc3.update(f32::MAX / 2.0);
}
assert_close(
acc3.entropy_nats(),
(3.0_f32).ln(),
"f32::MAX/2 uniform triple",
);
}
#[test]
fn test_entropy_large_spread() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(1e30);
acc.update(-1e30);
let h = acc.entropy_nats();
assert!(h.is_finite(), "large spread non-finite: {h}");
assert!(h < 0.01, "large spread should be near-peaked, got {h}");
}
#[test]
fn test_entropy_try_update_rejects_non_finite_logits() {
for logit in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
let mut acc = OnlineSoftmaxEntropy::new();
let err = acc
.try_update(logit)
.expect_err("non-finite logit must fail");
assert!(matches!(err, crate::error::InferenceError::InvalidInput(_)));
assert_eq!(acc.count(), 0);
assert_eq!(acc.entropy_nats(), 0.0);
}
}
#[test]
fn test_entropy_f32_extreme_tied_after_underflow() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(-f32::MAX);
acc.update(f32::MAX);
acc.update(f32::MAX);
assert_close(
acc.entropy_nats(),
std::f32::consts::LN_2,
"tied max after underflowed rescale",
);
}
#[test]
fn test_entropy_reset() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(1.0);
acc.update(2.0);
acc.reset();
assert_eq!(acc.count(), 0);
assert_eq!(acc.entropy_nats(), 0.0, "after reset, entropy is 0");
acc.update(1.0);
acc.update(1.0);
assert_close(
acc.entropy_nats(),
std::f32::consts::LN_2,
"after reset reuse",
);
}
#[test]
fn test_entropy_bits_and_normalized() {
let mut acc = OnlineSoftmaxEntropy::new();
for _ in 0..4 {
acc.update(0.0);
}
assert_close(acc.entropy_bits(), 2.0, "4-uniform bits");
assert_close(acc.normalized_entropy(), 1.0, "4-uniform normalized");
}
#[test]
fn test_l2_norm_known() {
let v = [3.0_f32, 4.0];
assert_close(l2_norm(&v), 5.0, "3-4-5 triangle");
let unit = [1.0_f32, 0.0, 0.0];
assert_close(l2_norm(&unit), 1.0, "unit vector");
let zeros = [0.0_f32; 16];
assert_close(l2_norm(&zeros), 0.0, "zero vector");
}
#[test]
fn test_l2_norm_larger() {
let v = vec![1.0_f32; 896];
let expected = (896.0_f32).sqrt();
assert_close(l2_norm(&v), expected, "896-dim unit-filled");
}
#[test]
fn test_l2_norm_large_finite_inputs() {
let h = l2_norm(&[1.0e20_f32, 1.0e20_f32]);
assert!(
h.is_finite(),
"representable norm should stay finite, got {h}"
);
assert!((h - 2.0_f32.sqrt() * 1.0e20_f32).abs() / h < 1e-6);
}
#[test]
fn test_layer_metrics_default_matches_adr061_schema() {
let lm = LayerMetrics::default();
assert_eq!(lm.mode, MetricsMode::Off);
assert_eq!(lm.layer_idx, 0);
assert_eq!(lm.latency_ns, 0);
assert_eq!(lm.input_norm, 0.0);
assert_eq!(lm.output_norm, 0.0);
assert_eq!(lm.update_ratio, 0.0);
assert_eq!(lm.block_influence, 0.0);
assert!(lm.entropy.is_none());
assert!(lm.sparsity.is_none());
assert!(lm.kv_page_mass.is_none());
assert!(lm.pattern_label.is_none());
}
#[test]
fn test_layer_metrics_entropy_is_per_head_vector() {
let lm = LayerMetrics {
mode: MetricsMode::AttentionProfile,
entropy: Some(vec![0.5, 0.8, 1.2, 0.3]),
..LayerMetrics::default()
};
assert_eq!(lm.entropy.as_ref().map(Vec::len), Some(4));
}
#[test]
fn test_l2_norm_empty() {
assert_eq!(l2_norm(&[]), 0.0_f32);
}
#[test]
fn test_l2_norm_single_element() {
assert_close(l2_norm(&[3.0_f32]), 3.0, "positive single");
assert_close(l2_norm(&[-4.0_f32]), 4.0, "negative single");
assert_eq!(l2_norm(&[0.0_f32]), 0.0, "zero single");
}
#[test]
fn test_entropy_split_distribution_vs_naive() {
let mut logits = vec![3.0_f32; 32];
logits.extend(std::iter::repeat_n(-3.0_f32, 32));
let mut acc = OnlineSoftmaxEntropy::new();
for &l in &logits {
acc.update(l);
}
let online_h = acc.entropy_nats();
let max_l = 3.0_f32;
let exps: Vec<f32> = logits.iter().map(|&l| (l - max_l).exp()).collect();
let sum_exp: f32 = exps.iter().sum();
let naive_h: f32 = exps
.iter()
.map(|&e| {
let p = e / sum_exp;
if p > 1e-30 { -p * p.ln() } else { 0.0 }
})
.sum();
let diff = (online_h - naive_h).abs();
assert!(
diff < 1e-4,
"split distribution: online={online_h} naive={naive_h} diff={diff} exceeds 1e-4"
);
assert!(online_h.is_finite(), "split entropy must be finite");
assert!(online_h >= 0.0, "split entropy must be non-negative");
}
#[test]
fn test_entropy_subnormal_inputs() {
let sub = f32::from_bits(1); let logits = [sub, sub, 0.0_f32, -sub];
let mut acc = OnlineSoftmaxEntropy::new();
for &l in &logits {
acc.update(l);
}
let h = acc.entropy_nats();
assert!(
h.is_finite(),
"subnormal inputs produced non-finite entropy: {h}"
);
assert!(h >= 0.0, "subnormal entropy must be non-negative, got {h}");
let max_l = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = logits.iter().map(|&l| (l - max_l).exp()).collect();
let sum_exp: f32 = exps.iter().sum();
let naive_h: f32 = exps
.iter()
.map(|&e| {
let p = e / sum_exp;
if p > 1e-30 { -p * p.ln() } else { 0.0 }
})
.sum();
let diff = (h - naive_h).abs();
assert!(
diff < 1e-4,
"subnormal: online={h} naive={naive_h} diff={diff} exceeds 1e-4"
);
}
#[test]
#[should_panic(expected = "logit must be finite")]
fn test_entropy_update_panics_on_nan() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(f32::NAN);
}
#[test]
#[should_panic(expected = "logit must be finite")]
fn test_entropy_update_panics_on_pos_inf() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(f32::INFINITY);
}
#[test]
#[should_panic(expected = "logit must be finite")]
fn test_entropy_update_panics_on_neg_inf() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(f32::NEG_INFINITY);
}
#[test]
fn test_entropy_try_update_valid_state_after_rejection() {
let mut acc = OnlineSoftmaxEntropy::new();
acc.update(1.0);
let _ = acc.try_update(f32::NAN); acc.update(1.0);
let mut fresh = OnlineSoftmaxEntropy::new();
fresh.update(1.0);
fresh.update(1.0);
assert_close(
acc.entropy_nats(),
fresh.entropy_nats(),
"state after NaN rejection must match fresh two-equal run",
);
assert_eq!(acc.count(), 2, "count must not increment on rejection");
}
#[test]
fn test_entropy_per_head_per_row_semantics() {
const H: usize = 4; const T: usize = 16; let logit_patterns: [fn(usize) -> f32; H] = [
|_| 0.0, |i| if i == 0 { 100.0 } else { -100.0 }, |i| i as f32, |i| (i as f32) * -0.1, ];
let mut head_entropies = Vec::with_capacity(H);
for pattern in &logit_patterns {
let mut acc = OnlineSoftmaxEntropy::new();
for i in 0..T {
acc.update(pattern(i));
}
head_entropies.push(acc.entropy_nats());
}
assert_close(
head_entropies[0],
(T as f32).ln(),
"head 0 (uniform) entropy",
);
assert!(
head_entropies[1] < 0.01,
"head 1 (peaked) entropy should be near 0, got {}",
head_entropies[1]
);
for (i, &h) in head_entropies.iter().enumerate() {
assert!(h.is_finite(), "head {i} entropy non-finite: {h}");
assert!(h >= 0.0, "head {i} entropy negative: {h}");
}
let lm = LayerMetrics {
mode: MetricsMode::AttentionProfile,
entropy: Some(head_entropies.clone()),
..LayerMetrics::default()
};
assert_eq!(lm.entropy.as_ref().map(Vec::len), Some(H));
}
}