use super::*;
#[derive(Debug, Clone)]
pub struct SaeRowLayout {
pub active_atoms: Vec<Vec<usize>>,
pub coord_starts: Vec<Vec<usize>>,
pub coord_offsets_full: Vec<usize>,
pub coord_dims: Vec<usize>,
}
impl SaeRowLayout {
pub(crate) fn from_assignment_state(
state: &crate::assignment_state::SaeAssignmentState,
) -> Result<Self, String> {
let mut coord_offsets_full = Vec::with_capacity(state.k_atoms());
let mut cursor = 0usize;
let mut coord_dims = Vec::with_capacity(state.k_atoms());
for atom in 0..state.k_atoms() {
coord_offsets_full.push(cursor);
let d = state.atom_coord_dim(atom);
coord_dims.push(d);
cursor += d;
}
let active_atoms = (0..state.n_obs())
.map(|row| {
state
.support_indices(row)
.iter()
.map(|&atom| atom as usize)
.collect::<Vec<_>>()
})
.collect::<Vec<_>>();
let mut coord_starts = Vec::with_capacity(state.n_obs());
for active in &active_atoms {
let mut row_cursor = 0usize;
let mut starts = Vec::with_capacity(active.len());
for &atom in active {
starts.push(row_cursor);
row_cursor += coord_dims[atom];
}
coord_starts.push(starts);
}
Ok(Self {
active_atoms,
coord_starts,
coord_offsets_full,
coord_dims,
})
}
pub(crate) fn from_topk_gates(
assignments: &[Array1<f64>],
support_size: usize,
coord_dims: Vec<usize>,
coord_offsets_full: Vec<usize>,
) -> Result<Self, String> {
if support_size == 0 {
return Err("SaeRowLayout::from_topk_gates requires positive support_size".to_string());
}
let mut per_row = Vec::with_capacity(assignments.len());
for (row, gates) in assignments.iter().enumerate() {
let mut active = Vec::with_capacity(support_size);
for (atom, &gate) in gates.iter().enumerate() {
if gate == 1.0 {
active.push(atom);
} else if gate != 0.0 {
return Err(format!(
"SaeRowLayout::from_topk_gates: row {row}, atom {atom} has non-binary gate {gate}"
));
}
}
if active.len() != support_size.min(gates.len()) {
return Err(format!(
"SaeRowLayout::from_topk_gates: row {row} has {} active atoms; expected {}",
active.len(),
support_size.min(gates.len())
));
}
per_row.push(active);
}
let mut coord_starts = Vec::with_capacity(per_row.len());
for active in &per_row {
let mut cursor = 0usize;
let mut starts = Vec::with_capacity(active.len());
for &atom in active {
starts.push(cursor);
cursor += coord_dims[atom];
}
coord_starts.push(starts);
}
Ok(Self {
active_atoms: per_row,
coord_starts,
coord_offsets_full,
coord_dims,
})
}
pub fn row_q_active(&self, row: usize) -> usize {
let active = &self.active_atoms[row];
let coord_sum: usize = active.iter().map(|&k| self.coord_dims[k]).sum();
coord_sum
}
pub fn expand_row(&self, row: usize, delta_t_row: &[f64], out: &mut [f64]) {
for v in out.iter_mut() {
*v = 0.0;
}
let active = &self.active_atoms[row];
let starts = &self.coord_starts[row];
for (pos, &k) in active.iter().enumerate() {
let d = self.coord_dims[k];
let full_off = self.coord_offsets_full[k];
for axis in 0..d {
out[full_off + axis] = delta_t_row[starts[pos] + axis];
}
}
}
}
#[cfg(test)]
mod support_state_tests {
use super::*;
use crate::assignment_state::{SaeAssignmentAtomSpec, SaeAssignmentState};
#[test]
fn direct_layout_preserves_heterogeneous_active_offsets() {
let state = SaeAssignmentState::from_topk_support_heterogeneous(
2,
4,
2,
vec![
SaeAssignmentAtomSpec::euclidean(1),
SaeAssignmentAtomSpec::euclidean(3),
SaeAssignmentAtomSpec::euclidean(2),
SaeAssignmentAtomSpec::euclidean(1),
],
vec![vec![2, 0], vec![3, 1]],
vec![vec![1.0; 2]; 2],
vec![vec![0.0; 3], vec![0.0; 4]],
)
.expect("state builds");
let layout = SaeRowLayout::from_assignment_state(&state).expect("layout builds");
assert_eq!(layout.active_atoms, vec![vec![0, 2], vec![1, 3]]);
assert_eq!(layout.coord_starts, vec![vec![0, 1], vec![0, 3]]);
assert_eq!(layout.coord_dims, vec![1, 3, 2, 1]);
assert_eq!(layout.coord_offsets_full, vec![0, 1, 4, 6]);
assert_eq!(layout.row_q_active(0), 3);
assert_eq!(layout.row_q_active(1), 4);
}
}