use std::sync::Arc;
use gam_terms::basis::monomial_exponents;
use ndarray::Array2;
use super::{
CylinderHarmonicEvaluator, DuchonCoordinateEvaluator, EuclideanPatchEvaluator,
MobiusHarmonicEvaluator, PeriodicHarmonicEvaluator, SaeAtomBasisKind, SaeBasisSecondJet,
SphereChartEvaluator, TorusHarmonicEvaluator,
};
pub const SAE_DEFAULT_TORUS_HARMONICS: usize = 3;
pub const SAE_SPHERE_BASIS_SIZE: usize = 7;
pub fn sae_duchon_atom_m(dim: usize) -> usize {
dim / 2 + 2
}
pub const SAE_EUCLIDEAN_PATCH_MAX_DEGREE: usize = 2;
pub const SAE_EUCLIDEAN_PATCH_RECOVERY_MAX_DEGREE: usize = 3;
pub const SAE_CYLINDER_LINE_DEGREE: usize = 2;
pub const SAE_MOBIUS_CIRCLE_HARMONICS: usize = 3;
pub const SAE_MOBIUS_WIDTH_DEGREE: usize = 2;
pub fn sae_cylinder_harmonics_degree(m: usize) -> Result<(usize, usize), String> {
let d_line = SAE_CYLINDER_LINE_DEGREE;
let ml = d_line + 1;
if ml == 0 || m == 0 || m % ml != 0 {
return Err(format!(
"sae_cylinder_harmonics_degree: basis size {m} is not (2H+1)·{ml} for a cylinder \
with line degree {d_line}"
));
}
let mc = m / ml;
if mc < 3 || mc % 2 == 0 {
return Err(format!(
"sae_cylinder_harmonics_degree: recovered circle width {mc} (= {m}/{ml}) is not an \
odd 2H+1 ≥ 3 for a cylinder"
));
}
Ok(((mc - 1) / 2, d_line))
}
pub const SAE_MAX_PERIODIC_HARMONICS: usize = 4096;
pub fn sae_periodic_basis_size(n_harmonics: usize) -> Result<usize, String> {
if n_harmonics > SAE_MAX_PERIODIC_HARMONICS {
return Err(format!(
"sae_build_periodic_atom: n_harmonics={n_harmonics} exceeds dense limit {SAE_MAX_PERIODIC_HARMONICS}"
));
}
n_harmonics
.checked_mul(2)
.and_then(|twice| twice.checked_add(1))
.ok_or_else(|| {
format!("sae_build_periodic_atom: basis size overflows for n_harmonics={n_harmonics}")
})
}
pub fn sae_torus_axis_basis_size(m: usize, d: usize) -> Result<usize, String> {
if d == 0 {
return Err("sae_torus_axis_basis_size: d must be >= 1".to_string());
}
if m == 0 {
return Err("sae_torus_axis_basis_size: m must be >= 1".to_string());
}
let mut axis_m: usize = 1;
loop {
let mut prod: usize = 1;
let mut overflow = false;
for _ in 0..d {
match prod.checked_mul(axis_m) {
Some(p) => prod = p,
None => {
overflow = true;
break;
}
}
}
if overflow || prod > m {
return Err(format!(
"sae_torus_axis_basis_size: m={m} is not a perfect d-th power for d={d}"
));
}
if prod == m {
if axis_m % 2 == 0 {
return Err(format!(
"sae_torus_axis_basis_size: m={m} = {axis_m}^{d} but axis size must be odd (2H+1)"
));
}
return Ok(axis_m);
}
axis_m += 1;
}
}
pub fn sae_euclidean_degree_for_basis_size(dim: usize, basis_size: usize) -> Result<usize, String> {
for degree in 0..=SAE_EUCLIDEAN_PATCH_RECOVERY_MAX_DEGREE {
if monomial_exponents(dim, degree).len() == basis_size {
return Ok(degree);
}
}
Err(format!(
"euclidean patch basis size {basis_size} is not a valid monomial width for latent_dim={dim} with max_degree<={SAE_EUCLIDEAN_PATCH_RECOVERY_MAX_DEGREE}"
))
}
pub fn build_sae_basis_evaluators(
basis_kinds: &[SaeAtomBasisKind],
basis_sizes: &[usize],
atom_dim: &[usize],
coord_blocks: &[Array2<f64>],
atom_centers: &[Option<Array2<f64>>],
) -> Result<Vec<Option<Arc<dyn SaeBasisSecondJet>>>, String> {
let k_atoms = basis_kinds.len();
if atom_dim.len() != k_atoms
|| basis_sizes.len() != k_atoms
|| coord_blocks.len() != k_atoms
|| atom_centers.len() != k_atoms
{
return Err(format!(
"build_sae_basis_evaluators: K-length metadata mismatch (kinds={k_atoms}, dims={}, sizes={}, coords={}, centers={})",
atom_dim.len(),
basis_sizes.len(),
coord_blocks.len(),
atom_centers.len()
));
}
let mut out: Vec<Option<Arc<dyn SaeBasisSecondJet>>> = Vec::with_capacity(k_atoms);
for k in 0..k_atoms {
let m = basis_sizes[k];
let d = atom_dim[k];
let evaluator: Arc<dyn SaeBasisSecondJet> = match &basis_kinds[k] {
SaeAtomBasisKind::Periodic if d == 1 && m % 2 == 1 => {
Arc::new(PeriodicHarmonicEvaluator::new(m)?)
}
SaeAtomBasisKind::Sphere if d == 2 && m == SAE_SPHERE_BASIS_SIZE => {
Arc::new(SphereChartEvaluator)
}
SaeAtomBasisKind::Torus if d >= 1 => {
let axis_m = sae_torus_axis_basis_size(m, d)?;
let h = (axis_m - 1) / 2;
Arc::new(TorusHarmonicEvaluator::new(d, h)?)
}
SaeAtomBasisKind::Duchon => {
let centers = atom_centers[k].as_ref().ok_or_else(|| {
format!(
"build_sae_basis_evaluators: Duchon atom {k} cannot refresh its basis without centers; \
build the atom through the SAE auto path so its Duchon centers are threaded in"
)
})?;
Arc::new(DuchonCoordinateEvaluator::new(
centers.clone(),
sae_duchon_atom_m(centers.ncols()),
)?)
}
SaeAtomBasisKind::Cylinder if d == 2 => {
let (h, d_line) = sae_cylinder_harmonics_degree(m)?;
Arc::new(CylinderHarmonicEvaluator::new(h, d_line)?)
}
SaeAtomBasisKind::Linear => Arc::new(EuclideanPatchEvaluator::new(
d,
sae_euclidean_degree_for_basis_size(d, m)?,
)?),
SaeAtomBasisKind::EuclideanPatch | SaeAtomBasisKind::Poincare => {
Arc::new(EuclideanPatchEvaluator::new(
d,
sae_euclidean_degree_for_basis_size(d, m)?,
)?)
}
SaeAtomBasisKind::Mobius if d == 2 => {
let evaluator = MobiusHarmonicEvaluator::new(
SAE_MOBIUS_CIRCLE_HARMONICS,
SAE_MOBIUS_WIDTH_DEGREE,
)?;
if evaluator.basis_size() != m {
return Err(format!(
"build_sae_basis_evaluators: Mobius atom {k} width {m} does not match \
the production deck-invariant layout ({} columns)",
evaluator.basis_size()
));
}
Arc::new(evaluator)
}
SaeAtomBasisKind::Mobius => {
return Err(format!(
"build_sae_basis_evaluators: Mobius atom {k} requires latent_dim == 2; got dim={d}, m={m}"
));
}
SaeAtomBasisKind::Cylinder => {
return Err(format!(
"build_sae_basis_evaluators: Cylinder atom {k} requires latent_dim == 2; got dim={d}, m={m}"
));
}
SaeAtomBasisKind::FiniteSet => {
return Err(format!(
"build_sae_basis_evaluators: atom {k} 'finite_set' is a discrete-anchor \
(categorical) candidate with no continuous Phi(t) refresh; it is not yet \
wired into the inner Newton latent-update path"
));
}
SaeAtomBasisKind::Precomputed(label) => {
return Err(format!(
"build_sae_basis_evaluators: atom {k} basis {label:?} is precomputed and has no \
analytic refresh routine; the inner Newton latent update requires a basis kind \
that can re-evaluate Phi(t)/dPhi/dt at updated coordinates"
));
}
SaeAtomBasisKind::Periodic => {
return Err(format!(
"build_sae_basis_evaluators: Periodic atom {k} requires latent_dim == 1 and odd basis size; got dim={d}, m={m}"
));
}
SaeAtomBasisKind::Sphere => {
return Err(format!(
"build_sae_basis_evaluators: Sphere atom {k} requires latent_dim == 2 and basis size {SAE_SPHERE_BASIS_SIZE}; got dim={d}, m={m}"
));
}
SaeAtomBasisKind::Torus => {
return Err(format!(
"build_sae_basis_evaluators: Torus atom {k} requires latent_dim >= 1; got dim={d}, m={m}"
));
}
};
out.push(Some(evaluator));
}
Ok(out)
}
pub fn sae_atom_basis_kind_from_str(value: &str) -> SaeAtomBasisKind {
let canonical = crate::atom_schema::canonical_basis_kind(value);
match canonical.as_str() {
"duchon" => SaeAtomBasisKind::Duchon,
"periodic" => SaeAtomBasisKind::Periodic,
"sphere" => SaeAtomBasisKind::Sphere,
"torus" => SaeAtomBasisKind::Torus,
"linear" => SaeAtomBasisKind::Linear,
"linear_block" => SaeAtomBasisKind::Linear,
"euclidean" => SaeAtomBasisKind::EuclideanPatch,
"poincare" => SaeAtomBasisKind::Poincare,
"cylinder" => SaeAtomBasisKind::Cylinder,
"mobius" => SaeAtomBasisKind::Mobius,
other => SaeAtomBasisKind::Precomputed(other.to_string()),
}
}
pub fn sae_atom_basis_kind_name(kind: &SaeAtomBasisKind) -> String {
match kind {
SaeAtomBasisKind::Periodic => "periodic".to_string(),
SaeAtomBasisKind::Duchon => "duchon".to_string(),
SaeAtomBasisKind::Sphere => "sphere".to_string(),
SaeAtomBasisKind::Torus => "torus".to_string(),
SaeAtomBasisKind::Linear => "linear".to_string(),
SaeAtomBasisKind::EuclideanPatch => "euclidean_patch".to_string(),
SaeAtomBasisKind::Poincare => "poincare".to_string(),
SaeAtomBasisKind::Cylinder => "cylinder".to_string(),
SaeAtomBasisKind::Mobius => "mobius".to_string(),
SaeAtomBasisKind::FiniteSet => "finite_set".to_string(),
SaeAtomBasisKind::Precomputed(name) => name.clone(),
}
}