use super::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct GraphEdge {
pub a: usize,
pub b: usize,
}
impl GraphEdge {
pub fn new(a: usize, b: usize) -> Result<Self, String> {
if a == b {
return Err(format!("GraphEdge cannot join vertex {a} to itself"));
}
Ok(if a < b {
Self { a, b }
} else {
Self { a: b, b: a }
})
}
}
pub fn graph_edge_rank_charge(n_eff: f64, fiber_rank: usize) -> f64 {
0.5 * fiber_rank as f64 * n_eff.max(2.0).ln()
}
pub fn coactivation_mi_nats(stats: &crate::atom_codes::CoactivationStats) -> f64 {
let n = stats.n_obs as f64;
if n <= 0.0 {
return 0.0;
}
let p1x = stats.n_a as f64 / n;
let px1 = stats.n_b as f64 / n;
let p11 = stats.n_joint as f64 / n;
let p10 = (p1x - p11).max(0.0);
let p01 = (px1 - p11).max(0.0);
let p00 = (1.0 - p11 - p10 - p01).max(0.0);
let cell = |p: f64, pa: f64, pb: f64| -> f64 {
if p > 0.0 && pa > 0.0 && pb > 0.0 {
p * (p / (pa * pb)).ln()
} else {
0.0
}
};
(cell(p11, p1x, px1)
+ cell(p10, p1x, 1.0 - px1)
+ cell(p01, 1.0 - p1x, px1)
+ cell(p00, 1.0 - p1x, 1.0 - px1))
.max(0.0)
}
pub fn coactivation_edge_evidence(
stats: &crate::atom_codes::CoactivationStats,
n_eff: f64,
) -> (f64, f64) {
let mi = coactivation_mi_nats(stats);
(mi, n_eff.max(0.0) * mi)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GraphTopologyReadout {
pub vertices: usize,
pub surviving_edges: usize,
pub b0: usize,
pub b1: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum GraphCompressionKind {
Circle,
Interval,
FiniteSet,
Cylinder,
Torus,
Sphere,
Graph,
}
#[derive(Debug, Clone, PartialEq)]
pub struct GraphCompressionReport {
pub kind: GraphCompressionKind,
pub name: &'static str,
pub generic_edge_bits: f64,
pub named_bits: f64,
pub bits_saved: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct GraphStructureSelection {
pub selected: bool,
pub total_edge_delta_loss: f64,
pub total_edge_charge: f64,
pub margin: f64,
pub topology: GraphTopologyReadout,
pub occupancy: OccupancyLaw,
pub compression: GraphCompressionReport,
}
impl GraphCompressionReport {
pub fn certified(
kind: GraphCompressionKind,
name: &'static str,
generic_edge_bits: f64,
named_bits: f64,
) -> Self {
Self {
kind,
name,
generic_edge_bits,
named_bits,
bits_saved: generic_edge_bits - named_bits,
}
}
pub fn unnamed(generic_edge_bits: f64) -> Self {
Self {
kind: GraphCompressionKind::Graph,
name: "structure without a standard name",
generic_edge_bits,
named_bits: generic_edge_bits,
bits_saved: 0.0,
}
}
pub fn earns_standard_name(&self) -> bool {
self.kind != GraphCompressionKind::Graph && self.bits_saved > 0.0
}
}
#[derive(Debug, Clone)]
pub struct LearnedGraphAtom {
anchor_embeddings: Array2<f64>,
candidate_edges: Vec<GraphEdge>,
edge_precisions: Vec<f64>,
edge_delta_loss: Vec<f64>,
surviving_edges: Vec<bool>,
n_eff: f64,
occupancy: OccupancyLaw,
}
impl LearnedGraphAtom {
pub fn derived_knn_k(anchors: usize) -> usize {
anchors.saturating_sub(1).min(2)
}
pub fn knn_candidate_edges(
anchor_embeddings: ArrayView2<'_, f64>,
) -> Result<Vec<GraphEdge>, String> {
validate_anchor_embeddings(anchor_embeddings)?;
let anchors = anchor_embeddings.nrows();
let k = Self::derived_knn_k(anchors);
if k == 0 {
return Ok(Vec::new());
}
let mut edges = Vec::new();
for a in 0..anchors {
let mut distances = Vec::<(f64, usize)>::with_capacity(anchors.saturating_sub(1));
for b in 0..anchors {
if a == b {
continue;
}
let mut dist2 = 0.0_f64;
for c in 0..anchor_embeddings.ncols() {
let d = anchor_embeddings[[a, c]] - anchor_embeddings[[b, c]];
dist2 += d * d;
}
distances.push((dist2, b));
}
distances.sort_by(|left, right| {
left.0
.partial_cmp(&right.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| left.1.cmp(&right.1))
});
for &(_, b) in distances.iter().take(k) {
let edge = GraphEdge::new(a, b)?;
if !edges.contains(&edge) {
edges.push(edge);
}
}
}
edges.sort_by_key(|edge| (edge.a, edge.b));
Ok(edges)
}
pub fn from_reml_knn_edges(
anchor_embeddings: ArrayView2<'_, f64>,
row_coordinates: &[f64],
n_eff: f64,
edge_precisions: &[f64],
edge_delta_loss: &[f64],
) -> Result<Self, String> {
let candidate_edges = Self::knn_candidate_edges(anchor_embeddings)?;
Self::from_reml_candidate_edges(
anchor_embeddings,
row_coordinates,
n_eff,
&candidate_edges,
edge_precisions,
edge_delta_loss,
)
}
pub fn from_reml_candidate_edges(
anchor_embeddings: ArrayView2<'_, f64>,
row_coordinates: &[f64],
n_eff: f64,
candidate_edges: &[GraphEdge],
edge_precisions: &[f64],
edge_delta_loss: &[f64],
) -> Result<Self, String> {
validate_anchor_embeddings(anchor_embeddings)?;
let anchors = anchor_embeddings.nrows();
let fiber_rank = anchor_embeddings.ncols();
if candidate_edges.is_empty() {
return Err("LearnedGraphAtom requires at least one candidate edge".to_string());
}
if edge_precisions.len() != candidate_edges.len() {
return Err(format!(
"LearnedGraphAtom edge_precisions length {} must equal candidate edges {}",
edge_precisions.len(),
candidate_edges.len()
));
}
if edge_delta_loss.len() != candidate_edges.len() {
return Err(format!(
"LearnedGraphAtom edge_delta_loss length {} must equal candidate edges {}",
edge_delta_loss.len(),
candidate_edges.len()
));
}
if !(n_eff.is_finite() && n_eff > 0.0) {
return Err(format!(
"LearnedGraphAtom n_eff must be finite and positive; got {n_eff}"
));
}
let mut normalized = Vec::<GraphEdge>::with_capacity(candidate_edges.len());
for (idx, edge) in candidate_edges.iter().enumerate() {
if edge.a >= anchors || edge.b >= anchors || edge.a == edge.b {
return Err(format!(
"LearnedGraphAtom candidate edge {idx} = ({}, {}) is invalid for {anchors} anchors",
edge.a, edge.b
));
}
let edge = GraphEdge::new(edge.a, edge.b)?;
if normalized.contains(&edge) {
return Err(format!(
"LearnedGraphAtom candidate edge {idx} duplicates ({}, {})",
edge.a, edge.b
));
}
normalized.push(edge);
}
for (edge, &precision) in edge_precisions.iter().enumerate() {
if !(precision.is_finite() && precision >= 0.0) {
return Err(format!(
"LearnedGraphAtom edge {edge} precision must be finite and nonnegative; got {precision}"
));
}
}
for (edge, &delta) in edge_delta_loss.iter().enumerate() {
if !delta.is_finite() {
return Err(format!(
"LearnedGraphAtom edge {edge} deletion loss must be finite; got {delta}"
));
}
}
let charge = graph_edge_rank_charge(n_eff, fiber_rank);
let surviving_edges = edge_precisions
.iter()
.zip(edge_delta_loss.iter())
.map(|(&precision, &delta)| precision > 0.0 && delta > charge)
.collect();
Ok(Self {
anchor_embeddings: anchor_embeddings.to_owned(),
candidate_edges: normalized,
edge_precisions: edge_precisions.to_vec(),
edge_delta_loss: edge_delta_loss.to_vec(),
surviving_edges,
n_eff,
occupancy: classify_occupancy(row_coordinates),
})
}
pub fn enroll_from_coactivation(
anchor_embeddings: ArrayView2<'_, f64>,
row_coordinates: &[f64],
n_eff: f64,
coactive_pairs: &[(usize, usize, crate::atom_codes::CoactivationStats)],
dependence_floor: f64,
) -> Result<Self, String> {
validate_anchor_embeddings(anchor_embeddings)?;
let anchors = anchor_embeddings.nrows();
let mut candidate_edges = Vec::new();
let mut edge_precisions = Vec::new();
let mut edge_delta_loss = Vec::new();
for (a, b, stats) in coactive_pairs {
if stats.dependence() < dependence_floor {
continue;
}
if *a >= anchors || *b >= anchors || a == b {
return Err(format!(
"enroll_from_coactivation: co-fire pair ({a}, {b}) is invalid for {anchors} atoms"
));
}
let edge = GraphEdge::new(*a, *b)?;
if candidate_edges.contains(&edge) {
continue;
}
let (precision, delta) = coactivation_edge_evidence(stats, n_eff);
candidate_edges.push(edge);
edge_precisions.push(precision);
edge_delta_loss.push(delta);
}
if candidate_edges.is_empty() {
return Err(format!(
"enroll_from_coactivation: no co-fire pair cleared the dependence floor {dependence_floor}"
));
}
Self::from_reml_candidate_edges(
anchor_embeddings,
row_coordinates,
n_eff,
&candidate_edges,
&edge_precisions,
&edge_delta_loss,
)
}
pub fn anchors(&self) -> usize {
self.anchor_embeddings.nrows()
}
pub fn fiber_rank(&self) -> usize {
self.anchor_embeddings.ncols()
}
pub fn n_eff(&self) -> f64 {
self.n_eff
}
pub fn one_edge_charge(&self) -> f64 {
graph_edge_rank_charge(self.n_eff, self.fiber_rank())
}
pub fn summed_edge_charge(&self) -> f64 {
self.one_edge_charge() * self.topology_readout().surviving_edges as f64
}
pub fn occupancy(&self) -> OccupancyLaw {
self.occupancy
}
pub fn candidate_edges(&self) -> &[GraphEdge] {
&self.candidate_edges
}
pub fn edge_precisions(&self) -> &[f64] {
&self.edge_precisions
}
pub fn edge_delta_loss(&self) -> &[f64] {
&self.edge_delta_loss
}
pub fn surviving_edges(&self) -> &[bool] {
&self.surviving_edges
}
pub fn surviving_degrees(&self) -> Vec<usize> {
let mut degrees = vec![0usize; self.anchors()];
for (idx, edge) in self.candidate_edges.iter().enumerate() {
if self.surviving_edges[idx] {
degrees[edge.a] += 1;
degrees[edge.b] += 1;
}
}
degrees
}
pub fn surviving_laplacian(&self) -> Array2<f64> {
self.weighted_laplacian_from_mask(&self.surviving_edges)
}
pub fn full_laplacian(&self) -> Array2<f64> {
let all_edges = vec![true; self.candidate_edges.len()];
self.weighted_laplacian_from_mask(&all_edges)
}
pub fn surviving_smoothness_value(&self) -> f64 {
self.smoothness_value_from_mask(&self.surviving_edges)
}
pub fn topology_readout(&self) -> GraphTopologyReadout {
let vertices = self.anchors();
let mut parent: Vec<usize> = (0..vertices).collect();
let mut surviving_edges = 0usize;
for (idx, edge) in self.candidate_edges.iter().enumerate() {
if self.surviving_edges[idx] {
surviving_edges += 1;
graph_union(&mut parent, edge.a, edge.b);
}
}
let mut roots = Vec::with_capacity(vertices);
for vertex in 0..vertices {
let root = graph_find(&mut parent, vertex);
if !roots.contains(&root) {
roots.push(root);
}
}
let b0 = roots.len();
let b1 = surviving_edges + b0 - vertices;
GraphTopologyReadout {
vertices,
surviving_edges,
b0,
b1,
}
}
pub fn certified_compression(&self) -> GraphCompressionReport {
let readout = self.topology_readout();
let degrees = self.surviving_degrees();
let max_edges = readout
.vertices
.saturating_mul(readout.vertices.saturating_sub(1))
/ 2;
let generic = crate::description_length::selection_bits(
max_edges as i64,
readout.surviving_edges as i64,
);
let log_vertices = (readout.vertices.max(2) as f64).log2();
let named = if readout.b0 == 1
&& readout.b1 == 1
&& degrees.iter().all(|&d| d == 2)
&& self.surviving_edge_weights_are_uniform()
{
Some((GraphCompressionKind::Circle, "circle", log_vertices))
} else if readout.b0 == 1
&& readout.b1 == 0
&& readout.vertices >= 2
&& degrees.iter().filter(|&&d| d == 1).count() == 2
&& degrees.iter().filter(|&&d| d == 2).count() == readout.vertices.saturating_sub(2)
&& self.surviving_edge_weights_are_uniform()
{
Some((
GraphCompressionKind::Interval,
"interval",
2.0 * log_vertices,
))
} else if matches!(
self.occupancy,
OccupancyLaw::Discrete { anchors } if anchors == readout.vertices
) && readout.b1 == 0
{
Some((GraphCompressionKind::FiniteSet, "finite_set", log_vertices))
} else {
None
};
if let Some((kind, name, named_bits)) = named {
let report = GraphCompressionReport::certified(kind, name, generic, named_bits);
if report.bits_saved > 0.0 {
return report;
}
}
GraphCompressionReport::unnamed(generic)
}
pub fn structure_selection(&self) -> GraphStructureSelection {
let topology = self.topology_readout();
let total_edge_delta_loss = self
.edge_delta_loss
.iter()
.zip(self.surviving_edges.iter())
.filter_map(|(&delta, &survives)| survives.then_some(delta))
.sum::<f64>();
let total_edge_charge = self.one_edge_charge() * topology.surviving_edges as f64;
let margin = total_edge_delta_loss - total_edge_charge;
GraphStructureSelection {
selected: topology.surviving_edges > 0 && margin > 0.0,
total_edge_delta_loss,
total_edge_charge,
margin,
topology,
occupancy: self.occupancy,
compression: self.certified_compression(),
}
}
pub fn surviving_penalty_op(
&self,
global_offset: usize,
beta_dim: usize,
) -> Arc<dyn gam_solve::arrow_schur::BetaPenaltyOp> {
Arc::new(IdentityRightKroneckerPenaltyOp {
factor_a: self.surviving_laplacian(),
p: self.fiber_rank(),
global_offset,
k: beta_dim,
})
}
fn weighted_laplacian_from_mask(&self, active_edges: &[bool]) -> Array2<f64> {
let anchors = self.anchors();
let mut laplacian = Array2::<f64>::zeros((anchors, anchors));
for (idx, edge) in self.candidate_edges.iter().enumerate() {
if !active_edges[idx] {
continue;
}
let w = self.edge_precisions[idx];
if w == 0.0 {
continue;
}
laplacian[[edge.a, edge.a]] += w;
laplacian[[edge.b, edge.b]] += w;
laplacian[[edge.a, edge.b]] -= w;
laplacian[[edge.b, edge.a]] -= w;
}
laplacian
}
fn smoothness_value_from_mask(&self, active_edges: &[bool]) -> f64 {
let fiber_rank = self.fiber_rank();
let mut value = 0.0_f64;
for (idx, edge) in self.candidate_edges.iter().enumerate() {
if !active_edges[idx] {
continue;
}
let w = self.edge_precisions[idx];
if w == 0.0 {
continue;
}
for channel in 0..fiber_rank {
let diff = self.anchor_embeddings[[edge.a, channel]]
- self.anchor_embeddings[[edge.b, channel]];
value += w * diff * diff;
}
}
value
}
fn surviving_edge_weights_are_uniform(&self) -> bool {
let mut min_weight = f64::INFINITY;
let mut max_weight = f64::NEG_INFINITY;
let mut count = 0usize;
for (idx, survives) in self.surviving_edges.iter().enumerate() {
if *survives {
let weight = self.edge_precisions[idx];
min_weight = min_weight.min(weight);
max_weight = max_weight.max(weight);
count += 1;
}
}
if count == 0 {
return false;
}
let scale = max_weight.abs().max(min_weight.abs()).max(1.0);
max_weight - min_weight <= f64::EPSILON * scale * count as f64
}
}
fn validate_anchor_embeddings(anchor_embeddings: ArrayView2<'_, f64>) -> Result<(), String> {
let anchors = anchor_embeddings.nrows();
let fiber_rank = anchor_embeddings.ncols();
if anchors < 2 {
return Err(format!(
"LearnedGraphAtom requires at least 2 anchors; got {anchors}"
));
}
if fiber_rank == 0 {
return Err("LearnedGraphAtom requires fiber_rank >= 1".to_string());
}
if anchor_embeddings.iter().any(|v| !v.is_finite()) {
return Err("LearnedGraphAtom anchor_embeddings contain a non-finite value".to_string());
}
Ok(())
}
fn graph_find(parent: &mut [usize], x: usize) -> usize {
let mut root = x;
while parent[root] != root {
root = parent[root];
}
let mut cur = x;
while parent[cur] != root {
let next = parent[cur];
parent[cur] = root;
cur = next;
}
root
}
fn graph_union(parent: &mut [usize], a: usize, b: usize) {
let ra = graph_find(parent, a);
let rb = graph_find(parent, b);
if ra != rb {
parent[rb] = ra;
}
}
pub const SPECTRAL_DECODE_MAX_Q: usize = 8;
pub fn spectral_decode_rank_charge(n_eff: f64, q: usize) -> f64 {
0.5 * q as f64 * n_eff.max(2.0).ln()
}
fn select_spectral_q(mu: &[f64], q_max: usize) -> usize {
let cap = q_max.min(mu.len());
if cap <= 1 {
return cap.max(1);
}
let tiny = f64::MIN_POSITIVE;
let mut best_q = 1usize;
let mut best_ratio = f64::NEG_INFINITY;
for k in 1..cap {
let ratio = mu[k] / mu[k - 1].max(tiny);
if ratio > best_ratio {
best_ratio = ratio;
best_q = k;
}
}
best_q
}
#[derive(Debug, Clone)]
pub struct GraphSpectralBasis {
basis_values: Array2<f64>,
eigenvalues: Vec<f64>,
anchor_embeddings: Array2<f64>,
bandwidth: f64,
n_eff: f64,
}
impl GraphSpectralBasis {
pub fn selected_q(&self) -> usize {
self.eigenvalues.len()
}
pub fn eigenvalues(&self) -> &[f64] {
&self.eigenvalues
}
pub fn vertex_basis(&self) -> ArrayView2<'_, f64> {
self.basis_values.view()
}
pub fn bandwidth(&self) -> f64 {
self.bandwidth
}
pub fn penalty(&self) -> Array2<f64> {
let q = self.selected_q();
let mut penalty = Array2::<f64>::zeros((q, q));
for k in 0..q {
penalty[[k, k]] = self.eigenvalues[k];
}
penalty
}
pub fn rank_charge_dof(&self) -> f64 {
spectral_decode_rank_charge(self.n_eff, self.selected_q())
}
pub fn nystrom_coordinate(&self, z: &[f64]) -> Result<(Vec<f64>, Array2<f64>), String> {
let r = self.anchor_embeddings.ncols();
if z.len() != r {
return Err(format!(
"GraphSpectralBasis::nystrom_coordinate: query has {} features but graph anchors have {r}",
z.len()
));
}
let query = ArrayView2::from_shape((1, r), z)
.map_err(|e| format!("GraphSpectralBasis::nystrom_coordinate: bad query shape: {e}"))?;
let (phi, jet) = self.nystrom_coordinates(query)?;
let q = self.selected_q();
let coord = phi.row(0).to_vec();
let mut jac = Array2::<f64>::zeros((q, r));
for k in 0..q {
for c in 0..r {
jac[[k, c]] = jet[[0, k, c]];
}
}
Ok((coord, jac))
}
pub fn nystrom_coordinates(
&self,
points: ArrayView2<'_, f64>,
) -> Result<(Array2<f64>, Array3<f64>), String> {
let anchors = self.anchor_embeddings.nrows();
let r = self.anchor_embeddings.ncols();
let q = self.selected_q();
if points.ncols() != r {
return Err(format!(
"GraphSpectralBasis::nystrom_coordinates: query has {} features but graph anchors have {r}",
points.ncols()
));
}
if points.iter().any(|v| !v.is_finite()) {
return Err(
"GraphSpectralBasis::nystrom_coordinates: query contains a non-finite value".into(),
);
}
let n = points.nrows();
let eps = self.bandwidth;
if !(eps > 0.0 && eps.is_finite()) {
return Err(format!(
"GraphSpectralBasis::nystrom_coordinates: non-positive bandwidth {eps}"
));
}
let mut phi = Array2::<f64>::zeros((n, q));
let mut jet = Array3::<f64>::zeros((n, q, r));
let mut w = vec![0.0_f64; anchors];
let mut n_k = vec![0.0_f64; q];
let mut dw = vec![0.0_f64; anchors];
for row in 0..n {
let mut s = 0.0_f64;
for i in 0..anchors {
let mut d2 = 0.0_f64;
for c in 0..r {
let d = points[[row, c]] - self.anchor_embeddings[[i, c]];
d2 += d * d;
}
let wi = (-d2 / eps).exp();
w[i] = wi;
s += wi;
}
if !(s > 0.0 && s.is_finite()) {
return Err(
"GraphSpectralBasis::nystrom_coordinates: query point has vanishing affinity \
to every anchor (bandwidth underflow)"
.into(),
);
}
for k in 0..q {
let mut acc = 0.0_f64;
for i in 0..anchors {
acc += w[i] * self.basis_values[[i, k]];
}
n_k[k] = acc;
phi[[row, k]] = acc / s;
}
for c in 0..r {
let mut ds_c = 0.0_f64;
for i in 0..anchors {
let dwi =
w[i] * (-2.0 * (points[[row, c]] - self.anchor_embeddings[[i, c]]) / eps);
dw[i] = dwi;
ds_c += dwi;
}
for k in 0..q {
let mut dn_kc = 0.0_f64;
for i in 0..anchors {
dn_kc += dw[i] * self.basis_values[[i, k]];
}
jet[[row, k, c]] = (dn_kc * s - n_k[k] * ds_c) / (s * s);
}
}
}
Ok((phi, jet))
}
pub fn evaluator(&self) -> Arc<dyn SaeBasisEvaluator> {
Arc::new(NystromSpectralEvaluator {
basis: self.clone(),
})
}
}
#[derive(Debug, Clone)]
pub struct NystromSpectralEvaluator {
basis: GraphSpectralBasis,
}
impl SaeBasisEvaluator for NystromSpectralEvaluator {
fn evaluate(&self, coords: ArrayView2<'_, f64>) -> Result<(Array2<f64>, Array3<f64>), String> {
self.basis.nystrom_coordinates(coords)
}
fn second_jet_dyn(&self, coords: ArrayView2<'_, f64>) -> Option<Result<Array4<f64>, String>> {
let r = self.basis.anchor_embeddings.ncols();
if coords.ncols() != r {
return Some(Err(format!(
"NystromSpectralEvaluator::second_jet_dyn: query has {} features but graph anchors have {r}",
coords.ncols()
)));
}
None
}
fn third_jet_dyn(
&self,
coords: ArrayView2<'_, f64>,
) -> Option<Result<ndarray::Array5<f64>, String>> {
let r = self.basis.anchor_embeddings.ncols();
if coords.ncols() != r {
return Some(Err(format!(
"NystromSpectralEvaluator::third_jet_dyn: query has {} features but graph anchors have {r}",
coords.ncols()
)));
}
None
}
}
pub struct SpectralGraphRaceCandidate {
pub basis_kind: SaeAtomBasisKind,
pub manifold: LatentManifold,
pub latent_dim: usize,
pub row_coords: Array2<f64>,
pub phi: Array2<f64>,
pub jet: Array3<f64>,
pub decoder: Array2<f64>,
pub penalty: Array2<f64>,
pub evaluator: Arc<dyn SaeBasisEvaluator>,
pub rank_charge_dof: f64,
}
impl LearnedGraphAtom {
fn surviving_edge_bandwidth(&self) -> Result<f64, String> {
let fiber_rank = self.fiber_rank();
let mut lengths = Vec::new();
for (idx, edge) in self.candidate_edges.iter().enumerate() {
if !self.surviving_edges[idx] {
continue;
}
let mut d2 = 0.0_f64;
for c in 0..fiber_rank {
let d = self.anchor_embeddings[[edge.a, c]] - self.anchor_embeddings[[edge.b, c]];
d2 += d * d;
}
lengths.push(d2);
}
if lengths.is_empty() {
return Err(
"LearnedGraphAtom::spectral_decode_basis: no surviving edge to set the Nyström \
bandwidth"
.into(),
);
}
lengths.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let median = lengths[lengths.len() / 2];
if !(median > 0.0 && median.is_finite()) {
return Err(format!(
"LearnedGraphAtom::spectral_decode_basis: degenerate surviving-edge bandwidth {median}"
));
}
Ok(median)
}
pub fn spectral_decode_basis(&self) -> Result<GraphSpectralBasis, String> {
let anchors = self.anchors();
let laplacian = self.surviving_laplacian();
let (evals, evecs) = laplacian.eigh(Side::Lower).map_err(|e| {
format!(
"LearnedGraphAtom::spectral_decode_basis: Laplacian eigendecomposition failed: {e}"
)
})?;
let mut order: Vec<usize> = (0..evals.len()).collect();
order.sort_by(|&a, &b| {
evals[a]
.partial_cmp(&evals[b])
.unwrap_or(std::cmp::Ordering::Equal)
});
let lambda_max = order.last().map(|&i| evals[i]).unwrap_or(0.0).max(0.0);
let zero_tol = (lambda_max * 1e-8).max(1e-12);
let nontrivial: Vec<usize> = order
.iter()
.copied()
.filter(|&i| evals[i] > zero_tol)
.collect();
if nontrivial.is_empty() {
return Err(
"LearnedGraphAtom::spectral_decode_basis: survived graph has no non-trivial \
Laplacian mode to decode (all vertices isolated or a single constant component)"
.into(),
);
}
let sorted_mu: Vec<f64> = nontrivial.iter().map(|&i| evals[i]).collect();
let q = select_spectral_q(&sorted_mu, SPECTRAL_DECODE_MAX_Q);
let mut basis_values = Array2::<f64>::zeros((anchors, q));
for (col, &eig_idx) in nontrivial.iter().take(q).enumerate() {
for v in 0..anchors {
basis_values[[v, col]] = evecs[[v, eig_idx]];
}
}
let eigenvalues: Vec<f64> = sorted_mu.iter().take(q).copied().collect();
let bandwidth = self.surviving_edge_bandwidth()?;
Ok(GraphSpectralBasis {
basis_values,
eigenvalues,
anchor_embeddings: self.anchor_embeddings.clone(),
bandwidth,
n_eff: self.n_eff,
})
}
pub fn spectral_race_candidate(
&self,
target: ArrayView2<'_, f64>,
row_embeddings: ArrayView2<'_, f64>,
) -> Result<SpectralGraphRaceCandidate, String> {
let basis = self.spectral_decode_basis()?;
let (phi, jet) = basis.nystrom_coordinates(row_embeddings)?;
let n = target.nrows();
if phi.nrows() != n {
return Err(format!(
"LearnedGraphAtom::spectral_race_candidate: {n} targets but {} row embeddings",
phi.nrows()
));
}
let penalty = basis.penalty();
let reml = gam_solve::gaussian_reml::gaussian_reml_multi_closed_form(
phi.view(),
target,
penalty.view(),
None,
None,
)
.map_err(|e| {
format!("LearnedGraphAtom::spectral_race_candidate: REML decode fit: {e:?}")
})?;
Ok(SpectralGraphRaceCandidate {
basis_kind: SaeAtomBasisKind::Precomputed("spectral_graph".to_string()),
manifold: LatentManifold::Euclidean,
latent_dim: row_embeddings.ncols(),
row_coords: row_embeddings.to_owned(),
phi,
jet,
decoder: reml.coefficients,
penalty,
evaluator: basis.evaluator(),
rank_charge_dof: basis.rank_charge_dof(),
})
}
}