use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use crate::manifold::SaeManifoldTerm;
use gam_linalg::faer_ndarray::{FaerSvd, fast_ab, fast_abt, fast_atb};
pub fn canonical_output_subspace_rows(raw: ArrayView2<'_, f64>) -> Result<Array2<f64>, String> {
let (r, p) = raw.dim();
if r == 0 || p == 0 {
return Err("canonical_output_subspace_rows: frame must be non-empty".to_string());
}
if r > p {
return Err(format!(
"canonical_output_subspace_rows: rank r={r} cannot exceed output dimension p={p}"
));
}
if raw.iter().any(|value| !value.is_finite()) {
return Err("canonical_output_subspace_rows: frame must be finite".to_string());
}
let frame = GrassmannFrame::polar_update(raw.t())?;
let singular_values = frame.gauge_singular_values();
let rank_resolution = f64::EPSILON * r.max(p) as f64 * singular_values[0];
if singular_values[r - 1] <= rank_resolution {
return Err(format!(
"canonical_output_subspace_rows: requested {r}-row frame is rank deficient \
(sigma_min={}, numerical resolution={rank_resolution})",
singular_values[r - 1]
));
}
Ok(frame.frame().t().to_owned())
}
#[derive(Clone, Debug)]
pub(crate) struct FrameProjection {
pub(crate) p: usize,
pub(crate) beta_offsets: Vec<usize>,
pub(crate) border_offsets: Vec<usize>,
pub(crate) basis_sizes: Vec<usize>,
pub(crate) ranks: Vec<usize>,
frames: Vec<Option<Array2<f64>>>,
}
impl FrameProjection {
pub(crate) fn new(term: &SaeManifoldTerm) -> Self {
Self {
p: term.output_dim(),
beta_offsets: term.beta_offsets(),
border_offsets: term.factored_border_offsets(),
basis_sizes: term.atoms.iter().map(|atom| atom.basis_size()).collect(),
ranks: term
.atoms
.iter()
.map(|atom| atom.border_frame_rank())
.collect(),
frames: term
.atoms
.iter()
.map(|atom| {
atom.decoder_frame
.as_ref()
.map(|frame| frame.frame().to_owned())
})
.collect(),
}
}
pub(crate) fn beta_dim(&self) -> usize {
self.basis_sizes.iter().sum::<usize>() * self.p
}
pub(crate) fn frames_owned(&self) -> Vec<Option<Array2<f64>>> {
self.frames.clone()
}
pub(crate) fn border_dim(&self) -> usize {
self.basis_sizes
.iter()
.zip(&self.ranks)
.map(|(m, r)| m * r)
.sum()
}
pub(crate) fn lift_border_vec(&self, border: ArrayView1<'_, f64>) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(self.beta_dim());
for atom in 0..self.basis_sizes.len() {
self.lift_atom_vec_into(atom, border, out.view_mut());
}
out
}
pub(crate) fn project_border_vec(&self, beta: ArrayView1<'_, f64>) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(self.border_dim());
for atom in 0..self.basis_sizes.len() {
self.project_atom_vec_into(atom, beta, out.view_mut(), 1.0);
}
out
}
pub(crate) fn lift_block(&self, atom: usize, block: ArrayView2<'_, f64>) -> Array2<f64> {
let m = self.basis_sizes[atom];
let r = self.ranks[atom];
if self.frames[atom].is_none() {
return block.to_owned();
}
let uk = self.frames[atom].as_ref().expect("framed atom has a frame");
let mut out = Array2::<f64>::zeros((m * self.p, m * self.p));
for b1 in 0..m {
for b2 in 0..m {
for c1 in 0..self.p {
for c2 in 0..self.p {
let mut acc = 0.0;
for j1 in 0..r {
for j2 in 0..r {
acc +=
uk[[c1, j1]] * block[[b1 * r + j1, b2 * r + j2]] * uk[[c2, j2]];
}
}
out[[b1 * self.p + c1, b2 * self.p + c2]] = acc;
}
}
}
}
out
}
pub(crate) fn project_block(&self, hbb: ArrayView2<'_, f64>) -> Array2<f64> {
let t = self.project_rows(hbb);
let mut out = Array2::<f64>::zeros((self.border_dim(), self.border_dim()));
for atom in 0..self.basis_sizes.len() {
self.project_block_left_atom(atom, t.view(), out.view_mut());
}
out
}
pub(crate) fn project_rows(&self, block: ArrayView2<'_, f64>) -> Array2<f64> {
let mut out = Array2::<f64>::zeros((block.nrows(), self.border_dim()));
for row in 0..block.nrows() {
let projected = self.project_border_vec(block.row(row));
out.row_mut(row).assign(&projected);
}
out
}
pub(crate) fn atom_border_range(&self, atom: usize) -> std::ops::Range<usize> {
let start = self.border_offsets[atom];
start..start + self.basis_sizes[atom] * self.ranks[atom]
}
pub(crate) fn lift_axis_into(
&self,
out: &mut Array1<f64>,
atom: usize,
basis_col: usize,
frame_col: usize,
) {
let base = self.beta_offsets[atom] + basis_col * self.p;
match &self.frames[atom] {
None => out[base + frame_col] = 1.0,
Some(uk) => {
for out_col in 0..self.p {
out[base + out_col] = uk[[out_col, frame_col]];
}
}
}
}
pub(crate) fn lift_local_axis_into(
&self,
out: &mut Array1<f64>,
atom: usize,
basis_col: usize,
frame_col: usize,
) {
let base = basis_col * self.p;
match &self.frames[atom] {
None => out[base + frame_col] = 1.0,
Some(uk) => {
for out_col in 0..self.p {
out[base + out_col] = uk[[out_col, frame_col]];
}
}
}
}
pub(crate) fn project_atom_vec_into(
&self,
atom: usize,
beta: ArrayView1<'_, f64>,
mut out: ndarray::ArrayViewMut1<'_, f64>,
scale: f64,
) {
let m = self.basis_sizes[atom];
let r = self.ranks[atom];
let ob = self.beta_offsets[atom];
let oc = self.border_offsets[atom];
for basis_col in 0..m {
let base_b = ob + basis_col * self.p;
let base_c = oc + basis_col * r;
match &self.frames[atom] {
None => {
for j in 0..r {
out[base_c + j] += scale * beta[base_b + j];
}
}
Some(uk) => {
for j in 0..r {
let mut acc = 0.0;
for i in 0..self.p {
acc += uk[[i, j]] * beta[base_b + i];
}
out[base_c + j] += scale * acc;
}
}
}
}
}
pub(crate) fn project_atom_block(&self, atom: usize, block: &[f64]) -> Vec<f64> {
let m = self.basis_sizes[atom];
let r = self.ranks[atom];
let mut out = vec![0.0_f64; m * r];
for basis_col in 0..m {
let base_b = basis_col * self.p;
let base_c = basis_col * r;
match &self.frames[atom] {
None => {
out[base_c..base_c + r].copy_from_slice(&block[base_b..base_b + r]);
}
Some(uk) => {
for j in 0..r {
let mut acc = 0.0;
for i in 0..self.p {
acc += uk[[i, j]] * block[base_b + i];
}
out[base_c + j] = acc;
}
}
}
}
out
}
pub(crate) fn project_local_atom_vec_into(
&self,
atom: usize,
beta: ArrayView1<'_, f64>,
out: ndarray::ArrayViewMut1<'_, f64>,
scale: f64,
) {
self.project_atom_vec_into_with_base(atom, beta, out, scale, 0);
}
pub(crate) fn project_atom_vec_into_with_base(
&self,
atom: usize,
beta: ArrayView1<'_, f64>,
mut out: ndarray::ArrayViewMut1<'_, f64>,
scale: f64,
beta_base_offset: usize,
) {
let m = self.basis_sizes[atom];
let r = self.ranks[atom];
let oc = self.border_offsets[atom];
for basis_col in 0..m {
let base_b = beta_base_offset + basis_col * self.p;
let base_c = oc + basis_col * r;
match &self.frames[atom] {
None => {
for j in 0..r {
out[base_c + j] += scale * beta[base_b + j];
}
}
Some(uk) => {
for j in 0..r {
let mut acc = 0.0;
for i in 0..self.p {
acc += uk[[i, j]] * beta[base_b + i];
}
out[base_c + j] += scale * acc;
}
}
}
}
}
pub(crate) fn lift_atom_vec_into(
&self,
atom: usize,
border: ArrayView1<'_, f64>,
mut out: ndarray::ArrayViewMut1<'_, f64>,
) {
let m = self.basis_sizes[atom];
let r = self.ranks[atom];
let ob = self.beta_offsets[atom];
let oc = self.border_offsets[atom];
for basis_col in 0..m {
let base_b = ob + basis_col * self.p;
let base_c = oc + basis_col * r;
match &self.frames[atom] {
None => {
for i in 0..self.p {
out[base_b + i] = border[base_c + i];
}
}
Some(uk) => {
for i in 0..self.p {
let mut acc = 0.0;
for j in 0..r {
acc += uk[[i, j]] * border[base_c + j];
}
out[base_b + i] = acc;
}
}
}
}
}
pub(crate) fn accumulate_output_project(
&self,
atom: usize,
c_base: usize,
output: usize,
value: f64,
out: &mut [f64],
) {
match &self.frames[atom] {
None => out[c_base + output] += value,
Some(uk) => {
let rank = self.ranks[atom];
let frame_row = uk.row(output);
let frame_slice = frame_row.as_slice().expect("frame rows are contiguous");
let out_slice = &mut out[c_base..c_base + rank];
for (slot, &u) in out_slice.iter_mut().zip(frame_slice.iter()) {
*slot += value * u;
}
}
}
}
pub(crate) fn project_jacobian_rows(
&self,
atom: usize,
jac: ArrayView2<'_, f64>,
) -> Option<Array2<f64>> {
self.frames[atom].as_ref().map(|uk| fast_ab(&jac, uk))
}
pub(crate) fn output_variance(
&self,
atom: usize,
cov_c: ArrayView2<'_, f64>,
basis: ArrayView1<'_, f64>,
output: usize,
) -> f64 {
let Some(uk) = &self.frames[atom] else {
return self.full_output_variance(atom, cov_c, basis, output);
};
let m = self.basis_sizes[atom];
let r = self.ranks[atom];
let mut var = 0.0;
for b1 in 0..m {
let phi1 = basis[b1];
if phi1 == 0.0 {
continue;
}
for b2 in 0..m {
let phi2 = basis[b2];
if phi2 == 0.0 {
continue;
}
for j1 in 0..r {
for j2 in 0..r {
var += phi1
* phi2
* uk[[output, j1]]
* cov_c[[b1 * r + j1, b2 * r + j2]]
* uk[[output, j2]];
}
}
}
}
var
}
pub(crate) fn full_output_variance(
&self,
atom: usize,
cov: ArrayView2<'_, f64>,
basis: ArrayView1<'_, f64>,
output: usize,
) -> f64 {
let m = self.basis_sizes[atom];
let mut var = 0.0;
for b1 in 0..m {
let phi1 = basis[b1];
if phi1 == 0.0 {
continue;
}
for b2 in 0..m {
var += phi1 * basis[b2] * cov[[b1 * self.p + output, b2 * self.p + output]];
}
}
var
}
pub(crate) fn project_block_left_atom(
&self,
atom: usize,
t: ArrayView2<'_, f64>,
mut out: ndarray::ArrayViewMut2<'_, f64>,
) {
let m = self.basis_sizes[atom];
let r = self.ranks[atom];
let ob = self.beta_offsets[atom];
let oc = self.border_offsets[atom];
for basis_col in 0..m {
let base_b = ob + basis_col * self.p;
let base_c = oc + basis_col * r;
match &self.frames[atom] {
None => {
for j in 0..r {
for c in 0..out.ncols() {
out[[base_c + j, c]] += t[[base_b + j, c]];
}
}
}
Some(uk) => {
for j in 0..r {
for c in 0..out.ncols() {
let mut acc = 0.0;
for i in 0..self.p {
acc += uk[[i, j]] * t[[base_b + i, c]];
}
out[[base_c + j, c]] += acc;
}
}
}
}
}
}
}
pub(crate) struct FramedDeviceArgs<'a> {
pub p: usize,
pub border_dim: usize,
pub border_offsets: &'a [usize],
pub ranks: &'a [usize],
pub basis_sizes: &'a [usize],
pub smooth_scaled_s: &'a [Array2<f64>],
pub frame_blocks: Vec<gam_solve::arrow_schur::FactoredFrameGBlock>,
pub rows: &'a [gam_solve::arrow_schur::ArrowRowBlock],
}
pub(crate) fn build_framed_device_sae_data(
args: FramedDeviceArgs<'_>,
) -> gam_solve::arrow_schur::DeviceSaePcgData {
use gam_solve::arrow_schur::{DeviceSaeFrameData, DeviceSaePcgData, DeviceSaeSmoothBlock};
let FramedDeviceArgs {
p,
border_dim,
border_offsets,
ranks,
basis_sizes,
smooth_scaled_s,
frame_blocks,
rows,
} = args;
let n_atoms = ranks.len();
let mut smooth_blocks = Vec::with_capacity(n_atoms);
let mut smooth_ranks = Vec::with_capacity(n_atoms);
for k in 0..n_atoms {
smooth_blocks.push(DeviceSaeSmoothBlock {
global_offset: border_offsets[k],
factor_a: smooth_scaled_s[k].clone(),
});
smooth_ranks.push(ranks[k]);
}
let row_htbeta: Vec<Vec<f64>> = rows
.iter()
.map(|row| {
let (qi, w) = row.htbeta.dim();
if w != border_dim {
return Vec::new();
}
let mut flat = vec![0.0_f64; qi * w];
for c in 0..qi {
for a in 0..w {
flat[c * w + a] = row.htbeta[[c, a]];
}
}
flat
})
.collect();
DeviceSaePcgData {
p,
beta_dim: border_dim,
a_phi: std::sync::Arc::from(Vec::new().into_boxed_slice()),
local_jac: std::sync::Arc::from(Vec::new().into_boxed_slice()),
smooth_blocks,
sparse_g_blocks: Vec::new(),
frame: Some(DeviceSaeFrameData {
ranks: ranks.to_vec(),
basis_sizes: basis_sizes.to_vec(),
border_offsets: border_offsets.to_vec(),
frame_blocks,
smooth_ranks,
row_htbeta,
}),
}
}
pub(crate) fn build_framed_device_sae_data_reusing(
args: FramedDeviceArgs<'_>,
recycled: Option<std::sync::Arc<gam_solve::arrow_schur::DeviceSaePcgData>>,
) -> std::sync::Arc<gam_solve::arrow_schur::DeviceSaePcgData> {
if let Some(mut allocation) = recycled {
if let Some(data) = std::sync::Arc::get_mut(&mut allocation) {
if data.frame.is_some() {
refresh_framed_device_sae_data(data, args);
return allocation;
}
}
return std::sync::Arc::new(build_framed_device_sae_data(args));
}
std::sync::Arc::new(build_framed_device_sae_data(args))
}
fn refresh_framed_device_sae_data(
data: &mut gam_solve::arrow_schur::DeviceSaePcgData,
args: FramedDeviceArgs<'_>,
) {
let FramedDeviceArgs {
p,
border_dim,
border_offsets,
ranks,
basis_sizes,
smooth_scaled_s,
frame_blocks,
rows,
} = args;
data.p = p;
data.beta_dim = border_dim;
if !data.a_phi.is_empty() {
data.a_phi = std::sync::Arc::from(Vec::new().into_boxed_slice());
}
if !data.local_jac.is_empty() {
data.local_jac = std::sync::Arc::from(Vec::new().into_boxed_slice());
}
if data.smooth_blocks.len() == smooth_scaled_s.len() {
for (atom_idx, (block, source)) in data
.smooth_blocks
.iter_mut()
.zip(smooth_scaled_s)
.enumerate()
{
block.global_offset = border_offsets[atom_idx];
block.factor_a.clone_from(source);
}
} else {
data.smooth_blocks = smooth_scaled_s
.iter()
.enumerate()
.map(
|(atom_idx, source)| gam_solve::arrow_schur::DeviceSaeSmoothBlock {
global_offset: border_offsets[atom_idx],
factor_a: source.clone(),
},
)
.collect();
}
data.sparse_g_blocks.clear();
let frame = data
.frame
.as_mut()
.expect("framed descriptor checked before refresh");
frame.ranks.clear();
frame.ranks.extend_from_slice(ranks);
frame.basis_sizes.clear();
frame.basis_sizes.extend_from_slice(basis_sizes);
frame.border_offsets.clear();
frame.border_offsets.extend_from_slice(border_offsets);
let stable_frame_block_layout = frame.frame_blocks.len() == frame_blocks.len()
&& frame
.frame_blocks
.iter()
.zip(&frame_blocks)
.all(|(current, source)| {
current.g.dim() == source.g.dim() && current.w.dim() == source.w.dim()
});
if stable_frame_block_layout {
for (current, source) in frame.frame_blocks.iter_mut().zip(frame_blocks) {
current.atom_i = source.atom_i;
current.atom_j = source.atom_j;
current.g.assign(&source.g);
current.w.assign(&source.w);
}
} else {
frame.frame_blocks = frame_blocks;
}
frame.smooth_ranks.clear();
frame.smooth_ranks.extend_from_slice(ranks);
frame.row_htbeta.resize_with(rows.len(), Vec::new);
for (target, row) in frame.row_htbeta.iter_mut().zip(rows) {
let (qi, width) = row.htbeta.dim();
if width != border_dim {
target.clear();
continue;
}
target.resize(qi * width, 0.0);
if let Some(source) = row.htbeta.as_slice() {
target.copy_from_slice(source);
} else {
for c in 0..qi {
for a in 0..width {
target[c * width + a] = row.htbeta[[c, a]];
}
}
}
}
}
pub(crate) const SAE_FRAME_RANK_CUTOFF: f64 = 1.0e-7;
pub(crate) const SAE_FRAME_MIN_AUTO_OUTPUT_DIM: usize = 12;
pub(crate) const SAE_FRAME_ACTIVATION_MARGIN: f64 = 0.25;
#[derive(Debug, Clone)]
pub struct GrassmannFrame {
frame: Array2<f64>,
gauge_singular_values: Array1<f64>,
}
impl GrassmannFrame {
pub fn output_dim(&self) -> usize {
self.frame.nrows()
}
pub fn rank(&self) -> usize {
self.frame.ncols()
}
pub fn gauge_singular_values(&self) -> &Array1<f64> {
&self.gauge_singular_values
}
pub fn frame(&self) -> ArrayView2<'_, f64> {
self.frame.view()
}
pub fn from_orthonormal(
frame: Array2<f64>,
gauge_singular_values: Array1<f64>,
) -> Result<Self, String> {
let (p, r) = frame.dim();
if p == 0 || r == 0 {
return Err("GrassmannFrame::from_orthonormal: frame must be non-empty".to_string());
}
if r > p {
return Err(format!(
"GrassmannFrame::from_orthonormal: frame rank r={r} cannot exceed output dim p={p}"
));
}
if gauge_singular_values.len() != r {
return Err(format!(
"GrassmannFrame::from_orthonormal: gauge length {} must equal rank {r}",
gauge_singular_values.len()
));
}
for i in 0..r {
let value = gauge_singular_values[i];
if !(value.is_finite() && value >= 0.0) {
return Err(format!(
"GrassmannFrame::from_orthonormal: gauge value {i} must be finite and non-negative, got {value}"
));
}
if i > 0 && gauge_singular_values[i - 1] < value {
return Err(
"GrassmannFrame::from_orthonormal: gauge values must be descending".to_string(),
);
}
}
let tol = 1.0e-8_f64;
for a in 0..r {
for b in a..r {
let mut dot = 0.0_f64;
for row in 0..p {
dot += frame[[row, a]] * frame[[row, b]];
}
let target = if a == b { 1.0 } else { 0.0 };
if (dot - target).abs() > tol {
return Err(format!(
"GrassmannFrame::from_orthonormal: frame columns are not orthonormal at ({a}, {b}); dot={dot}"
));
}
}
}
Ok(Self::from_oriented(frame, gauge_singular_values))
}
pub fn manifold_dimension(&self) -> usize {
let r = self.rank();
let p = self.output_dim();
r * (p - r)
}
pub(crate) fn from_oriented(
mut frame: Array2<f64>,
gauge_singular_values: Array1<f64>,
) -> Self {
let (p, r) = frame.dim();
for col in 0..r {
let mut pivot_abs = 0.0_f64;
let mut pivot_val = 0.0_f64;
for row in 0..p {
let v = frame[[row, col]];
if v.abs() > pivot_abs {
pivot_abs = v.abs();
pivot_val = v;
}
}
if pivot_val < 0.0 {
for row in 0..p {
frame[[row, col]] = -frame[[row, col]];
}
}
}
Self {
frame,
gauge_singular_values,
}
}
pub fn polar_update(cross_moment: ArrayView2<'_, f64>) -> Result<Self, String> {
let (p, r) = cross_moment.dim();
if p == 0 || r == 0 {
return Err("GrassmannFrame::polar_update: cross-moment must be non-empty".into());
}
if r > p {
return Err(format!(
"GrassmannFrame::polar_update: frame rank r={r} cannot exceed output dim p={p}"
));
}
let owned = cross_moment.to_owned();
let (u_opt, sv, vt_opt) = owned
.svd(true, true)
.map_err(|e| format!("GrassmannFrame::polar_update: SVD failed: {e}"))?;
let w = u_opt.ok_or_else(|| {
"GrassmannFrame::polar_update: thin SVD returned no left factor".to_string()
})?;
let vt = vt_opt.ok_or_else(|| {
"GrassmannFrame::polar_update: thin SVD returned no right factor".to_string()
})?;
let polar = fast_ab(&w, &vt);
Ok(Self::from_oriented(polar, sv))
}
pub fn reconstruct_decoder(&self, coords: ArrayView2<'_, f64>) -> Result<Array2<f64>, String> {
if coords.ncols() != self.rank() {
return Err(format!(
"GrassmannFrame::reconstruct_decoder: coord cols {} must equal frame rank {}",
coords.ncols(),
self.rank()
));
}
Ok(fast_abt(&coords.to_owned(), &self.frame))
}
pub fn project_decoder(&self, decoder: ArrayView2<'_, f64>) -> Result<Array2<f64>, String> {
if decoder.ncols() != self.output_dim() {
return Err(format!(
"GrassmannFrame::project_decoder: decoder cols {} must equal output dim {}",
decoder.ncols(),
self.output_dim()
));
}
Ok(fast_ab(&decoder.to_owned(), &self.frame))
}
pub fn max_principal_angle(&self, other: ArrayView2<'_, f64>) -> Result<f64, String> {
if other.nrows() != self.output_dim() {
return Err(format!(
"GrassmannFrame::max_principal_angle: other rows {} must equal output dim {}",
other.nrows(),
self.output_dim()
));
}
if other.ncols() != self.rank() {
return Ok(std::f64::consts::FRAC_PI_2);
}
let other_owned = other.to_owned();
let overlap = fast_atb(&self.frame, &other_owned);
let (_u, sv_cos, _vt) = overlap
.svd(false, false)
.map_err(|e| format!("GrassmannFrame::max_principal_angle: cos-SVD failed: {e}"))?;
let u_overlap = fast_ab(&self.frame, &overlap);
let v_perp = &other_owned - &u_overlap;
let (_u, sv_sin, _vt) = v_perp
.svd(false, false)
.map_err(|e| format!("GrassmannFrame::max_principal_angle: sin-SVD failed: {e}"))?;
let min_cos = sv_cos
.iter()
.copied()
.fold(1.0_f64, f64::min)
.clamp(0.0, 1.0);
let max_sin = sv_sin
.iter()
.copied()
.fold(0.0_f64, f64::max)
.clamp(0.0, 1.0);
Ok(max_sin.atan2(min_cos))
}
pub fn from_decoder_row_space(decoder: ArrayView2<'_, f64>) -> Option<Self> {
let (m, p) = decoder.dim();
if m == 0 || p == 0 {
return None;
}
let (_u, sv, vt_opt) = decoder.to_owned().svd(false, true).ok()?;
let vt = vt_opt?;
let sigma_max = sv.iter().copied().fold(0.0_f64, f64::max);
if !(sigma_max.is_finite() && sigma_max > 0.0) {
return None;
}
let cutoff = SAE_FRAME_RANK_CUTOFF * sigma_max;
let r = sv.iter().take(vt.nrows()).filter(|&&s| s >= cutoff).count();
if r == 0 {
return None;
}
let mut frame = Array2::<f64>::zeros((p, r));
for j in 0..r {
for row in 0..p {
frame[[row, j]] = vt[[j, row]];
}
}
let gauge = Array1::from_iter(sv.iter().take(r).copied());
Some(Self::from_oriented(frame, gauge))
}
}
#[derive(Debug, Clone)]
pub struct GrassmannCrossMoment {
moment: Array2<f64>,
}
impl GrassmannCrossMoment {
pub fn new(output_dim: usize, rank: usize) -> Self {
Self {
moment: Array2::<f64>::zeros((output_dim, rank)),
}
}
pub fn accumulate(
&mut self,
targets: ArrayView2<'_, f64>,
coords: ArrayView2<'_, f64>,
) -> Result<(), String> {
if targets.ncols() != self.moment.nrows() || coords.ncols() != self.moment.ncols() {
return Err(format!(
"GrassmannCrossMoment::accumulate: expected targets (·,{}) and coords (·,{}); \
got (·,{}) and (·,{})",
self.moment.nrows(),
self.moment.ncols(),
targets.ncols(),
coords.ncols()
));
}
if targets.nrows() != coords.nrows() {
return Err(format!(
"GrassmannCrossMoment::accumulate: targets rows {} must equal coords rows {}",
targets.nrows(),
coords.nrows()
));
}
let block = fast_atb(&targets.to_owned(), &coords.to_owned());
self.moment += █
Ok(())
}
pub fn moment(&self) -> ArrayView2<'_, f64> {
self.moment.view()
}
pub fn polar_frame(&self) -> Result<GrassmannFrame, String> {
GrassmannFrame::polar_update(self.moment.view())
}
}
pub fn grassmann_recover_planted_span_angle(
targets: ArrayView2<'_, f64>,
coords: ArrayView2<'_, f64>,
planted: ArrayView2<'_, f64>,
) -> Result<f64, String> {
let p = targets.ncols();
let r = coords.ncols();
if planted.dim() != (p, r) {
return Err(format!(
"grassmann_recover_planted_span_angle: planted frame must be ({p}, {r}); got {:?}",
planted.dim()
));
}
let mut cross = GrassmannCrossMoment::new(p, r);
cross.accumulate(targets, coords)?;
let frame = cross.polar_frame()?;
frame.max_principal_angle(planted)
}
pub fn grassmann_assert_border_dim_invariant(term: &SaeManifoldTerm) -> Result<(), String> {
let expected: usize = term
.atoms
.iter()
.map(|a| a.basis_size() * a.border_frame_rank())
.sum();
let got = term.factored_border_dim();
if got != expected {
return Err(format!(
"grassmann border-dim invariant violated: factored_border_dim() = {got}, \
expected Σ M_k·r_k = {expected}"
));
}
Ok(())
}