use ndarray::ArrayView2;
use super::atom::SaeAtomBasisKind;
use super::term::SaeManifoldTerm;
pub const EMBEDDEDNESS_CENTER_NODES: usize = 512;
pub const EMBEDDEDNESS_SEPARATION_NODES: usize = 257;
#[derive(Debug, Clone)]
pub struct AtomEmbeddednessCertificate {
pub harmonics: usize,
pub grid_min: f64,
pub grid_correction: f64,
pub certified_min: f64,
pub scale: f64,
pub relative_margin: f64,
pub embedded: bool,
pub center_nodes: usize,
pub separation_nodes: usize,
}
pub fn certify_periodic_decoder_embeddedness(
decoder: ArrayView2<'_, f64>,
) -> Result<AtomEmbeddednessCertificate, String> {
let m = decoder.nrows();
if m == 0 || m % 2 == 0 {
return Err(format!(
"certify_periodic_decoder_embeddedness: periodic decoder needs an odd \
row count 2H+1; got {m}"
));
}
let harmonics = (m - 1) / 2;
if harmonics == 0 {
return Ok(AtomEmbeddednessCertificate {
harmonics: 0,
grid_min: 0.0,
grid_correction: 0.0,
certified_min: 0.0,
scale: 0.0,
relative_margin: 0.0,
embedded: false,
center_nodes: EMBEDDEDNESS_CENTER_NODES,
separation_nodes: EMBEDDEDNESS_SEPARATION_NODES,
});
}
let sin_row = |h: usize| decoder.row(2 * h - 1);
let cos_row = |h: usize| decoder.row(2 * h);
let mut g_aa = vec![0.0_f64; harmonics * harmonics];
let mut g_ab = vec![0.0_f64; harmonics * harmonics];
let mut g_bb = vec![0.0_f64; harmonics * harmonics];
for h in 1..=harmonics {
for k in 1..=harmonics {
let idx = (h - 1) * harmonics + (k - 1);
g_aa[idx] = sin_row(h).dot(&sin_row(k));
g_ab[idx] = sin_row(h).dot(&cos_row(k));
g_bb[idx] = cos_row(h).dot(&cos_row(k));
}
}
let mut f_bar = 0.0_f64;
let mut c_bar = 0.0_f64;
let mut x_bar = 0.0_f64;
for h in 1..=harmonics {
let idx = (h - 1) * harmonics + (h - 1);
let a_h = (g_aa[idx] + g_bb[idx]).max(0.0).sqrt();
let hf = h as f64;
f_bar += hf * a_h;
c_bar += std::f64::consts::TAU * hf * hf * a_h;
x_bar += (hf - 1.0) * hf * (hf + 1.0) / 3.0 * a_h;
}
let center_nodes = EMBEDDEDNESS_CENTER_NODES;
let separation_nodes = EMBEDDEDNESS_SEPARATION_NODES;
let half_c = 0.5 / center_nodes as f64;
let half_x = if separation_nodes > 1 {
1.0 / (separation_nodes - 1) as f64
} else {
1.0
};
let grid_correction = half_c * 2.0 * f_bar * c_bar + half_x * 2.0 * f_bar * x_bar;
let mut cheb = vec![0.0_f64; harmonics];
let mut quad = vec![0.0_f64; harmonics * harmonics];
let mut cos_c = vec![0.0_f64; harmonics];
let mut sin_c = vec![0.0_f64; harmonics];
let mut grid_min = f64::INFINITY;
for ci in 0..center_nodes {
let c = ci as f64 / center_nodes as f64;
for h in 1..=harmonics {
let angle = std::f64::consts::TAU * h as f64 * c;
cos_c[h - 1] = angle.cos();
sin_c[h - 1] = angle.sin();
}
for h in 1..=harmonics {
for k in 1..=harmonics {
let idx = (h - 1) * harmonics + (k - 1);
let idx_t = (k - 1) * harmonics + (h - 1);
quad[idx] = cos_c[h - 1] * cos_c[k - 1] * g_aa[idx]
- cos_c[h - 1] * sin_c[k - 1] * g_ab[idx]
- sin_c[h - 1] * cos_c[k - 1] * g_ab[idx_t]
+ sin_c[h - 1] * sin_c[k - 1] * g_bb[idx];
}
}
for xi in 0..separation_nodes {
let x = if separation_nodes > 1 {
-1.0 + 2.0 * xi as f64 / (separation_nodes - 1) as f64
} else {
0.0
};
for h in 1..=harmonics {
cheb[h - 1] = match h {
1 => 1.0,
2 => 2.0 * x,
_ => 2.0 * x * cheb[h - 2] - cheb[h - 3],
};
}
let mut value = 0.0_f64;
for h in 0..harmonics {
for k in 0..harmonics {
value += cheb[h] * cheb[k] * quad[h * harmonics + k];
}
}
grid_min = grid_min.min(value);
}
}
let certified_min = grid_min - grid_correction;
let scale = f_bar * f_bar;
let relative_margin = if scale > 0.0 {
certified_min / scale
} else {
0.0
};
Ok(AtomEmbeddednessCertificate {
harmonics,
grid_min,
grid_correction,
certified_min,
scale,
relative_margin,
embedded: certified_min > 0.0,
center_nodes,
separation_nodes,
})
}
pub fn atom_decoder_embeddedness(
term: &SaeManifoldTerm,
atom_idx: usize,
) -> Result<Option<AtomEmbeddednessCertificate>, String> {
let Some(atom) = term.atoms.get(atom_idx) else {
return Err(format!(
"atom_decoder_embeddedness: atom {atom_idx} is not in the term"
));
};
if atom.latent_dim() != 1 || !matches!(atom.basis_kind(), SaeAtomBasisKind::Periodic) {
return Ok(None);
}
let decoder = atom.decoder_coefficients();
if decoder.nrows() % 2 == 0 {
return Ok(None);
}
certify_periodic_decoder_embeddedness(decoder.view()).map(Some)
}