use gam_problem::LatentRetractionRegistry;
use gam_terms::latent::{LatentIdMode, LatentManifold};
use ndarray::Array1;
use crate::assignment::{AssignmentMode, SaeAssignment};
#[derive(Debug, Clone)]
struct AtomCoordMeta {
latent_dim: usize,
manifold: LatentManifold,
retraction: LatentRetractionRegistry,
}
#[derive(Debug, Clone)]
pub struct SaeAssignmentAtomSpec {
pub latent_dim: usize,
pub id_mode: LatentIdMode,
pub manifold: LatentManifold,
pub retraction: LatentRetractionRegistry,
pub latent_id: u64,
}
impl SaeAssignmentAtomSpec {
#[must_use]
pub fn euclidean(latent_dim: usize) -> Self {
Self {
latent_dim,
id_mode: LatentIdMode::None,
manifold: LatentManifold::Euclidean,
retraction: LatentRetractionRegistry::all_euclidean(),
latent_id: 0,
}
}
}
#[derive(Debug, Clone)]
pub struct SaeAssignmentState {
n_obs: usize,
k_atoms: usize,
indices: Vec<Vec<u32>>,
gate_params: Vec<Vec<f64>>,
coords: Vec<Vec<f64>>,
atom_coord_meta: Vec<AtomCoordMeta>,
mode: AssignmentMode,
}
impl SaeAssignmentState {
#[must_use = "state build error must be handled"]
pub fn from_topk_support_heterogeneous(
n_obs: usize,
k_atoms: usize,
support_k: usize,
atom_specs: Vec<SaeAssignmentAtomSpec>,
mut indices: Vec<Vec<u32>>,
mut gate_params: Vec<Vec<f64>>,
mut coords: Vec<Vec<f64>>,
) -> Result<Self, String> {
if support_k == 0 || support_k > k_atoms {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: support_k must satisfy 1 <= s <= K={k_atoms}; got {support_k}"
));
}
if atom_specs.len() != k_atoms {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: atom_specs length {} must equal K={k_atoms}",
atom_specs.len()
));
}
for (atom, spec) in atom_specs.iter().enumerate() {
if spec.latent_dim == 0 {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: atom {atom} latent_dim must be positive"
));
}
let manifold_dim = spec.manifold.ambient_dim(spec.latent_dim);
if manifold_dim != spec.latent_dim {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: atom {atom} manifold ambient dimension {manifold_dim} != latent_dim {}",
spec.latent_dim
));
}
spec.retraction.validate_dim(
spec.latent_dim,
"SaeAssignmentState::from_topk_support_heterogeneous",
)?;
}
if indices.len() != n_obs || gate_params.len() != n_obs || coords.len() != n_obs {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: per-row arrays must all have length N={n_obs}; \
got indices={}, gate_params={}, coords={}",
indices.len(),
gate_params.len(),
coords.len()
));
}
for i in 0..n_obs {
if indices[i].len() > support_k
|| indices[i].is_empty()
|| gate_params[i].len() != indices[i].len()
{
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: row {i} widths must be indices={support_k}, gate_params={support_k}; got {}, {}",
indices[i].len(),
gate_params[i].len(),
));
}
if gate_params[i].iter().any(|value| !value.is_finite()) {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: row {i} contains a non-finite gate parameter"
));
}
if coords[i].iter().any(|value| !value.is_finite()) {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: row {i} contains a non-finite coordinate"
));
}
let mut coord_cursor = 0usize;
let mut slots = Vec::with_capacity(indices[i].len());
for slot in 0..indices[i].len() {
let atom = indices[i][slot] as usize;
if atom >= k_atoms {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: row {i} atom index {atom} out of range K={k_atoms}"
));
}
let d = atom_specs[atom].latent_dim;
let end = coord_cursor.saturating_add(d);
if end > coords[i].len() {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: row {i} coordinate width {} is too short for its declared support",
coords[i].len()
));
}
slots.push((
atom as u32,
gate_params[i][slot],
coords[i][coord_cursor..end].to_vec(),
));
coord_cursor = end;
}
if coord_cursor != coords[i].len() {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: row {i} coordinate width {} != support-implied width {coord_cursor}",
coords[i].len()
));
}
slots.sort_by_key(|slot| slot.0);
if slots.windows(2).any(|pair| pair[0].0 == pair[1].0) {
return Err(format!(
"SaeAssignmentState::from_topk_support_heterogeneous: row {i} support contains a duplicate atom"
));
}
indices[i] = slots.iter().map(|slot| slot.0).collect();
gate_params[i] = slots.iter().map(|slot| slot.1).collect();
coords[i] = slots.into_iter().flat_map(|slot| slot.2).collect();
}
let atom_coord_meta = atom_specs
.into_iter()
.map(|spec| AtomCoordMeta {
latent_dim: spec.latent_dim,
manifold: spec.manifold,
retraction: spec.retraction,
})
.collect();
Ok(Self {
n_obs,
k_atoms,
indices,
gate_params,
coords,
atom_coord_meta,
mode: AssignmentMode::top_k_support(support_k),
})
}
pub fn n_obs(&self) -> usize {
self.n_obs
}
pub fn k_atoms(&self) -> usize {
self.k_atoms
}
pub fn mode(&self) -> AssignmentMode {
self.mode
}
pub fn support_indices(&self, row: usize) -> &[u32] {
&self.indices[row]
}
pub fn gate_params(&self, row: usize) -> &[f64] {
&self.gate_params[row]
}
pub fn coords_row(&self, row: usize) -> &[f64] {
&self.coords[row]
}
pub fn atom_coord_dim(&self, atom: usize) -> usize {
self.atom_coord_meta[atom].latent_dim
}
pub fn atom_axis_periods(&self, atom: usize) -> Vec<Option<f64>> {
let meta = &self.atom_coord_meta[atom];
let periods = if meta.manifold.is_euclidean() {
meta.retraction.axis_periods(meta.latent_dim)
} else {
meta.manifold.axis_periods()
};
assert_eq!(
periods.len(),
meta.latent_dim,
"SaeAssignmentState atom {atom} axis-period count {} != latent dimension {}",
periods.len(),
meta.latent_dim,
);
periods
}
pub fn coords_for_slot(&self, row: usize, slot: usize) -> &[f64] {
let start: usize = self.indices[row][..slot]
.iter()
.map(|&atom| self.atom_coord_meta[atom as usize].latent_dim)
.sum();
let atom = self.indices[row][slot] as usize;
&self.coords[row][start..start + self.atom_coord_meta[atom].latent_dim]
}
pub fn take_coords(&mut self) -> Vec<Vec<f64>> {
std::mem::take(&mut self.coords)
}
pub fn restore_coords(&mut self, coords: Vec<Vec<f64>>) -> Result<(), String> {
if coords.len() != self.n_obs {
return Err(format!(
"SaeAssignmentState::restore_coords: {} rows != N={}",
coords.len(),
self.n_obs
));
}
self.coords = coords;
Ok(())
}
pub fn project_row_tangent(
&self,
row: usize,
coords: &[f64],
delta: &mut [f64],
) -> Result<(), String> {
if row >= self.n_obs {
return Err(format!(
"SaeAssignmentState::project_row_tangent: row {row} out of range N={}",
self.n_obs
));
}
if delta.len() != coords.len() {
return Err(format!(
"SaeAssignmentState::project_row_tangent: row {row} delta width {} != compact coordinate width {}",
delta.len(),
coords.len()
));
}
let mut cursor = 0usize;
for &atom in &self.indices[row] {
let meta = &self.atom_coord_meta[atom as usize];
let end = cursor + meta.latent_dim;
let point = Array1::from_vec(coords[cursor..end].to_vec());
let vector = Array1::from_vec(delta[cursor..end].to_vec());
let projected = meta
.manifold
.project_to_tangent(point.view(), vector.view());
delta[cursor..end].copy_from_slice(
projected
.as_slice()
.expect("tangent projection is contiguous"),
);
cursor = end;
}
Ok(())
}
pub fn retract_row_coords(
&self,
row: usize,
coords: &mut [f64],
delta: &[f64],
) -> Result<(), String> {
if delta.len() != coords.len() {
return Err(format!(
"SaeAssignmentState::retract_row_coords: row {row} delta width {} != compact coordinate width {}",
delta.len(),
coords.len()
));
}
let mut cursor = 0usize;
for slot in 0..self.indices[row].len() {
let atom = self.indices[row][slot] as usize;
let meta = &self.atom_coord_meta[atom];
let end = cursor + meta.latent_dim;
let mut current = Array1::from_vec(coords[cursor..end].to_vec());
let step = Array1::from_vec(delta[cursor..end].to_vec());
if meta.retraction.is_all_euclidean() {
current = meta.manifold.retract(current.view(), step.view());
} else {
meta.retraction
.retract(&mut current.view_mut(), step.view());
}
coords[cursor..end]
.copy_from_slice(current.as_slice().expect("retraction is contiguous"));
cursor = end;
}
Ok(())
}
pub fn project_row_coords(
&self,
row: usize,
values: &[f64],
coords: &mut [f64],
) -> Result<(), String> {
if row >= self.n_obs {
return Err(format!(
"SaeAssignmentState::project_row_coords: row {row} out of range N={}",
self.n_obs
));
}
if values.len() != coords.len() {
return Err(format!(
"SaeAssignmentState::project_row_coords: row {row} value width {} != compact coordinate width {}",
values.len(),
coords.len()
));
}
if values.iter().any(|value| !value.is_finite()) {
return Err(format!(
"SaeAssignmentState::project_row_coords: row {row} contains a non-finite coordinate"
));
}
let mut cursor = 0usize;
for &atom in &self.indices[row] {
let meta = &self.atom_coord_meta[atom as usize];
let end = cursor + meta.latent_dim;
let candidate = Array1::from_vec(values[cursor..end].to_vec());
let projected = meta.manifold.project_point(candidate.view());
coords[cursor..end]
.copy_from_slice(projected.as_slice().expect("projection is contiguous"));
cursor = end;
}
Ok(())
}
pub fn convert_atom_to_euclidean(&mut self, atom: usize) -> Result<(), String> {
if atom >= self.k_atoms {
return Err(format!(
"SaeAssignmentState::convert_atom_to_euclidean: atom {atom} out of range K={}",
self.k_atoms
));
}
let meta = &mut self.atom_coord_meta[atom];
meta.manifold = LatentManifold::Euclidean;
meta.retraction = LatentRetractionRegistry::all_euclidean();
Ok(())
}
pub fn set_slot_coords(
&mut self,
row: usize,
slot: usize,
values: &[f64],
) -> Result<(), String> {
if row >= self.n_obs || slot >= self.indices[row].len() {
return Err(format!(
"SaeAssignmentState::set_slot_coords: row {row} slot {slot} out of range"
));
}
let atom = self.indices[row][slot] as usize;
let width = self.atom_coord_meta[atom].latent_dim;
if values.len() != width {
return Err(format!(
"SaeAssignmentState::set_slot_coords: value width {} != atom width {width}",
values.len()
));
}
if values.iter().any(|value| !value.is_finite()) {
return Err(format!(
"SaeAssignmentState::set_slot_coords: row {row} slot {slot} non-finite coordinate"
));
}
let start: usize = self.indices[row][..slot]
.iter()
.map(|&prior| self.atom_coord_meta[prior as usize].latent_dim)
.sum();
let candidate = Array1::from_vec(values.to_vec());
let projected = self.atom_coord_meta[atom]
.manifold
.project_point(candidate.view());
self.coords[row][start..start + width]
.copy_from_slice(projected.as_slice().expect("projection is contiguous"));
Ok(())
}
pub fn set_row_coords(&mut self, row: usize, values: &[f64]) -> Result<(), String> {
if row >= self.n_obs {
return Err(format!(
"SaeAssignmentState::set_row_coords: row {row} out of range N={}",
self.n_obs
));
}
if values.len() != self.coords[row].len() {
return Err(format!(
"SaeAssignmentState::set_row_coords: row {row} value width {} != compact coordinate width {}",
values.len(),
self.coords[row].len()
));
}
if values.iter().any(|value| !value.is_finite()) {
return Err(format!(
"SaeAssignmentState::set_row_coords: row {row} contains a non-finite coordinate"
));
}
let mut cursor = 0usize;
for &atom in &self.indices[row] {
let meta = &self.atom_coord_meta[atom as usize];
let end = cursor + meta.latent_dim;
let candidate = Array1::from_vec(values[cursor..end].to_vec());
let projected = meta.manifold.project_point(candidate.view());
self.coords[row][cursor..end]
.copy_from_slice(projected.as_slice().expect("projection is contiguous"));
cursor = end;
}
Ok(())
}
}
impl SaeAssignment {
}