use std::collections::BTreeMap;
use super::error::ImatrixError;
#[derive(Debug, Clone)]
pub struct Accumulator {
pub name: String,
pub n_per_row: usize,
pub n_mat: usize,
pub values: Vec<f32>,
pub counts: Vec<i64>,
}
impl Accumulator {
pub fn new(name: String, n_per_row: usize, n_mat: usize) -> Self {
Self {
name,
n_per_row,
n_mat,
values: vec![0.0; n_per_row * n_mat],
counts: vec![0; n_mat],
}
}
pub fn absorb_dense(&mut self, row: &[f32]) -> Result<(), ImatrixError> {
if self.n_mat != 1 {
return Err(ImatrixError::CorpusRead {
path: format!("<accumulator:{}>", self.name),
detail: format!(
"absorb_dense called on MoE-shaped accumulator (n_mat={})",
self.n_mat
),
});
}
if row.len() != self.n_per_row {
return Err(ImatrixError::CorpusRead {
path: format!("<accumulator:{}>", self.name),
detail: format!(
"row length mismatch: got {}, expected {}",
row.len(),
self.n_per_row
),
});
}
for (j, &x) in row.iter().enumerate() {
self.values[j] += x * x;
}
self.counts[0] += 1;
Ok(())
}
pub fn absorb_moe(&mut self, expert_id: usize, row: &[f32]) -> Result<(), ImatrixError> {
if expert_id >= self.n_mat {
return Err(ImatrixError::CorpusRead {
path: format!("<accumulator:{}>", self.name),
detail: format!(
"expert_id={} out of range (n_mat={})",
expert_id, self.n_mat
),
});
}
if row.len() != self.n_per_row {
return Err(ImatrixError::CorpusRead {
path: format!("<accumulator:{}>", self.name),
detail: format!(
"row length mismatch: got {}, expected {}",
row.len(),
self.n_per_row
),
});
}
let e_start = expert_id * self.n_per_row;
for (j, &x) in row.iter().enumerate() {
self.values[e_start + j] += x * x;
}
self.counts[expert_id] += 1;
Ok(())
}
pub fn has_data(&self) -> bool {
self.counts.iter().any(|&c| c > 0)
}
}
#[derive(Debug, Default)]
pub struct AccumulatorRegistry {
inner: BTreeMap<String, Accumulator>,
}
impl AccumulatorRegistry {
pub fn new() -> Self {
Self {
inner: BTreeMap::new(),
}
}
pub fn register(
&mut self,
name: &str,
n_per_row: usize,
n_mat: usize,
) -> Result<&mut Accumulator, ImatrixError> {
if let Some(existing) = self.inner.get(name) {
if existing.n_per_row != n_per_row || existing.n_mat != n_mat {
return Err(ImatrixError::CorpusRead {
path: format!("<accumulator:{name}>"),
detail: format!(
"inconsistent shape on re-register: existing ({},{}) vs new ({},{})",
existing.n_per_row, existing.n_mat, n_per_row, n_mat
),
});
}
} else {
self.inner.insert(
name.to_string(),
Accumulator::new(name.to_string(), n_per_row, n_mat),
);
}
Ok(self
.inner
.get_mut(name)
.expect("entry was just registered or shape-checked"))
}
pub fn get(&self, name: &str) -> Option<&Accumulator> {
self.inner.get(name)
}
pub fn get_mut(&mut self, name: &str) -> Option<&mut Accumulator> {
self.inner.get_mut(name)
}
pub fn iter(&self) -> impl Iterator<Item = (&str, &Accumulator)> {
self.inner.iter().map(|(k, v)| (k.as_str(), v))
}
pub fn len(&self) -> usize {
self.inner.len()
}
pub fn is_empty(&self) -> bool {
self.inner.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn absorb_dense_canonical_formula() {
let mut acc = Accumulator::new("test.weight".to_string(), 3, 1);
acc.absorb_dense(&[1.0, 2.0, 3.0]).unwrap();
acc.absorb_dense(&[4.0, 5.0, 6.0]).unwrap();
assert_eq!(acc.values, vec![1.0 + 16.0, 4.0 + 25.0, 9.0 + 36.0]);
assert_eq!(acc.counts, vec![2]);
}
#[test]
fn absorb_moe_layout() {
let mut acc = Accumulator::new("blk.0.ffn_gate_exps.weight".to_string(), 3, 4);
acc.absorb_moe(2, &[1.0, 2.0, 3.0]).unwrap();
acc.absorb_moe(2, &[1.0, 1.0, 1.0]).unwrap();
acc.absorb_moe(0, &[0.5, 0.5, 0.5]).unwrap();
assert_eq!(acc.counts, vec![1, 0, 2, 0]);
assert_eq!(acc.values[0..3], [0.25, 0.25, 0.25]);
assert_eq!(acc.values[3..6], [0.0, 0.0, 0.0]);
assert_eq!(acc.values[6..9], [2.0, 5.0, 10.0]);
assert_eq!(acc.values[9..12], [0.0, 0.0, 0.0]);
}
#[test]
fn row_length_mismatch_errors() {
let mut acc = Accumulator::new("t".to_string(), 3, 1);
let err = acc.absorb_dense(&[1.0, 2.0]).unwrap_err();
assert!(matches!(err, ImatrixError::CorpusRead { .. }));
}
#[test]
fn dense_on_moe_errors() {
let mut acc = Accumulator::new("t".to_string(), 3, 4);
let err = acc.absorb_dense(&[1.0, 2.0, 3.0]).unwrap_err();
assert!(matches!(err, ImatrixError::CorpusRead { .. }));
}
#[test]
fn moe_expert_out_of_range_errors() {
let mut acc = Accumulator::new("t".to_string(), 3, 4);
let err = acc.absorb_moe(5, &[1.0, 2.0, 3.0]).unwrap_err();
assert!(matches!(err, ImatrixError::CorpusRead { .. }));
}
#[test]
fn registry_register_get() {
let mut reg = AccumulatorRegistry::new();
let acc = reg.register("a.weight", 4, 1).unwrap();
acc.absorb_dense(&[1.0, 2.0, 3.0, 4.0]).unwrap();
let _ = reg.register("b.weight", 8, 2).unwrap();
assert_eq!(reg.len(), 2);
assert!(reg.get("a.weight").is_some());
assert!(reg.get("b.weight").is_some());
assert!(reg.get("nonexistent").is_none());
assert_eq!(reg.get("a.weight").unwrap().counts, vec![1]);
}
#[test]
fn registry_inconsistent_shape_errors() {
let mut reg = AccumulatorRegistry::new();
let _ = reg.register("a.weight", 4, 1).unwrap();
let err = reg.register("a.weight", 8, 1).unwrap_err();
assert!(matches!(err, ImatrixError::CorpusRead { .. }));
}
#[test]
fn registry_iter_sorted() {
let mut reg = AccumulatorRegistry::new();
reg.register("c.weight", 4, 1).unwrap();
reg.register("a.weight", 4, 1).unwrap();
reg.register("b.weight", 4, 1).unwrap();
let names: Vec<&str> = reg.iter().map(|(n, _)| n).collect();
assert_eq!(names, vec!["a.weight", "b.weight", "c.weight"]);
}
}