use ndarray::{Array1, Array2, ArrayView2, Axis};
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
use gam_geometry::constant_curvature::{ConstantCurvature, distance_kappa_jet};
use super::{
ActivePenalty, BasisBuildResult, BasisError, BasisMetadata, BasisPsiDerivativeBundle,
BasisPsiDerivativeResult, BasisPsiSecondDerivativeResult, CenterStrategy, CenterStrategyKind,
ConstructiveQuadratic, PenaltyCandidate, PenaltySource, center_strategy_kind,
filter_penalty_candidates, normalize_penalty, select_centers_by_strategy,
weighted_coefficient_sum_to_zero_transform,
};
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub enum ConstantCurvatureIdentifiability {
#[default]
CenterSumToZero,
FrozenTransform { transform: Array2<f64> },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConstantCurvatureBasisSpec {
pub center_strategy: CenterStrategy,
pub kappa: f64,
#[serde(default)]
pub kappa_fixed: bool,
pub length_scale: f64,
#[serde(default)]
pub length_scale_fixed: bool,
pub double_penalty: bool,
#[serde(default)]
pub identifiability: ConstantCurvatureIdentifiability,
}
impl Default for ConstantCurvatureBasisSpec {
fn default() -> Self {
Self {
center_strategy: CenterStrategy::FarthestPoint { num_centers: 50 },
kappa: 0.0,
kappa_fixed: false,
length_scale: 0.0,
length_scale_fixed: false,
double_penalty: false,
identifiability: ConstantCurvatureIdentifiability::CenterSumToZero,
}
}
}
pub(crate) fn validate_chart_points(
points: ArrayView2<'_, f64>,
kappa: f64,
what: &str,
) -> Result<(), BasisError> {
for (i, row) in points.outer_iter().enumerate() {
let mut nx2 = 0.0_f64;
for &v in row.iter() {
if !v.is_finite() {
crate::bail_invalid_basis!(
"constant-curvature {what} row {i} has a non-finite coordinate"
);
}
nx2 += v * v;
}
if 1.0 + kappa * nx2 <= 0.0 {
crate::bail_invalid_basis!(
"constant-curvature {what} row {i} lies outside the κ-stereographic chart \
(need 1 + κ·‖x‖² > 0; got κ = {kappa}, ‖x‖² = {nx2}); for κ < 0 the chart is \
the open ball ‖x‖ < 1/√(−κ)"
);
}
}
Ok(())
}
#[inline]
fn eta_shape(u: f64) -> f64 {
if u >= 0.5 {
return (-u).exp() * (1.0 + u) - 1.0;
}
let mut term = u * u; let mut factorial = 2.0_f64;
let mut sum = 0.0_f64;
let mut sign = 1.0_f64; for m in 2..=20u32 {
if m > 2 {
term *= u;
factorial *= f64::from(m);
sign = -sign;
}
sum -= sign * f64::from(m - 1) * term / factorial;
}
sum
}
#[inline]
fn eta2_shape(u: f64) -> f64 {
if u >= 0.5 {
return (-u).exp() * (1.0 + u + u * u) - 1.0;
}
let mut term = u * u;
let mut factorial = 2.0_f64;
let mut sum = 0.0_f64;
let mut sign = 1.0_f64;
for m in 2..=20u32 {
if m > 2 {
term *= u;
factorial *= f64::from(m);
sign = -sign;
}
let weight = f64::from(m - 1) * f64::from(m - 1);
sum += sign * weight * term / factorial;
}
sum
}
#[inline]
pub(crate) fn constant_curvature_kernel_scalar(distance: f64, length_scale: f64) -> f64 {
length_scale * (-distance / length_scale).exp_m1()
}
pub fn constant_curvature_kernel_matrix(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
kappa: f64,
length_scale: f64,
) -> Result<Array2<f64>, BasisError> {
if data.ncols() != centers.ncols() {
crate::bail_dim_basis!(
"constant-curvature kernel dimension mismatch: data d={} centers d={}",
data.ncols(),
centers.ncols()
);
}
if !(length_scale.is_finite() && length_scale > 0.0) {
crate::bail_invalid_basis!(
"constant-curvature kernel needs a positive finite length_scale; got {length_scale}"
);
}
validate_chart_points(data, kappa, "data")?;
validate_chart_points(centers, kappa, "centers")?;
let manifold = ConstantCurvature::new(data.ncols(), kappa);
let mut out = Array2::<f64>::zeros((data.nrows(), centers.nrows()));
out.axis_iter_mut(Axis(0))
.into_par_iter()
.enumerate()
.try_for_each(|(i, mut row)| -> Result<(), BasisError> {
for (j, c) in centers.outer_iter().enumerate() {
let d = manifold.distance(data.row(i), c).map_err(|e| {
BasisError::InvalidInput(format!(
"constant-curvature distance failed at (row {i}, center {j}): {e}"
))
})?;
row[j] = constant_curvature_kernel_scalar(d, length_scale);
}
Ok(())
})?;
Ok(out)
}
#[derive(Clone, Debug)]
pub struct ConstantCurvatureKernelPsiJets {
pub value: Array2<f64>,
pub d_kappa: Array2<f64>,
pub d_eta: Array2<f64>,
pub d_kappa2: Array2<f64>,
pub d_kappa_eta: Array2<f64>,
pub d_eta2: Array2<f64>,
}
pub fn constant_curvature_kernel_psi_jets(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
kappa: f64,
length_scale: f64,
) -> Result<ConstantCurvatureKernelPsiJets, BasisError> {
if data.ncols() != centers.ncols() {
crate::bail_dim_basis!(
"constant-curvature kernel-jet dimension mismatch: data d={} centers d={}",
data.ncols(),
centers.ncols()
);
}
if !(length_scale.is_finite() && length_scale > 0.0) {
crate::bail_invalid_basis!(
"constant-curvature kernel jets need a positive finite length_scale; got {length_scale}"
);
}
validate_chart_points(data, kappa, "data")?;
validate_chart_points(centers, kappa, "centers")?;
let manifold = ConstantCurvature::new(data.ncols(), kappa);
let n = data.nrows();
let m = centers.nrows();
let mut jets = ConstantCurvatureKernelPsiJets {
value: Array2::<f64>::zeros((n, m)),
d_kappa: Array2::<f64>::zeros((n, m)),
d_eta: Array2::<f64>::zeros((n, m)),
d_kappa2: Array2::<f64>::zeros((n, m)),
d_kappa_eta: Array2::<f64>::zeros((n, m)),
d_eta2: Array2::<f64>::zeros((n, m)),
};
let rows: Vec<(usize, Vec<[f64; 6]>)> = (0..n)
.into_par_iter()
.map(|i| -> Result<(usize, Vec<[f64; 6]>), BasisError> {
let mut row = Vec::with_capacity(m);
for (j, c) in centers.outer_iter().enumerate() {
let (d, d1, d2) = distance_kappa_jet(&manifold, data.row(i), c).map_err(|e| {
BasisError::InvalidInput(format!(
"constant-curvature distance κ-jet failed at (row {i}, center {j}): {e}"
))
})?;
let q = d / length_scale;
let decay = (-q).exp();
row.push([
constant_curvature_kernel_scalar(d, length_scale),
-d1 * decay,
length_scale * eta_shape(q),
decay * (d1 * d1 / length_scale - d2),
-d1 * q * decay,
length_scale * eta2_shape(q),
]);
}
Ok((i, row))
})
.collect::<Result<Vec<_>, BasisError>>()?;
for (i, row) in rows {
for (j, entry) in row.into_iter().enumerate() {
jets.value[(i, j)] = entry[0];
jets.d_kappa[(i, j)] = entry[1];
jets.d_eta[(i, j)] = entry[2];
jets.d_kappa2[(i, j)] = entry[3];
jets.d_kappa_eta[(i, j)] = entry[4];
jets.d_eta2[(i, j)] = entry[5];
}
}
Ok(jets)
}
pub fn constant_curvature_kernel_kappa_jets(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
kappa: f64,
length_scale: f64,
) -> Result<(Array2<f64>, Array2<f64>, Array2<f64>), BasisError> {
let jets = constant_curvature_kernel_psi_jets(data, centers, kappa, length_scale)?;
Ok((jets.value, jets.d_kappa, jets.d_kappa2))
}
pub fn realized_constant_curvature_length_scale(
centers: ArrayView2<'_, f64>,
spec_length_scale: f64,
) -> Result<f64, BasisError> {
if spec_length_scale.is_finite() && spec_length_scale > 0.0 {
return Ok(spec_length_scale);
}
if spec_length_scale != 0.0 {
crate::bail_invalid_basis!(
"constant-curvature length_scale must be positive (or 0.0 for auto); got {spec_length_scale}"
);
}
let dists = center_chart_gauge_distances(centers)?;
let median = dists[dists.len() / 2];
if !(median.is_finite() && median > 0.0) {
crate::bail_invalid_basis!(
"constant-curvature auto length_scale failed: centers are degenerate \
(median pairwise chart distance = {median})"
);
}
Ok(median)
}
fn center_chart_gauge_distances(centers: ArrayView2<'_, f64>) -> Result<Vec<f64>, BasisError> {
let m = centers.nrows();
if m < 2 {
return Err(BasisError::InsufficientColumnsForConstraint { found: m });
}
let mut dists: Vec<f64> = Vec::with_capacity(m * (m - 1) / 2);
for i in 0..m {
for j in (i + 1)..m {
let mut s = 0.0_f64;
for k in 0..centers.ncols() {
let dlt = centers[(i, k)] - centers[(j, k)];
s += dlt * dlt;
}
dists.push(2.0 * s.sqrt());
}
}
dists.sort_by(|a, b| a.partial_cmp(b).expect("finite chart distances"));
Ok(dists)
}
pub fn constant_curvature_evaluated_scale_span(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
) -> Result<(f64, f64), BasisError> {
if data.ncols() != centers.ncols() {
crate::bail_dim_basis!(
"constant-curvature scale span dimension mismatch: data d={} centers d={}",
data.ncols(),
centers.ncols()
);
}
let mut lo = f64::INFINITY;
let mut hi = 0.0_f64;
let mut observe = |a: ndarray::ArrayView1<'_, f64>, b: ndarray::ArrayView1<'_, f64>| {
let mut sum = 0.0_f64;
for k in 0..a.len() {
let delta = a[k] - b[k];
sum += delta * delta;
}
let d = 2.0 * sum.sqrt();
if d.is_finite() && d > 0.0 {
lo = lo.min(d);
hi = hi.max(d);
}
};
for x in data.outer_iter() {
for c in centers.outer_iter() {
observe(x, c);
}
}
for i in 0..centers.nrows() {
for j in (i + 1)..centers.nrows() {
observe(centers.row(i), centers.row(j));
}
}
if !(lo.is_finite() && lo > 0.0 && hi.is_finite() && hi >= lo) {
crate::bail_invalid_basis!(
"constant-curvature range window is undefined: the evaluated pairs carry no \
positive chart distance (d_min = {lo}, d_max = {hi})"
);
}
Ok((lo, hi))
}
pub fn constant_curvature_length_scale_bounds(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
) -> Result<(f64, f64), BasisError> {
let (d_min, d_max) = constant_curvature_evaluated_scale_span(data, centers)?;
let gram_resolvable_efolds = 0.5 * -f64::EPSILON.ln();
let lo = d_max / gram_resolvable_efolds;
let hi = d_max / (2.0 * f64::EPSILON.sqrt());
if !(lo.is_finite() && lo > 0.0 && hi.is_finite() && hi > lo) {
crate::bail_invalid_basis!(
"constant-curvature range box collapsed: [{lo}, {hi}] from an evaluated span of \
[{d_min}, {d_max}]"
);
}
Ok((lo, hi))
}
pub fn build_constant_curvature_basis(
data: ArrayView2<'_, f64>,
spec: &ConstantCurvatureBasisSpec,
) -> Result<BasisBuildResult, BasisError> {
if data.ncols() == 0 {
crate::bail_invalid_basis!("constant-curvature smooth needs at least one feature column");
}
if !spec.kappa.is_finite() {
crate::bail_invalid_basis!("constant-curvature smooth needs a finite kappa");
}
validate_chart_points(data, spec.kappa, "data")?;
let centers = select_constant_curvature_centers(data, &spec.center_strategy)?;
if centers.nrows() < 2 {
return Err(BasisError::InsufficientColumnsForConstraint {
found: centers.nrows(),
});
}
validate_chart_points(centers.view(), spec.kappa, "centers")?;
let length_scale = realized_constant_curvature_length_scale(centers.view(), spec.length_scale)?;
let raw_penalty =
constant_curvature_kernel_matrix(centers.view(), centers.view(), spec.kappa, length_scale)?;
let z = match &spec.identifiability {
ConstantCurvatureIdentifiability::FrozenTransform { transform } => {
if transform.nrows() != centers.nrows() {
crate::bail_dim_basis!(
"frozen constant-curvature identifiability transform mismatch: {} centers but transform has {} rows",
centers.nrows(),
transform.nrows()
);
}
transform.clone()
}
ConstantCurvatureIdentifiability::CenterSumToZero => {
let weights = Array1::<f64>::ones(centers.nrows());
weighted_coefficient_sum_to_zero_transform(weights.view())?
}
};
let gauge = gam_problem::Gauge::from_block_transforms(&[z.clone()]);
let penalty = ConstructiveQuadratic::try_from_dense_psd(
symmetrize(&gauge.restrict_penalty(&raw_penalty)),
"constant-curvature restricted RKHS penalty",
)?;
let raw_design =
constant_curvature_kernel_matrix(data, centers.view(), spec.kappa, length_scale)?;
let design = gam_linalg::matrix::DesignMatrix::Dense(
gam_linalg::matrix::DenseDesignMatrix::from(gauge.restrict_design(&raw_design)),
);
let mut candidates = vec![PenaltyCandidate {
matrix: penalty,
source: PenaltySource::Primary,
normalization_scale: 1.0,
kronecker_factors: None,
op: None,
}];
if spec.double_penalty {
let ridge = Array2::<f64>::eye(design.ncols());
let (ridge_norm, c_ridge) = normalize_penalty(&ridge);
candidates.push(PenaltyCandidate {
matrix: ConstructiveQuadratic::try_from_dense_psd(
ridge_norm,
"constant-curvature whole-function ridge",
)?,
source: PenaltySource::DoublePenaltyNullspace,
normalization_scale: c_ridge,
kronecker_factors: None,
op: None,
});
}
let filtered = filter_penalty_candidates(candidates)?;
Ok(BasisBuildResult {
design,
affine_offset: None,
active_penalties: filtered.active,
dropped_penalties: filtered.dropped,
metadata: BasisMetadata::ConstantCurvature {
centers,
kappa: spec.kappa,
length_scale,
constraint_transform: Some(z),
},
kronecker_factored: None,
joint_null_rotation: None,
})
}
pub fn constant_curvature_center_chart_radius2(
data: ArrayView2<'_, f64>,
feature_cols: &[usize],
strategy: &CenterStrategy,
) -> f64 {
match strategy {
CenterStrategy::Auto(inner) => {
constant_curvature_center_chart_radius2(data, feature_cols, inner)
}
CenterStrategy::DuchonSpectral { knots, .. } => {
constant_curvature_center_chart_radius2(data, feature_cols, knots)
}
CenterStrategy::UserProvided(centers) => {
let mut max_r2 = 0.0_f64;
for row in centers.outer_iter() {
let mut r2 = 0.0_f64;
for &v in row.iter() {
if v.is_finite() {
r2 += v * v;
}
}
max_r2 = max_r2.max(r2);
}
max_r2
}
CenterStrategy::UniformGrid { .. } => {
let mut corner_r2 = 0.0_f64;
for &c in feature_cols.iter() {
let mut lo = f64::INFINITY;
let mut hi = f64::NEG_INFINITY;
for row in data.outer_iter() {
if let Some(&v) = row.get(c)
&& v.is_finite()
{
lo = lo.min(v);
hi = hi.max(v);
}
}
if lo.is_finite() && hi.is_finite() {
let extreme = lo.abs().max(hi.abs());
corner_r2 += extreme * extreme;
}
}
corner_r2
}
CenterStrategy::EqualMass { .. }
| CenterStrategy::EqualMassCovarRepresentative { .. }
| CenterStrategy::FarthestPoint { .. }
| CenterStrategy::KMeans { .. } => {
constant_curvature_data_chart_radius2(data, feature_cols)
}
}
}
pub fn constant_curvature_data_chart_radius2(
data: ArrayView2<'_, f64>,
feature_cols: &[usize],
) -> f64 {
let mut max_r2 = 0.0_f64;
for row in data.outer_iter() {
let mut r2 = 0.0_f64;
for &c in feature_cols.iter() {
if let Some(&v) = row.get(c)
&& v.is_finite()
{
r2 += v * v;
}
}
max_r2 = max_r2.max(r2);
}
max_r2
}
fn select_constant_curvature_centers(
data: ArrayView2<'_, f64>,
strategy: &CenterStrategy,
) -> Result<Array2<f64>, BasisError> {
let mut centers = select_centers_by_strategy(data, strategy)?;
match strategy {
CenterStrategy::UserProvided(_) => return Ok(centers),
CenterStrategy::Auto(inner) => {
if matches!(inner.as_ref(), CenterStrategy::UserProvided(_)) {
return Ok(centers);
}
}
CenterStrategy::DuchonSpectral { knots, .. } => {
if center_strategy_kind(knots) == CenterStrategyKind::UserProvided {
return Ok(centers);
}
}
CenterStrategy::EqualMass { .. }
| CenterStrategy::EqualMassCovarRepresentative { .. }
| CenterStrategy::FarthestPoint { .. }
| CenterStrategy::KMeans { .. }
| CenterStrategy::UniformGrid { .. } => {}
}
if centers.nrows() == 0 || centers.ncols() == 0 {
return Ok(centers);
}
let (closest, _) = centers
.outer_iter()
.enumerate()
.map(|(i, row)| (i, row.dot(&row)))
.min_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
.expect("centers has at least one row; the empty case returned above");
for j in 0..centers.ncols() {
centers[(closest, j)] = 0.0;
}
Ok(centers)
}
pub fn constant_curvature_realized_centers(
data: ArrayView2<'_, f64>,
spec: &ConstantCurvatureBasisSpec,
) -> Result<Array2<f64>, BasisError> {
let centers = select_constant_curvature_centers(data, &spec.center_strategy)?;
if centers.nrows() < 2 {
return Err(BasisError::InsufficientColumnsForConstraint {
found: centers.nrows(),
});
}
Ok(centers)
}
pub(crate) fn symmetrize(m: &Array2<f64>) -> Array2<f64> {
gam_linalg::matrix::symmetrize(m)
}
pub(crate) fn active_constant_curvature_penalty_derivatives(
penalties: &[ActivePenalty],
primary_derivative: &Array2<f64>,
) -> Result<Vec<Array2<f64>>, BasisError> {
penalties
.iter()
.map(|penalty| match &penalty.info.source {
PenaltySource::Primary => Ok(primary_derivative.clone()),
PenaltySource::DoublePenaltyNullspace => {
Ok(Array2::<f64>::zeros(primary_derivative.raw_dim()))
}
other => Err(BasisError::InvalidInput(format!(
"unexpected constant-curvature penalty source in κ-derivative path: {other:?}"
))),
})
.collect()
}
#[derive(Clone, Debug)]
pub struct ConstantCurvaturePsiJets {
pub design_kappa: Array2<f64>,
pub design_eta: Array2<f64>,
pub design_kappa2: Array2<f64>,
pub design_kappa_eta: Array2<f64>,
pub design_eta2: Array2<f64>,
pub penalties_kappa: Vec<Array2<f64>>,
pub penalties_eta: Vec<Array2<f64>>,
pub penalties_kappa2: Vec<Array2<f64>>,
pub penalties_kappa_eta: Vec<Array2<f64>>,
pub penalties_eta2: Vec<Array2<f64>>,
}
pub fn build_constant_curvature_basis_psi_derivatives(
data: ArrayView2<'_, f64>,
spec: &ConstantCurvatureBasisSpec,
) -> Result<ConstantCurvaturePsiJets, BasisError> {
if data.ncols() == 0 {
crate::bail_invalid_basis!("constant-curvature smooth needs at least one feature column");
}
if !spec.kappa.is_finite() {
crate::bail_invalid_basis!("constant-curvature smooth needs a finite kappa");
}
validate_chart_points(data, spec.kappa, "data")?;
let centers = select_constant_curvature_centers(data, &spec.center_strategy)?;
if centers.nrows() < 2 {
return Err(BasisError::InsufficientColumnsForConstraint {
found: centers.nrows(),
});
}
validate_chart_points(centers.view(), spec.kappa, "centers")?;
let length_scale = realized_constant_curvature_length_scale(centers.view(), spec.length_scale)?;
let z = match &spec.identifiability {
ConstantCurvatureIdentifiability::FrozenTransform { transform } => {
if transform.nrows() != centers.nrows() {
crate::bail_dim_basis!(
"frozen constant-curvature identifiability transform mismatch: {} centers but transform has {} rows",
centers.nrows(),
transform.nrows()
);
}
transform.clone()
}
ConstantCurvatureIdentifiability::CenterSumToZero => {
let weights = Array1::<f64>::ones(centers.nrows());
weighted_coefficient_sum_to_zero_transform(weights.view())?
}
};
let gauge = gam_problem::Gauge::from_block_transforms(&[z.clone()]);
let dc = constant_curvature_kernel_psi_jets(data, centers.view(), spec.kappa, length_scale)?;
let cc = constant_curvature_kernel_psi_jets(
centers.view(),
centers.view(),
spec.kappa,
length_scale,
)?;
let base = build_constant_curvature_basis(data, spec)?;
let penalty_block = |raw: &Array2<f64>| -> Result<Vec<Array2<f64>>, BasisError> {
let restricted = symmetrize(&gauge.restrict_penalty(raw));
active_constant_curvature_penalty_derivatives(&base.active_penalties, &restricted)
};
Ok(ConstantCurvaturePsiJets {
design_kappa: gauge.restrict_design(&dc.d_kappa),
design_eta: gauge.restrict_design(&dc.d_eta),
design_kappa2: gauge.restrict_design(&dc.d_kappa2),
design_kappa_eta: gauge.restrict_design(&dc.d_kappa_eta),
design_eta2: gauge.restrict_design(&dc.d_eta2),
penalties_kappa: penalty_block(&cc.d_kappa)?,
penalties_eta: penalty_block(&cc.d_eta)?,
penalties_kappa2: penalty_block(&cc.d_kappa2)?,
penalties_kappa_eta: penalty_block(&cc.d_kappa_eta)?,
penalties_eta2: penalty_block(&cc.d_eta2)?,
})
}
pub fn build_constant_curvature_basis_kappa_derivatives(
data: ArrayView2<'_, f64>,
spec: &ConstantCurvatureBasisSpec,
) -> Result<BasisPsiDerivativeBundle, BasisError> {
let jets = build_constant_curvature_basis_psi_derivatives(data, spec)?;
Ok(BasisPsiDerivativeBundle {
first: BasisPsiDerivativeResult {
design_derivative: jets.design_kappa,
penalties_derivative: jets.penalties_kappa,
implicit_operator: None,
},
second: BasisPsiSecondDerivativeResult {
designsecond_derivative: jets.design_kappa2,
penaltiessecond_derivative: jets.penalties_kappa2,
implicit_operator: None,
},
implicit_operator: None,
})
}
#[cfg(test)]
mod tests {
use super::*;
use gam_linalg::faer_ndarray::FaerEigh;
#[test]
pub(crate) fn kernel_spread_collapses_with_kappa_at_frozen_length_scale() {
let centers = ndarray::array![
[0.10, 0.05],
[-0.20, 0.15],
[0.30, -0.10],
[-0.05, -0.25],
[0.22, 0.20],
[-0.30, -0.05],
[0.05, 0.30],
[-0.15, 0.10],
];
let ell_frozen = realized_constant_curvature_length_scale(centers.view(), 0.0)
.expect("fixture centers span a positive pairwise distance");
let spread = |kappa: f64, ell: f64| -> f64 {
let k = constant_curvature_kernel_matrix(centers.view(), centers.view(), kappa, ell)
.expect("fixture centers are distinct and the length scale is positive");
let m = k.nrows();
let mut s = 0.0;
let mut cnt = 0.0;
for i in 0..m {
for j in 0..m {
if i != j {
s += k[(i, j)];
cnt += 1.0;
}
}
}
1.0 - s / cnt
};
let s_neg = spread(-2.0, ell_frozen);
let s_zero = spread(0.0, ell_frozen);
let s_pos = spread(2.0, ell_frozen);
eprintln!(
"[κ-collapse] frozen ℓ={ell_frozen:.4}: spread κ=-2 {s_neg:.4} | κ=0 {s_zero:.4} | κ=+2 {s_pos:.4}"
);
assert!(
s_pos < s_zero && s_zero < s_neg,
"expected kernel spread to shrink with κ at frozen ℓ: κ=-2 {s_neg} κ=0 {s_zero} κ=+2 {s_pos}"
);
let weights = Array1::<f64>::ones(centers.nrows());
let z = weighted_coefficient_sum_to_zero_transform(weights.view())
.expect("fixture weights are positive, so the sum-to-zero transform exists");
let logdet_norm_penalty = |kappa: f64, ell: f64| -> f64 {
let k = constant_curvature_kernel_matrix(centers.view(), centers.view(), kappa, ell)
.expect("fixture centers are distinct and the length scale is positive");
let s_raw = symmetrize(&z.t().dot(&k).dot(&z));
let (s_norm, _c) = normalize_penalty(&s_raw);
let sym = symmetrize(&s_norm);
let (evals, _v) = FaerEigh::eigh(&sym, faer::Side::Lower)
.expect("the fixture Gram is symmetric, so eigh converges");
let max = evals.iter().cloned().fold(0.0_f64, f64::max);
let tol = max * 1e-9;
evals
.iter()
.filter(|&&e| e > tol)
.map(|&e| e.ln())
.sum::<f64>()
};
let l_neg = logdet_norm_penalty(-2.0, ell_frozen);
let l_zero = logdet_norm_penalty(0.0, ell_frozen);
let l_pos = logdet_norm_penalty(2.0, ell_frozen);
eprintln!(
"[κ-collapse] log|S~|_+ (frozen ℓ): κ=-2 {l_neg:.4} | κ=0 {l_zero:.4} | κ=+2 {l_pos:.4}"
);
let geo_median_ell = |kappa: f64| -> f64 {
let m = centers.nrows();
let manifold = ConstantCurvature::new(centers.ncols(), kappa);
let mut dists = Vec::with_capacity(m * (m - 1) / 2);
for i in 0..m {
for j in (i + 1)..m {
dists.push(
manifold
.distance(centers.row(i), centers.row(j))
.expect("fixture centers lie on the manifold"),
);
}
}
dists.sort_by(|a, b| a.partial_cmp(b).expect("pairwise distances are finite"));
dists[dists.len() / 2]
};
let gs_neg = spread(-2.0, geo_median_ell(-2.0));
let gs_zero = spread(0.0, geo_median_ell(0.0));
let gs_pos = spread(2.0, geo_median_ell(2.0));
let gl_neg = logdet_norm_penalty(-2.0, geo_median_ell(-2.0));
let gl_zero = logdet_norm_penalty(0.0, geo_median_ell(0.0));
let gl_pos = logdet_norm_penalty(2.0, geo_median_ell(2.0));
eprintln!(
"[κ-collapse] geodesic ℓ: spread κ=-2 {gs_neg:.4} | κ=0 {gs_zero:.4} | κ=+2 {gs_pos:.4}"
);
eprintln!(
"[κ-collapse] geodesic ℓ: log|S~|_+ κ=-2 {gl_neg:.4} | κ=0 {gl_zero:.4} | κ=+2 {gl_pos:.4}"
);
let logdet_raw = |kappa: f64, ell: f64, c0: f64| -> f64 {
let k = constant_curvature_kernel_matrix(centers.view(), centers.view(), kappa, ell)
.expect("fixture centers are distinct and the length scale is positive");
let s_raw = symmetrize(&z.t().dot(&k).dot(&z));
let scaled = s_raw.mapv(|v| v / c0);
let (evals, _v) = FaerEigh::eigh(&scaled, faer::Side::Lower)
.expect("the fixture Gram is symmetric, so eigh converges");
let max = evals.iter().cloned().fold(0.0_f64, f64::max);
let tol = max * 1e-9;
evals
.iter()
.filter(|&&e| e > tol)
.map(|&e| e.ln())
.sum::<f64>()
};
let k0 = constant_curvature_kernel_matrix(centers.view(), centers.view(), 0.0, ell_frozen)
.expect("fixture centers are distinct and the length scale is positive");
let s_raw0 = symmetrize(&z.t().dot(&k0).dot(&z));
let c0 = s_raw0.iter().map(|v| v * v).sum::<f64>().sqrt();
let r_neg = logdet_raw(-2.0, ell_frozen, c0);
let r_zero = logdet_raw(0.0, ell_frozen, c0);
let r_pos = logdet_raw(2.0, ell_frozen, c0);
eprintln!(
"[κ-collapse] frozen-c₀ log|S_raw/c₀|_+ (frozen ℓ): κ=-2 {r_neg:.4} | κ=0 {r_zero:.4} | κ=+2 {r_pos:.4}"
);
eprint!("[κ-collapse] frozen-c₀ grid:");
for kk in [-2.0, -1.0, -0.5, 0.0, 0.5, 1.0, 2.0] {
eprint!(" κ={kk}:{:.4}", logdet_raw(kk, ell_frozen, c0));
}
eprintln!();
}
pub(crate) fn oracle_disk_design_centers() -> (Array2<f64>, Array2<f64>) {
let centers = ndarray::array![
[0.10, 0.05],
[-0.20, 0.15],
[0.30, -0.10],
[-0.05, -0.25],
[0.22, 0.20],
[-0.30, -0.05],
[0.05, 0.30],
[-0.15, 0.10],
];
let mut state = 0x2545_f491_4f6c_dd1d_u64;
let mut next = || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
((state >> 11) as f64 / (1u64 << 53) as f64 - 0.5) * 0.84
};
let n = 60usize;
let mut data = Array2::<f64>::zeros((n, 2));
for i in 0..n {
data[(i, 0)] = next();
data[(i, 1)] = next();
}
(data, centers)
}
#[test]
pub(crate) fn range_box_refuses_a_degenerate_geometry() {
let coincident = ndarray::array![[0.3_f64, -0.2], [0.3, -0.2]];
let error = constant_curvature_length_scale_bounds(coincident.view(), coincident.view())
.expect_err("a cloud with no positive pairwise distance has no range box");
let message = format!("{error}");
assert!(
message.contains("no") && message.contains("positive"),
"the refusal must name what is missing; got {message}"
);
let pair = ndarray::array![[0.0_f64, 0.0], [0.2, 0.0]];
let (lo, hi) = constant_curvature_length_scale_bounds(pair.view(), pair.view())
.expect("one positive pairwise distance is enough");
let format_width = (0.5 * -f64::EPSILON.ln()) / (2.0 * f64::EPSILON.sqrt());
assert!(
lo > 0.0 && (hi / lo - format_width).abs() <= 1.0e-9 * format_width,
"a one-distance geometry's box width is the format's, {format_width}; got \
[{lo}, {hi}] with ratio {}",
hi / lo
);
}
#[test]
pub(crate) fn kernel_psi_jets_match_central_differences_in_both_coordinates() {
let (data, centers) = oracle_disk_design_centers();
let ell0 = realized_constant_curvature_length_scale(centers.view(), 0.0)
.expect("fixture centers span a positive pairwise distance");
let eta0 = ell0.ln();
let at = |kappa: f64, eta: f64| {
constant_curvature_kernel_psi_jets(data.view(), centers.view(), kappa, eta.exp())
.expect("the fixture disk is inside every probed chart")
};
let rel = |exact: &Array2<f64>, fd: &Array2<f64>| -> f64 {
let mut err = 0.0_f64;
let mut scale = 0.0_f64;
for (&a, &b) in exact.iter().zip(fd.iter()) {
err = err.max((a - b).abs());
scale = scale.max(a.abs()).max(b.abs());
}
err / scale.max(1.0)
};
let h = 1.0e-5_f64;
for &kappa in &[-1.5_f64, -0.5, -1e-7, 0.0, 1e-7, 0.8, 1.7] {
for &eta in &[eta0 - 0.7, eta0, eta0 + 0.7] {
let jets = at(kappa, eta);
let kp = at(kappa + h, eta);
let km = at(kappa - h, eta);
let ep = at(kappa, eta + h);
let em = at(kappa, eta - h);
let central = |plus: &Array2<f64>, minus: &Array2<f64>| -> Array2<f64> {
(plus - minus) / (2.0 * h)
};
let checks = [
("∂K/∂κ", &jets.d_kappa, central(&kp.value, &km.value), 1e-6),
("∂K/∂η", &jets.d_eta, central(&ep.value, &em.value), 1e-6),
(
"∂²K/∂κ²",
&jets.d_kappa2,
central(&kp.d_kappa, &km.d_kappa),
1e-5,
),
(
"∂²K/∂κ∂η",
&jets.d_kappa_eta,
central(&ep.d_kappa, &em.d_kappa),
1e-5,
),
("∂²K/∂η²", &jets.d_eta2, central(&ep.d_eta, &em.d_eta), 1e-5),
];
for (label, exact, fd, tol) in checks {
let error = rel(exact, &fd);
assert!(
error < tol,
"κ={kappa} η={eta}: {label} disagrees with its central difference: rel={error:.6e}"
);
}
}
}
}
#[test]
pub(crate) fn design_and_penalty_are_one_gram_at_one_range() {
let (data, centers) = oracle_disk_design_centers();
for kappa in [-1.2_f64, -0.4, 0.0, 0.4, 1.2] {
let spec = ConstantCurvatureBasisSpec {
center_strategy: CenterStrategy::UserProvided(centers.clone()),
kappa,
length_scale: 1.3,
..Default::default()
};
let built = build_constant_curvature_basis(data.view(), &spec).expect("build");
let BasisMetadata::ConstantCurvature {
length_scale,
constraint_transform,
..
} = &built.metadata
else {
panic!("expected ConstantCurvature metadata");
};
assert_eq!(
*length_scale, 1.3,
"the realized range is the spec's range, not a κ-remapped one"
);
let z = constraint_transform.as_ref().expect("constraint transform");
let k_dc =
constant_curvature_kernel_matrix(data.view(), centers.view(), kappa, *length_scale)
.expect("design kernel");
let k_cc = constant_curvature_kernel_matrix(
centers.view(),
centers.view(),
kappa,
*length_scale,
)
.expect("penalty kernel");
let design = built.design.to_dense();
for (a, b) in design.iter().zip(k_dc.dot(z).iter()) {
assert!(
(a - b).abs() < 1e-12,
"κ={kappa}: design != K(ℓ)·z ({a} vs {b})"
);
}
let gram = symmetrize(&z.t().dot(&k_cc).dot(z));
let primary = built
.active_penalties
.iter()
.find(|penalty| matches!(penalty.info.source, PenaltySource::Primary))
.expect("primary RKHS penalty");
for (a, b) in gram.iter().zip(primary.matrix.iter()) {
assert!(
(a - b).abs() < 1e-12,
"κ={kappa}: penalty != zᵀK(ℓ)z at the SAME ℓ ({a} vs {b})"
);
}
}
}
#[test]
pub(crate) fn range_box_is_the_gram_conditioning_wall_and_contains_the_scale_span() {
let (data, centers) = oracle_disk_design_centers();
let (span_lo, span_hi) =
constant_curvature_evaluated_scale_span(data.view(), centers.view())
.expect("the fixture carries positive evaluated distances");
let (lo, hi) = constant_curvature_length_scale_bounds(data.view(), centers.view())
.expect("box is derivable");
let seed = realized_constant_curvature_length_scale(centers.view(), 0.0).expect("seed");
assert!(
span_hi < hi && lo < seed && seed < hi,
"the box [{lo}, {hi}] must contain the coarse end of the evaluated span \
[{span_lo}, {span_hi}] and the auto seed {seed}"
);
let gram_range = |ell: f64| (-2.0 * span_hi / ell).exp();
assert!(
1.0 + gram_range(lo) != 1.0,
"at ℓ_lo the Gram's far entries must still perturb the diagonal; got {}",
gram_range(lo)
);
assert!(
1.0 + gram_range(lo / 2.0) == 1.0,
"half an ℓ_lo below, the Gram's far entries must round into the diagonal"
);
let limit_departure = |ell: f64| {
let k = constant_curvature_kernel_scalar(span_hi, ell);
((k + span_hi) / span_hi).abs()
};
let root_eps = f64::EPSILON.sqrt();
assert!(
limit_departure(hi) <= root_eps * 1.01,
"at ℓ_hi the widest evaluated pair must BE its own limit to √ε; departure {:.3e} \
against {root_eps:.3e}",
limit_departure(hi)
);
assert!(
limit_departure(hi / 10.0) > root_eps,
"an order of magnitude below ℓ_hi the limit must still be resolvable, or the \
chart is being truncated earlier than the model justifies; departure {:.3e}",
limit_departure(hi / 10.0)
);
let retired_top = span_lo / root_eps;
assert!(
(hi - retired_top).abs() > 0.1 * retired_top,
"ℓ_hi must no longer be the retired cancellation wall {retired_top:.4e}; got {hi:.4e}"
);
}
#[test]
fn constant_curvature_gram_is_full_rank_so_identity_is_the_only_double_penalty() {
let centers = ndarray::array![
[0.10, 0.05],
[-0.20, 0.15],
[0.30, -0.10],
[-0.05, -0.25],
[0.22, 0.20],
[-0.30, -0.05],
[0.05, 0.30],
[-0.15, 0.10],
];
let weights = Array1::<f64>::ones(centers.nrows());
let z = weighted_coefficient_sum_to_zero_transform(weights.view())
.expect("fixture weights are positive, so the sum-to-zero transform exists");
let ell = realized_constant_curvature_length_scale(centers.view(), 0.0)
.expect("fixture centers span a positive pairwise distance");
for &kappa in &[-2.0_f64, -0.5, 0.0, 0.5, 2.0] {
let k = constant_curvature_kernel_matrix(centers.view(), centers.view(), kappa, ell)
.expect("fixture centers are distinct and the length scale is positive");
let raw = symmetrize(&z.t().dot(&k).dot(&z));
let (evals, _v) = FaerEigh::eigh(&raw, faer::Side::Lower)
.expect("the fixture Gram is symmetric, so eigh converges");
let max = evals.iter().cloned().fold(0.0_f64, f64::max);
let min = evals.iter().cloned().fold(f64::INFINITY, f64::min);
assert!(
max > 0.0 && min > max * 1e-9,
"constant-curvature Gram must be full-rank PD at κ={kappa}: \
min eig {min:e}, max eig {max:e}"
);
}
}
#[test]
fn geodesic_exponential_kernel_has_unbounded_curvature_at_a_center() {
let kappa = 0.0_f64;
let ell = 1.0_f64;
let center = ndarray::arr2(&[[0.0_f64, 0.0]]);
let curvature_at = |base: [f64; 2], h: f64| -> f64 {
let probe = ndarray::arr2(&[
[base[0] - h, base[1]],
[base[0], base[1]],
[base[0] + h, base[1]],
]);
let k = constant_curvature_kernel_matrix(probe.view(), center.view(), kappa, ell)
.expect("fixture centers are distinct and the length scale is positive");
(k[[0, 0]] - 2.0 * k[[1, 0]] + k[[2, 0]]) / (h * h)
};
let mut at_center = Vec::new();
let mut off_center = Vec::new();
for &h in &[1.0e-2_f64, 1.0e-3, 1.0e-4] {
at_center.push(curvature_at([0.0, 0.0], h).abs());
off_center.push(curvature_at([0.5, 0.0], h).abs());
}
assert!(
(off_center[2] - off_center[1]).abs() <= 1.0e-3 * off_center[1].max(1.0),
"the kernel must be C² away from its centers; got {off_center:?}"
);
assert!(
at_center[2] > 10.0 * at_center[0],
"a Matérn-½ cusp must make the second difference diverge as h -> 0; \
got {at_center:?}"
);
assert!(
at_center[2] > 100.0 * off_center[2],
"the at-center curvature must dwarf the smooth-region curvature; \
at-center {:?} vs off-center {:?}",
at_center[2],
off_center[2]
);
}
#[test]
fn kappa_second_derivatives_match_a_central_difference_of_the_first_2458() {
let data = ndarray::array![
[0.10, 0.05],
[-0.20, 0.15],
[0.30, -0.10],
[-0.05, -0.25],
[0.22, 0.20],
[-0.30, -0.05],
[0.05, 0.30],
[-0.15, 0.10],
];
let spec = ConstantCurvatureBasisSpec {
center_strategy: CenterStrategy::FarthestPoint { num_centers: 6 },
..Default::default()
};
let kappa0 = 0.35_f64;
let first_at = |kappa: f64| {
let mut probe = spec.clone();
probe.kappa = kappa;
let bundle = build_constant_curvature_basis_kappa_derivatives(data.view(), &probe)
.expect("fixture points lie inside the chart for every probed kappa");
(
bundle.first.design_derivative,
bundle.first.penalties_derivative,
)
};
let mut exact_spec = spec.clone();
exact_spec.kappa = kappa0;
let analytic = build_constant_curvature_basis_kappa_derivatives(data.view(), &exact_spec)
.expect("fixture points lie inside the chart at kappa0");
let design_second = analytic.second.designsecond_derivative;
let penalty_second = analytic.second.penaltiessecond_derivative;
let max_rel_error_at = |h: f64| -> (f64, f64) {
let (x_plus, s_plus) = first_at(kappa0 + h);
let (x_minus, s_minus) = first_at(kappa0 - h);
let mut design_error = 0.0_f64;
let mut design_scale = 0.0_f64;
for ((&plus, &minus), &exact) in
x_plus.iter().zip(x_minus.iter()).zip(design_second.iter())
{
let fd = (plus - minus) / (2.0 * h);
design_error = design_error.max((fd - exact).abs());
design_scale = design_scale.max(exact.abs()).max(fd.abs());
}
assert_eq!(
s_plus.len(),
penalty_second.len(),
"penalty block count must not depend on kappa"
);
let mut penalty_error = 0.0_f64;
let mut penalty_scale = 0.0_f64;
for ((block_plus, block_minus), block_exact) in
s_plus.iter().zip(s_minus.iter()).zip(penalty_second.iter())
{
for ((&plus, &minus), &exact) in block_plus
.iter()
.zip(block_minus.iter())
.zip(block_exact.iter())
{
let fd = (plus - minus) / (2.0 * h);
penalty_error = penalty_error.max((fd - exact).abs());
penalty_scale = penalty_scale.max(exact.abs()).max(fd.abs());
}
}
(
design_error / design_scale.max(1.0),
penalty_error / penalty_scale.max(1.0),
)
};
let h = 1.0e-4_f64;
let (design_rel, penalty_rel) = max_rel_error_at(h);
eprintln!(
"[2458-second-fd] h={h:.1e}: design rel={design_rel:.3e} penalty rel={penalty_rel:.3e}"
);
assert!(
design_rel < 1.0e-6,
"d2X/dkappa2 disagrees with a central difference of dX/dkappa: rel={design_rel:.6e}"
);
assert!(
penalty_rel < 1.0e-6,
"d2S/dkappa2 disagrees with a central difference of dS/dkappa: rel={penalty_rel:.6e}"
);
let (design_rel_half, penalty_rel_half) = max_rel_error_at(0.5 * h);
eprintln!(
"[2458-second-fd] h={:.1e}: design rel={design_rel_half:.3e} penalty rel={penalty_rel_half:.3e}",
0.5 * h
);
assert!(
design_rel_half <= design_rel.max(1.0e-9) * 2.0,
"halving h must not inflate the design disagreement (missing-term signature): \
{design_rel:.6e} -> {design_rel_half:.6e}"
);
assert!(
penalty_rel_half <= penalty_rel.max(1.0e-9) * 2.0,
"halving h must not inflate the penalty disagreement (missing-term signature): \
{penalty_rel:.6e} -> {penalty_rel_half:.6e}"
);
}
#[test]
fn the_contrast_gauge_is_the_same_model_and_the_exp_gauge_loses_it_2747() {
let (data, centers) = oracle_disk_design_centers();
let manifold = ConstantCurvature::new(2, 0.6);
let exp_gauge_design = |ell: f64| -> Array2<f64> {
let mut raw = Array2::<f64>::zeros((data.nrows(), centers.nrows()));
for i in 0..data.nrows() {
for j in 0..centers.nrows() {
let d = manifold
.distance(data.row(i), centers.row(j))
.expect("fixture disk is inside the κ = 0.6 chart");
raw[(i, j)] = (-d / ell).exp();
}
}
let weights = Array1::<f64>::ones(centers.nrows());
let z = weighted_coefficient_sum_to_zero_transform(weights.view())
.expect("uniform sum-to-zero frame");
raw.dot(&z).mapv(|value| value * ell)
};
let realized_design = |ell: f64| -> Array2<f64> {
let spec = ConstantCurvatureBasisSpec {
center_strategy: CenterStrategy::UserProvided(centers.clone()),
kappa: 0.6,
length_scale: ell,
..Default::default()
};
build_constant_curvature_basis(data.view(), &spec)
.expect("build")
.design
.to_dense()
};
let disagreement = |ell: f64| -> f64 {
let a = realized_design(ell);
let b = exp_gauge_design(ell);
let mut num = 0.0_f64;
let mut den = 0.0_f64;
for (&x, &y) in a.iter().zip(b.iter()) {
num += (x - y) * (x - y);
den += x * x;
}
(num / den).sqrt()
};
for ell in [0.25_f64, 1.0, 4.0, 16.0] {
let error = disagreement(ell);
assert!(
error < 1.0e-12,
"the two gauges are one model: at ℓ={ell} they differ by {error:.3e}"
);
}
let (_, ell_hi) = constant_curvature_length_scale_bounds(data.view(), centers.view())
.expect("the fixture geometry has a range box");
let (d_min, d_max) = constant_curvature_evaluated_scale_span(data.view(), centers.view())
.expect("the fixture geometry has an evaluated scale span");
let mildest = f64::EPSILON * ell_hi / d_max;
let worst = f64::EPSILON * ell_hi / d_min;
let at_wall = disagreement(ell_hi);
assert!(
at_wall >= mildest && at_wall <= worst,
"at the box top ℓ={ell_hi:.4e} the `exp` gauge's error must be the cancellation \
`ε·ℓ/d`, bracketed by the evaluated span [{d_min:.4e}, {d_max:.4e}] as \
[{mildest:.3e}, {worst:.3e}]; measured {at_wall:.3e}"
);
assert!(
at_wall > 1.0e3 * disagreement(16.0),
"and it must GROW with the range: {at_wall:.3e} at the box top against \
{:.3e} mid-box",
disagreement(16.0)
);
}
#[test]
fn the_range_limit_is_the_geodesic_distance_kernel_2747() {
let (data, centers) = oracle_disk_design_centers();
let weights = Array1::<f64>::ones(centers.nrows());
let z = weighted_coefficient_sum_to_zero_transform(weights.view())
.expect("uniform sum-to-zero frame");
for kappa in [-1.1_f64, 0.0, 0.9] {
let manifold = ConstantCurvature::new(2, kappa);
let distances = |a: ArrayView2<'_, f64>, b: ArrayView2<'_, f64>| -> Array2<f64> {
let mut out = Array2::<f64>::zeros((a.nrows(), b.nrows()));
for i in 0..a.nrows() {
for j in 0..b.nrows() {
out[(i, j)] = manifold
.distance(a.row(i), b.row(j))
.expect("fixture disk is inside every probed chart");
}
}
out
};
let limit_design = distances(data.view(), centers.view())
.mapv(|d| -d)
.dot(&z);
let limit_penalty = symmetrize(
&z.t()
.dot(&distances(centers.view(), centers.view()).mapv(|d| -d))
.dot(&z),
);
let mut previous = f64::INFINITY;
for ell in [1.0e3_f64, 1.0e5, 1.0e7, 1.0e9] {
let spec = ConstantCurvatureBasisSpec {
center_strategy: CenterStrategy::UserProvided(centers.clone()),
kappa,
length_scale: ell,
..Default::default()
};
let built = build_constant_curvature_basis(data.view(), &spec).expect("build");
let design = built.design.to_dense();
let penalty = built.active_penalties[0].matrix.clone();
let gap = |a: &Array2<f64>, b: &Array2<f64>| -> f64 {
let mut num = 0.0_f64;
let mut den = 0.0_f64;
for (&x, &y) in a.iter().zip(b.iter()) {
num += (x - y) * (x - y);
den += y * y;
}
(num / den).sqrt()
};
let design_gap = gap(&design, &limit_design);
let penalty_gap = gap(&penalty, &limit_penalty);
assert!(
design_gap < previous / 50.0,
"κ={kappa}: the design must converge to the distance kernel; \
at ℓ={ell:.0e} the gap is {design_gap:.3e} against {previous:.3e} before"
);
previous = design_gap;
assert!(
penalty_gap < 1.0e-2,
"κ={kappa}: the penalty must converge too; gap {penalty_gap:.3e} at ℓ={ell:.0e}"
);
let (evals, _) = FaerEigh::eigh(&penalty, faer::Side::Lower).expect("penalty spectrum");
let smallest = evals.iter().cloned().fold(f64::INFINITY, f64::min);
let largest = evals.iter().cloned().fold(0.0_f64, f64::max);
assert!(
smallest > 1.0e-10 * largest,
"κ={kappa}: the restricted Gram must stay strictly PD at ℓ={ell:.0e}; \
spectrum spans [{smallest:.3e}, {largest:.3e}]"
);
}
assert!(
previous < 1.0e-8,
"κ={kappa}: at ℓ=1e9 the design must BE the distance kernel; gap {previous:.3e}"
);
}
}
}