use std::collections::{BTreeMap, HashMap};
use std::sync::Mutex;
use frink_core::WeightMatrix;
use super::file::Stats;
#[derive(Debug, Clone)]
struct Slot {
name: String,
expert: Option<(usize, usize)>,
cols: usize,
}
#[derive(Default)]
pub struct Collector {
slots: HashMap<usize, Slot>,
stats: Mutex<BTreeMap<String, Stats>>,
non_finite: Mutex<Option<String>>,
}
impl Collector {
fn key(m: &WeightMatrix) -> usize {
m as *const WeightMatrix as usize
}
pub fn register_dense(&mut self, m: &WeightMatrix, name: &str) {
self.slots.insert(
Self::key(m),
Slot {
name: name.to_string(),
expert: None,
cols: m.cols(),
},
);
self.stats
.get_mut()
.unwrap_or_else(|e| e.into_inner())
.entry(name.to_string())
.or_insert_with(|| Stats {
values: vec![0.0; m.cols()],
counts: vec![0],
});
}
pub fn register_expert(
&mut self,
m: &WeightMatrix,
name: &str,
expert: usize,
n_experts: usize,
) {
assert!(expert < n_experts);
self.slots.insert(
Self::key(m),
Slot {
name: name.to_string(),
expert: Some((expert, n_experts)),
cols: m.cols(),
},
);
self.stats
.get_mut()
.unwrap_or_else(|e| e.into_inner())
.entry(name.to_string())
.or_insert_with(|| Stats {
values: vec![0.0; m.cols() * n_experts],
counts: vec![0; n_experts],
});
}
pub fn n_registered(&self) -> usize {
self.slots.len()
}
pub fn observe(&self, m: &WeightMatrix, rows: &[f32], n_rows: usize) {
let Some(slot) = self.slots.get(&Self::key(m)) else {
return;
};
let cols = slot.cols;
debug_assert_eq!(rows.len(), n_rows * cols);
let mut stats = self.stats.lock().unwrap_or_else(|e| e.into_inner());
let e = stats
.get_mut(&slot.name)
.expect("every registered slot has an entry");
let (mat, count_idx) = match slot.expert {
Some((ex, _)) => (ex, ex),
None => (0, 0),
};
let start = mat * cols;
let values = &mut e.values[start..start + cols];
for row in rows.chunks_exact(cols) {
for (v, &x) in values.iter_mut().zip(row) {
*v = x.mul_add(x, *v);
}
}
e.counts[count_idx] += n_rows as i64;
if let Some(bad) = values.iter().find(|v| !v.is_finite()) {
let mut nf = self.non_finite.lock().unwrap_or_else(|e| e.into_inner());
if nf.is_none() {
*nf = Some(format!("{bad} detected in {}", slot.name));
}
}
}
pub fn non_finite(&self) -> Option<String> {
self.non_finite
.lock()
.unwrap_or_else(|e| e.into_inner())
.clone()
}
pub fn stats(&self) -> BTreeMap<String, Stats> {
self.stats.lock().unwrap_or_else(|e| e.into_inner()).clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use frink_core::tensor::Tensor;
fn matrix(rows: usize, cols: usize) -> WeightMatrix {
WeightMatrix::F32(Tensor::new(vec![0.5; rows * cols], vec![rows, cols]))
}
#[test]
fn dense_sums_squares_per_column_and_counts_rows() {
let m = matrix(3, 4);
let mut c = Collector::default();
c.register_dense(&m, "blk.0.attn_q.weight");
c.observe(&m, &[1.0, 2.0, 3.0, 4.0, 1.0, 1.0, 1.0, 1.0], 2);
c.observe(&m, &[0.0, 0.0, 0.0, 2.0, 3.0, 0.0, 0.0, 0.0], 2);
let s = &c.stats()["blk.0.attn_q.weight"];
assert_eq!(s.values, vec![11.0, 5.0, 10.0, 21.0]);
assert_eq!(s.counts, vec![4]);
}
#[test]
fn an_unregistered_matrix_is_ignored() {
let m = matrix(2, 4);
let other = matrix(2, 4);
let mut c = Collector::default();
c.register_dense(&m, "blk.0.attn_q.weight");
c.observe(&other, &[9.0; 4], 1);
let s = c.stats();
assert_eq!(s.len(), 1);
assert_eq!(s["blk.0.attn_q.weight"].counts, vec![0]);
}
#[test]
fn experts_accumulate_into_their_own_matrix_and_count() {
let e0 = matrix(2, 3);
let e1 = matrix(2, 3);
let mut c = Collector::default();
c.register_expert(&e0, "blk.0.ffn_gate_exps.weight", 0, 2);
c.register_expert(&e1, "blk.0.ffn_gate_exps.weight", 1, 2);
c.observe(&e1, &[1.0, 2.0, 3.0, 1.0, 1.0, 1.0], 2);
c.observe(&e0, &[2.0, 0.0, 0.0], 1);
let s = &c.stats()["blk.0.ffn_gate_exps.weight"];
assert_eq!(s.values, vec![4.0, 0.0, 0.0, 2.0, 5.0, 10.0]);
assert_eq!(s.counts, vec![1, 2]);
}
#[test]
fn a_non_finite_sum_is_reported_with_its_entry() {
let m = matrix(1, 2);
let mut c = Collector::default();
c.register_dense(&m, "blk.3.ffn_down.weight");
c.observe(&m, &[f32::MAX, 1.0], 1);
c.observe(&m, &[f32::MAX, 1.0], 1);
let msg = c.non_finite().expect("overflow must be reported");
assert!(msg.contains("blk.3.ffn_down.weight"), "{msg}");
}
}