use super::*;
use super::invariant_tie_break::resolve_sorted_profile_tie;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) struct ConstraintNullspaceCacheKey {
pub(crate) centersrows: usize,
pub(crate) centers_cols: usize,
pub(crate) centers_hash: u64,
pub(crate) order: ConstraintNullspaceOrderKey,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) enum ConstraintNullspaceOrderKey {
Duchon(DuchonNullspaceOrder),
ThinPlate,
}
#[derive(Default, Clone, Debug)]
pub(crate) struct ConstraintNullspaceCache {
pub(crate) map: HashMap<ConstraintNullspaceCacheKey, Arc<Array2<f64>>>,
pub(crate) order: Vec<ConstraintNullspaceCacheKey>,
}
pub(crate) const CONSTRAINT_NULLSPACE_CACHE_MAX_ENTRIES: usize = 32;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) struct OwnedDataCacheKey {
pub(crate) rows: usize,
pub(crate) cols: usize,
pub(crate) ptr: usize,
pub(crate) stride0: isize,
pub(crate) stride1: isize,
}
#[derive(Debug)]
pub(crate) struct BasisCacheContext {
pub(crate) constraint_nullspace: ConstraintNullspaceCache,
pub(crate) owned_data: gam_runtime::resource::ByteLruCache<OwnedDataCacheKey, Arc<Array2<f64>>>,
}
impl BasisCacheContext {
pub(crate) fn with_policy(policy: &gam_runtime::resource::ResourcePolicy) -> Self {
Self {
constraint_nullspace: ConstraintNullspaceCache::default(),
owned_data: gam_runtime::resource::ByteLruCache::with_max_entries(
policy.max_owned_data_cache_bytes,
gam_runtime::resource::OWNED_DATA_CACHE_MAX_ENTRIES,
),
}
}
}
impl Default for BasisCacheContext {
fn default() -> Self {
Self::with_policy(&gam_runtime::resource::ResourcePolicy::default_library())
}
}
#[derive(Debug)]
pub struct BasisWorkspace {
pub(crate) cache: BasisCacheContext,
pub(crate) policy: gam_runtime::resource::ResourcePolicy,
}
impl BasisWorkspace {
pub fn new() -> Self {
Self::default()
}
pub fn with_policy(policy: gam_runtime::resource::ResourcePolicy) -> Self {
Self {
cache: BasisCacheContext::with_policy(&policy),
policy,
}
}
pub fn default_library() -> Self {
Self::with_policy(gam_runtime::resource::ResourcePolicy::default_library())
}
pub fn policy(&self) -> &gam_runtime::resource::ResourcePolicy {
&self.policy
}
}
impl Default for BasisWorkspace {
fn default() -> Self {
Self::default_library()
}
}
pub(crate) fn hash_arrayview2(values: ArrayView2<'_, f64>) -> u64 {
let mut hasher = DefaultHasher::new();
values.nrows().hash(&mut hasher);
values.ncols().hash(&mut hasher);
for v in values {
v.to_bits().hash(&mut hasher);
}
hasher.finish()
}
pub(crate) fn shared_owned_data_matrix(
data: ArrayView2<'_, f64>,
cache: &BasisCacheContext,
) -> Arc<Array2<f64>> {
let key = OwnedDataCacheKey {
rows: data.nrows(),
cols: data.ncols(),
ptr: data.as_ptr() as usize,
stride0: data.strides()[0],
stride1: data.strides()[1],
};
if let Some(hit) = cache.owned_data.get(&key) {
return hit;
}
let owned = Arc::new(data.to_owned());
if let Some(hit) = cache.owned_data.get(&key) {
return hit;
}
cache.owned_data.insert(key, owned.clone());
owned
}
#[inline]
pub(crate) fn shared_owned_data_matrix_from_view(data: ArrayView2<'_, f64>) -> Arc<Array2<f64>> {
Arc::new(data.to_owned())
}
#[inline]
pub(crate) fn shared_owned_centers_matrix_from_view(
centers: ArrayView2<'_, f64>,
) -> Arc<Array2<f64>> {
Arc::new(centers.to_owned())
}
pub(crate) fn kernel_constraint_nullspace(
centers: ArrayView2<'_, f64>,
order: DuchonNullspaceOrder,
cache: &mut BasisCacheContext,
) -> Result<Array2<f64>, BasisError> {
let effective_order = duchon_effective_nullspace_order(centers, order);
let centers_centered = mean_centered_centers(centers);
let centers = centers_centered.view();
let key = ConstraintNullspaceCacheKey {
centersrows: centers.nrows(),
centers_cols: centers.ncols(),
centers_hash: hash_arrayview2(centers),
order: ConstraintNullspaceOrderKey::Duchon(effective_order),
};
if let Some(hit) = cache.constraint_nullspace.map.get(&key) {
return Ok((**hit).clone());
}
let z = Arc::new(duchon_constraint_nullspace_of_centered(
centers,
order,
effective_order,
)?);
if let Some(hit) = cache.constraint_nullspace.map.get(&key) {
return Ok((**hit).clone());
}
cache.constraint_nullspace.map.insert(key, z.clone());
cache.constraint_nullspace.order.push(key);
while cache.constraint_nullspace.map.len() > CONSTRAINT_NULLSPACE_CACHE_MAX_ENTRIES {
if cache.constraint_nullspace.order.is_empty() {
break;
}
let oldkey = cache.constraint_nullspace.order.remove(0);
cache.constraint_nullspace.map.remove(&oldkey);
}
Ok((*z).clone())
}
fn mean_centered_centers(centers: ArrayView2<'_, f64>) -> Array2<f64> {
let k = centers.nrows();
let d = centers.ncols();
let center_mean: Vec<f64> = (0..d)
.map(|c| centers.column(c).sum() / (k.max(1) as f64))
.collect();
let mut centers_centered = centers.to_owned();
for c in 0..d {
let mu = center_mean[c];
centers_centered.column_mut(c).mapv_inplace(|v| v - mu);
}
centers_centered
}
fn duchon_constraint_nullspace_of_centered(
centers: ArrayView2<'_, f64>,
order: DuchonNullspaceOrder,
effective_order: DuchonNullspaceOrder,
) -> Result<Array2<f64>, BasisError> {
let degraded = effective_order != order;
let p_k = polynomial_block_from_order(centers, effective_order);
kernel_constraint_nullspace_from_matrix(p_k.view()).map_err(|err| {
if degraded {
BasisError::InvalidInput(format!(
"Duchon degraded from order={:?} to order={:?} due to insufficient centers ({} in dim={}); order={:?} construction then failed: {err}",
order,
effective_order,
centers.nrows(),
centers.ncols(),
effective_order,
))
} else {
err
}
})
}
pub fn duchon_kernel_constraint_nullspace(
centers: ArrayView2<'_, f64>,
order: DuchonNullspaceOrder,
) -> Result<Array2<f64>, BasisError> {
let effective_order = duchon_effective_nullspace_order(centers, order);
let centers_centered = mean_centered_centers(centers);
duchon_constraint_nullspace_of_centered(centers_centered.view(), order, effective_order)
}
pub(crate) fn thin_plate_kernel_constraint_nullspace(
centers: ArrayView2<'_, f64>,
cache: &mut BasisCacheContext,
) -> Result<Array2<f64>, BasisError> {
let key = ConstraintNullspaceCacheKey {
centersrows: centers.nrows(),
centers_cols: centers.ncols(),
centers_hash: hash_arrayview2(centers),
order: ConstraintNullspaceOrderKey::ThinPlate,
};
if let Some(hit) = cache.constraint_nullspace.map.get(&key) {
return Ok((**hit).clone());
}
let p_k = thin_plate_polynomial_block(centers);
if centers.nrows() < p_k.ncols() {
crate::bail_invalid_basis!(
"thin-plate spline requires at least {} centers to span the degree-{} polynomial null space in dimension {}; got {}",
p_k.ncols(),
thin_plate_polynomial_degree(centers.ncols()),
centers.ncols(),
centers.nrows()
);
}
let (z, rank) =
rrqr_nullspace_basis(&p_k, default_rrqr_rank_alpha()).map_err(BasisError::LinalgError)?;
if rank != p_k.ncols() {
crate::bail_invalid_basis!(
"thin-plate spline polynomial block is rank deficient at the selected centers: expected rank {}, got {}; choose geometrically independent centers for dimension {}",
p_k.ncols(),
rank,
centers.ncols()
);
}
let z = Arc::new(z);
if let Some(hit) = cache.constraint_nullspace.map.get(&key) {
return Ok((**hit).clone());
}
cache.constraint_nullspace.map.insert(key, z.clone());
cache.constraint_nullspace.order.push(key);
while cache.constraint_nullspace.map.len() > CONSTRAINT_NULLSPACE_CACHE_MAX_ENTRIES {
if cache.constraint_nullspace.order.is_empty() {
break;
}
let oldkey = cache.constraint_nullspace.order.remove(0);
cache.constraint_nullspace.map.remove(&oldkey);
}
Ok((*z).clone())
}
pub(crate) fn matern_identifiability_transform(
centers: ArrayView2<'_, f64>,
identifiability: &MaternIdentifiability,
) -> Result<Option<Array2<f64>>, BasisError> {
let k = centers.nrows();
match identifiability {
MaternIdentifiability::None => Ok(None),
MaternIdentifiability::CenterSumToZero => {
let q = Array2::<f64>::ones((k, 1));
Ok(Some(kernel_constraint_nullspace_from_matrix(q.view())?))
}
MaternIdentifiability::CenterLinearOrthogonal => {
let effective_order =
duchon_effective_nullspace_order(centers, DuchonNullspaceOrder::Linear);
let q = polynomial_block_from_order(centers, effective_order);
Ok(Some(kernel_constraint_nullspace_from_matrix(q.view())?))
}
MaternIdentifiability::FrozenTransform { transform, .. } => {
if transform.nrows() != k {
crate::bail_dim_basis!(
"frozen Matérn identifiability transform mismatch: centers={k}, transform rows={}",
transform.nrows()
);
}
Ok(Some(transform.clone()))
}
}
}
pub(crate) fn build_matern_operator_penalty_candidates(
centers: ArrayView2<'_, f64>,
length_scale: f64,
nu: MaternNu,
include_intercept: bool,
z_opt: Option<&Array2<f64>>,
aniso_log_scales: Option<&[f64]>,
) -> Result<Vec<PenaltyCandidate>, BasisError> {
let ops = build_matern_collocation_operator_matrices(
centers,
None,
length_scale,
nu,
include_intercept,
z_opt.map(|z| z.view()),
aniso_log_scales,
)?;
let matern_spec = DuchonOperatorPenaltySpec::matern_for_smoothness(nu, centers.ncols());
operator_penalty_candidates_from_collocation(&ops.d0, &ops.d1, &ops.d2, &matern_spec)
}
fn matrix_all_finite(m: &Array2<f64>) -> bool {
m.iter().all(|v| v.is_finite())
}
pub(crate) fn matern_center_function_gram(
embedded_kernel: &Array2<f64>,
include_intercept: bool,
full_transform: Option<&Array2<f64>>,
) -> Result<Array2<f64>, BasisError> {
if embedded_kernel.nrows() != embedded_kernel.ncols() {
crate::bail_dim_basis!("Matérn embedded kernel penalty must be square");
}
let total = embedded_kernel.nrows();
let k = total
.checked_sub(usize::from(include_intercept))
.ok_or_else(|| BasisError::InvalidInput("Matérn basis width underflow".to_string()))?;
if k == 0 {
crate::bail_invalid_basis!("Matérn function metric requires at least one center");
}
let mut center_design = Array2::<f64>::zeros((k, total));
center_design
.slice_mut(s![.., 0..k])
.assign(&embedded_kernel.slice(s![0..k, 0..k]));
if include_intercept {
center_design.column_mut(k).fill(1.0);
}
let center_design = match full_transform {
Some(transform) => fast_ab(¢er_design, transform),
None => center_design,
};
Ok(symmetrize_penalty(&fast_ata(¢er_design)))
}
pub(crate) fn matern_double_penalty_candidates(
primary: &Array2<f64>,
function_gram: &Array2<f64>,
include_intercept: bool,
) -> Result<Vec<PenaltyCandidate>, BasisError> {
if !matrix_all_finite(primary) {
crate::bail_invalid_basis!(
"Matérn double-penalty primary kernel Gram is non-finite; the projected \
kernel `Zᵀ K Z` could not be formed at this length scale (degenerate \
geometry). Widen the data spread, change the length scale, or drop the term."
);
}
if primary.dim() != function_gram.dim() || !matrix_all_finite(function_gram) {
crate::bail_invalid_basis!(
"Matérn center function Gram is non-finite or does not match the primary penalty"
);
}
let mut candidates = vec![normalize_penalty_candidate(
primary.clone(),
PenaltySource::Primary,
)?];
if include_intercept {
let p = primary.nrows();
let mut intercept_frame = Array2::<f64>::zeros((p, 1));
intercept_frame[[p - 1, 0]] = 1.0;
let shrinkage = function_space_subspace_shrinkage(&intercept_frame, function_gram)?;
candidates.push(normalize_penalty_candidate(
shrinkage,
PenaltySource::DoublePenaltyNullspace,
)?);
}
Ok(candidates)
}
pub(crate) fn build_matern_double_penalty_candidates(
spline: &MaternSplineBasis,
full_transform: Option<&Array2<f64>>,
) -> Result<Vec<PenaltyCandidate>, BasisError> {
let primary = project_penalty_matrix(&spline.penalty_kernel, full_transform);
let include_intercept = spline.num_polynomial_basis == 1;
let function_gram =
matern_center_function_gram(&spline.penalty_kernel, include_intercept, full_transform)?;
matern_double_penalty_candidates(&primary, &function_gram, include_intercept)
}
pub fn create_matern_spline_basiswithworkspace(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
length_scale: f64,
nu: MaternNu,
include_intercept: bool,
aniso_log_scales: Option<&[f64]>,
workspace: &mut BasisWorkspace,
) -> Result<MaternSplineBasis, BasisError> {
let n = data.nrows();
let d = data.ncols();
let k = centers.nrows();
let total_cols = k + usize::from(include_intercept);
let dense_bytes = dense_design_bytes(n, total_cols);
if dense_bytes > workspace.policy().max_single_materialization_bytes {
crate::bail_invalid_basis!(
"Matérn basis dense design exceeds resource policy: n={n}, p={total_cols}, dense={:.1} MiB, cap={:.1} MiB",
dense_bytes as f64 / (1024.0 * 1024.0),
workspace.policy().max_single_materialization_bytes as f64 / (1024.0 * 1024.0),
);
}
if d == 0 {
crate::bail_invalid_basis!("Matérn basis requires at least one covariate dimension");
}
if k == 0 {
crate::bail_invalid_basis!("Matérn basis requires at least one center");
}
if centers.ncols() != d {
crate::bail_dim_basis!(
"Matérn basis dimension mismatch: data has {d} columns, centers have {}",
centers.ncols()
);
}
if data.iter().any(|v| !v.is_finite()) || centers.iter().any(|v| !v.is_finite()) {
crate::bail_invalid_basis!("Matérn basis requires finite data and center values");
}
validate_matern_length_scale(length_scale)?;
if let Some(eta) = aniso_log_scales {
if eta.len() != d {
crate::bail_dim_basis!(
"aniso_log_scales length {} does not match data dimension {d}",
eta.len()
);
}
if eta.iter().any(|v| !v.is_finite()) {
crate::bail_invalid_basis!("aniso_log_scales must contain finite values");
}
}
let warn_bounds = if let Some(eta) = aniso_log_scales {
let y_centers = points_in_aniso_y_space(centers, eta);
pairwise_distance_bounds(y_centers.view())
} else {
pairwise_distance_bounds(centers)
};
if let Some((r_min, r_max)) = warn_bounds {
let kappa = 1.0 / length_scale.max(1e-300);
let kappa_lo = 1e-2 / r_max;
let kappa_hi = 1e2 / r_min;
if kappa < kappa_lo || kappa > kappa_hi {
log::debug!(
"Matérn κ={} is outside recommended range [{}, {}] derived from centers (r_min={}, r_max={}); kernel conditioning may degrade",
kappa,
kappa_lo,
kappa_hi,
r_min,
r_max
);
}
}
let mut kernel_block = Array2::<f64>::zeros((n, k));
let mut center_kernel = Array2::<f64>::zeros((k, k));
let axis_scales = aniso_log_scales.map(aniso_axis_scales);
let kernel_result: Result<(), BasisError> = kernel_block
.axis_iter_mut(Axis(0))
.into_par_iter()
.enumerate()
.try_for_each(|(i, mut row)| {
for j in 0..k {
let r = if let Some(scales) = axis_scales.as_deref() {
aniso_distance_rows_with_scales(data, i, centers, j, scales)
} else {
euclidean_distance_rows(data, i, centers, j)
};
row[j] = matern_kernel_from_distance(r, length_scale, nu)?;
}
Ok(())
});
kernel_result?;
fill_symmetric_from_row_kernel(&mut center_kernel, |i, j| {
let r = if let Some(scales) = axis_scales.as_deref() {
aniso_distance_rows_with_scales(centers, i, centers, j, scales)
} else {
euclidean_distance_rows(centers, i, centers, j)
};
matern_kernel_from_distance(r, length_scale, nu)
})?;
let mut basis = Array2::<f64>::zeros((n, total_cols));
basis.slice_mut(s![.., 0..k]).assign(&kernel_block);
if include_intercept {
basis.column_mut(k).fill(1.0);
}
let mut penalty_kernel = Array2::<f64>::zeros((total_cols, total_cols));
penalty_kernel
.slice_mut(s![0..k, 0..k])
.assign(¢er_kernel);
let function_gram = matern_center_function_gram(&penalty_kernel, include_intercept, None)?;
let penalty_ridge = if include_intercept {
let mut intercept_frame = Array2::<f64>::zeros((total_cols, 1));
intercept_frame[[total_cols - 1, 0]] = 1.0;
function_space_subspace_shrinkage(&intercept_frame, &function_gram)?
} else {
Array2::<f64>::zeros((total_cols, total_cols))
};
Ok(MaternSplineBasis {
basis,
penalty_kernel,
penalty_ridge,
num_kernel_basis: k,
num_polynomial_basis: usize::from(include_intercept),
dimension: d,
})
}
#[inline]
pub(crate) fn validate_lat_lon_matrix(
data: ArrayView2<'_, f64>,
context: &str,
radians: bool,
) -> Result<(), BasisError> {
if data.ncols() != 2 {
crate::bail_dim_basis!(
"{context} requires exactly two columns: latitude and longitude; got {}",
data.ncols()
);
}
if data.nrows() == 0 {
crate::bail_invalid_basis!("{context} requires at least one row");
}
let (lat_lo, lat_hi, unit) = if radians {
(
-std::f64::consts::FRAC_PI_2,
std::f64::consts::FRAC_PI_2,
"radians",
)
} else {
(-90.0, 90.0, "degrees")
};
for (i, row) in data.outer_iter().enumerate() {
let lat = row[0];
let lon = row[1];
if !lat.is_finite() || !lon.is_finite() {
crate::bail_invalid_basis!(
"{context} requires finite latitude/longitude; row {i} has ({lat}, {lon})"
);
}
if !(lat_lo..=lat_hi).contains(&lat) {
crate::bail_invalid_basis!(
"{context} latitude must be in [{lat_lo}, {lat_hi}] {unit}; row {i} has {lat}"
);
}
}
Ok(())
}
fn validate_spherical_wahba_gram_request(
penalty_order: usize,
kernel: SphereWahbaKernel,
) -> Result<(), BasisError> {
if !(1..=4).contains(&penalty_order) {
crate::bail_invalid_basis!(
"spherical spline penalty_order must be one of 1, 2, 3, 4; got {penalty_order}"
);
}
if matches!(kernel, SphereWahbaKernel::Sobolev) && penalty_order == 1 {
crate::bail_invalid_basis!(
"the m = 1 Sobolev sphere kernel is log-singular at coincident points, so its Gram \
diagonal does not exist and any finite value is a choice of resolution rather than a \
limit; use SobolevTruncated {{ lmax }} (the same kernel with the resolution stated, \
diagonal ~ ln(lmax)/2pi) or penalty_order >= 2, whose diagonals are finite closed \
forms (1/(4pi) at m = 2, (2*zeta3 - 2)/(4pi) at m = 3)"
);
}
Ok(())
}
pub fn spherical_wahba_kernel_matrix(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
penalty_order: usize,
radians: bool,
) -> Result<Array2<f64>, BasisError> {
spherical_wahba_kernel_matrix_with_kind(
data,
centers,
penalty_order,
radians,
SphereWahbaKernel::Sobolev,
)
}
pub fn spherical_wahba_kernel_matrix_with_kind(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
penalty_order: usize,
radians: bool,
kernel: SphereWahbaKernel,
) -> Result<Array2<f64>, BasisError> {
validate_spherical_wahba_gram_request(penalty_order, kernel)?;
validate_lat_lon_matrix(data, "spherical spline data", radians)?;
validate_lat_lon_matrix(centers, "spherical spline centers", radians)?;
if let Some(gpu_result) = crate::basis::sphere_gpu::try_build_truncated_kernel_matrix_gpu(
data,
centers,
penalty_order,
radians,
kernel,
) {
let gpu_matrix = gpu_result.map_err(|err| {
BasisError::InvalidInput(format!(
"spherical spline GPU truncated kernel was admitted but failed on device: {err}"
))
})?;
return Ok(gpu_matrix);
}
spherical_wahba_kernel_matrix_cpu_validated(data, centers, penalty_order, radians, kernel)
}
pub fn spherical_wahba_kernel_matrix_cpu(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
penalty_order: usize,
radians: bool,
kernel: SphereWahbaKernel,
) -> Result<Array2<f64>, BasisError> {
validate_spherical_wahba_gram_request(penalty_order, kernel)?;
validate_lat_lon_matrix(data, "spherical spline data", radians)?;
validate_lat_lon_matrix(centers, "spherical spline centers", radians)?;
spherical_wahba_kernel_matrix_cpu_validated(data, centers, penalty_order, radians, kernel)
}
fn spherical_wahba_kernel_matrix_cpu_validated(
data: ArrayView2<'_, f64>,
centers: ArrayView2<'_, f64>,
penalty_order: usize,
radians: bool,
kernel: SphereWahbaKernel,
) -> Result<Array2<f64>, BasisError> {
let n = data.nrows();
let k = centers.nrows();
let deg = if radians {
1.0
} else {
std::f64::consts::PI / 180.0
};
let mut sin_lat_c = Vec::<f64>::with_capacity(k);
let mut cos_lat_c = Vec::<f64>::with_capacity(k);
let mut sin_lon_c = Vec::<f64>::with_capacity(k);
let mut cos_lon_c = Vec::<f64>::with_capacity(k);
for c in centers.outer_iter() {
let trig = SphereTrig::from_radians(c[0] * deg, c[1] * deg);
sin_lat_c.push(trig.sin_lat);
cos_lat_c.push(trig.cos_lat);
sin_lon_c.push(trig.sin_lon);
cos_lon_c.push(trig.cos_lon);
}
let mut out = Array2::<f64>::zeros((n, k));
let err_flag = std::sync::atomic::AtomicBool::new(false);
out.axis_chunks_iter_mut(ndarray::Axis(0), 256)
.into_par_iter()
.enumerate()
.for_each(|(chunk_idx, mut block)| {
use wide::f64x4;
let row_offset = chunk_idx * 256;
let chunks = k / 4;
let tail = k % 4;
for (local_i, mut out_row) in block.outer_iter_mut().enumerate() {
let i = row_offset + local_i;
let row = SphereTrig::from_radians(data[(i, 0)] * deg, data[(i, 1)] * deg);
let row_v = SphereTrig {
sin_lat: f64x4::from(row.sin_lat),
cos_lat: f64x4::from(row.cos_lat),
sin_lon: f64x4::from(row.sin_lon),
cos_lon: f64x4::from(row.cos_lon),
};
for cidx in 0..chunks {
let base = cidx * 4;
let center_v = SphereTrig {
sin_lat: f64x4::from([
sin_lat_c[base],
sin_lat_c[base + 1],
sin_lat_c[base + 2],
sin_lat_c[base + 3],
]),
cos_lat: f64x4::from([
cos_lat_c[base],
cos_lat_c[base + 1],
cos_lat_c[base + 2],
cos_lat_c[base + 3],
]),
sin_lon: f64x4::from([
sin_lon_c[base],
sin_lon_c[base + 1],
sin_lon_c[base + 2],
sin_lon_c[base + 3],
]),
cos_lon: f64x4::from([
cos_lon_c[base],
cos_lon_c[base + 1],
cos_lon_c[base + 2],
cos_lon_c[base + 3],
]),
};
let (u, v) = half_angle_separation(row_v, center_v);
let vals = wahba_sphere_kernel_simd_kind(u, v, penalty_order, kernel);
let arr = vals.to_array();
for lane in 0..4 {
if !arr[lane].is_finite() {
err_flag.store(true, std::sync::atomic::Ordering::Relaxed);
return;
}
out_row[base + lane] = arr[lane];
}
}
let tail_start = chunks * 4;
for t in 0..tail {
let j = tail_start + t;
let center = SphereTrig {
sin_lat: sin_lat_c[j],
cos_lat: cos_lat_c[j],
sin_lon: sin_lon_c[j],
cos_lon: cos_lon_c[j],
};
let sep = half_angle_separation_scalar(row, center);
match wahba_sphere_kernel_kind(sep, penalty_order, kernel) {
Ok(v) => out_row[j] = v,
Err(_) => {
err_flag.store(true, std::sync::atomic::Ordering::Relaxed);
return;
}
}
}
}
});
if err_flag.load(std::sync::atomic::Ordering::Relaxed) {
crate::bail_invalid_basis!("spherical spline kernel produced a non-finite value");
}
Ok(out)
}
#[cfg(test)]
mod spherical_wahba_kernel_contract_2475_tests {
use super::*;
use ndarray::array;
fn assert_sobolev_m1_refusal(entry_point: &str, result: Result<Array2<f64>, BasisError>) {
let error = result.expect_err("untruncated Sobolev m=1 has no Gram diagonal");
let message = error.to_string();
assert!(
message.contains("log-singular") && message.contains("SobolevTruncated"),
"{entry_point} must identify both the mathematical defect and the explicit-resolution \
remedy; got: {message}"
);
}
#[test]
fn all_public_matrix_entry_points_refuse_untruncated_sobolev_m1() {
let data = array![[0.0, 0.0]];
let centers = array![[35.0, 70.0]];
assert_sobolev_m1_refusal(
"spherical_wahba_kernel_matrix",
spherical_wahba_kernel_matrix(data.view(), centers.view(), 1, false),
);
assert_sobolev_m1_refusal(
"spherical_wahba_kernel_matrix_with_kind",
spherical_wahba_kernel_matrix_with_kind(
data.view(),
centers.view(),
1,
false,
SphereWahbaKernel::Sobolev,
),
);
assert_sobolev_m1_refusal(
"spherical_wahba_kernel_matrix_cpu",
spherical_wahba_kernel_matrix_cpu(
data.view(),
centers.view(),
1,
false,
SphereWahbaKernel::Sobolev,
),
);
}
#[test]
fn explicit_resolution_and_finite_diagonal_m1_kernels_remain_available() {
let point = array![[0.0, 0.0]];
let pseudo = spherical_wahba_kernel_matrix_with_kind(
point.view(),
point.view(),
1,
false,
SphereWahbaKernel::Pseudo,
)
.expect("pseudo-Wahba m=1 has a finite analytic coincident-point value");
assert_eq!(
pseudo[(0, 0)],
1.0 / (4.0 * std::f64::consts::PI),
"the refusal must not absorb valid pseudo-Wahba m=1"
);
let truncated = spherical_wahba_kernel_matrix_with_kind(
point.view(),
point.view(),
1,
false,
SphereWahbaKernel::SobolevTruncated { lmax: 16 },
)
.expect("explicitly truncated Sobolev m=1 has a stated finite resolution");
assert!(
truncated[(0, 0)].is_finite(),
"a stated spectral resolution must produce a finite Gram diagonal"
);
spherical_wahba_kernel_matrix(point.view(), point.view(), 2, false)
.expect("untruncated Sobolev m=2 has a finite closed-form diagonal");
}
}
pub(crate) fn weighted_coefficient_sum_to_zero_transform(
weights: ArrayView1<'_, f64>,
) -> Result<Array2<f64>, BasisError> {
let k = weights.len();
if k < 2 {
return Err(BasisError::InsufficientColumnsForConstraint { found: k });
}
if weights.iter().any(|w| !w.is_finite() || *w < 0.0) {
crate::bail_invalid_basis!(
"sphere coefficient constraint weights must be finite and non-negative"
);
}
let norm = weights.iter().map(|w| w * w).sum::<f64>().sqrt();
if norm <= 0.0 {
crate::bail_invalid_basis!("sphere coefficient constraint weights cannot all be zero");
}
let c = Array2::from_shape_vec((k, 1), weights.iter().map(|w| *w / norm).collect())
.map_err(|e| BasisError::InvalidInput(format!("invalid sphere constraint weights: {e}")))?;
let (z, rank) =
rrqr_nullspace_basis(&c, default_rrqr_rank_alpha()).map_err(BasisError::LinalgError)?;
if rank >= k {
return Err(BasisError::ConstraintNullspaceCollapsed {
site: "weighted_coefficient_sum_to_zero_transform",
cross_rank: rank,
coeff_dim: k,
cross_frobenius: 1.0,
gram_spectrum: "not computed (structural rank collapse before Gram eigendecomposition)"
.to_string(),
});
}
Ok(z)
}
const SPHERICAL_CENTER_COINCIDENT_TOL: f64 = 1.0e-12;
#[inline]
fn spherical_center_dot(a: &[f64; 3], b: &[f64; 3]) -> f64 {
a[0] * b[0] + a[1] * b[1] + a[2] * b[2]
}
fn resolve_spherical_profile_tie<F>(
units: &[[f64; 3]],
tied: &[usize],
on_profile_builds: &mut F,
) -> Vec<usize>
where
F: FnMut(usize),
{
resolve_sorted_profile_tie(
units.len(),
tied,
|anchor, row| spherical_center_dot(&units[anchor], &units[row]),
on_profile_builds,
)
}
fn distinct_spherical_orbit(
units: &[[f64; 3]],
candidates: &[usize],
already_selected: &[usize],
) -> Vec<usize> {
let mut distinct = Vec::with_capacity(candidates.len());
'candidate: for &candidate in candidates {
for &selected in already_selected.iter().chain(distinct.iter()) {
if spherical_center_dot(&units[candidate], &units[selected])
>= 1.0 - SPHERICAL_CENTER_COINCIDENT_TOL
{
continue 'candidate;
}
}
distinct.push(candidate);
}
distinct
}
pub fn select_spherical_farthest_point_centers(
data: ArrayView2<'_, f64>,
num_centers: usize,
radians: bool,
) -> Result<Array2<f64>, BasisError> {
let chosen = select_spherical_farthest_point_center_rows(data, num_centers, radians)?;
log::debug!(
"spherical farthest-point centers: {} of {} rows, {} sorted dot profile(s) built",
chosen.rows.len(),
data.nrows(),
chosen.profile_builds
);
Ok(Array2::from_shape_fn((chosen.rows.len(), 2), |(r, c)| {
data[[chosen.rows[r], c]]
}))
}
pub(crate) struct SphericalCenterSelection {
pub(crate) rows: Vec<usize>,
pub(crate) profile_builds: usize,
}
fn select_spherical_farthest_point_center_rows(
data: ArrayView2<'_, f64>,
num_centers: usize,
radians: bool,
) -> Result<SphericalCenterSelection, BasisError> {
use rayon::prelude::*;
let mut profile_builds = 0usize;
validate_lat_lon_matrix(data, "spherical farthest-point centers", radians)?;
if num_centers == 0 {
crate::bail_invalid_basis!("spherical farthest-point center count must be positive");
}
let n = data.nrows();
if n < 2 {
return Err(BasisError::InsufficientColumnsForConstraint { found: n });
}
if num_centers > n {
crate::bail_invalid_basis!(
"requested {num_centers} spherical farthest-point centers but only {n} rows are available"
);
}
let to_rad = if radians {
1.0
} else {
std::f64::consts::PI / 180.0
};
let units: Vec<[f64; 3]> = (0..n)
.into_par_iter()
.map(|i| {
let lat = data[[i, 0]] * to_rad;
let lon = data[[i, 1]] * to_rad;
let cos_lat = lat.cos();
[cos_lat * lon.cos(), cos_lat * lon.sin(), lat.sin()]
})
.collect();
let mut sum = [0.0_f64; 3];
for (c, sum_c) in sum.iter_mut().enumerate() {
let mut col: Vec<f64> = units.par_iter().map(|u| u[c]).collect();
col.par_sort_by(|a, b| a.total_cmp(b));
*sum_c = col.iter().sum();
}
let dot_to_sum: Vec<f64> = units
.par_iter()
.map(|u| spherical_center_dot(u, &sum))
.collect();
let seed_key = dot_to_sum.par_iter().copied().reduce(
|| f64::NEG_INFINITY,
|a, b| if b.total_cmp(&a).is_gt() { b } else { a },
);
let seed_tied: Vec<usize> = (0..n)
.into_par_iter()
.filter(|&i| dot_to_sum[i].total_cmp(&seed_key).is_eq())
.collect();
let target = num_centers;
let seed_class = resolve_spherical_profile_tie(&units, &seed_tied, &mut |built: usize| {
profile_builds += built;
});
let seed_orbit = distinct_spherical_orbit(&units, &seed_class, &[]);
if seed_orbit.len() > target {
crate::bail_invalid_basis!(
"spherical farthest-point seed symmetry orbit has {} distinct directions, exceeding the requested center budget {target}; use a budget at least as large as the orbit or the harmonic sphere basis",
seed_orbit.len()
);
}
let mut selected = Vec::with_capacity(target);
let mut chosen = vec![false; n];
let mut max_dot = vec![f64::NEG_INFINITY; n];
for &i in &seed_class {
chosen[i] = true;
}
selected.extend(seed_orbit);
max_dot.par_iter_mut().enumerate().for_each(|(i, slot)| {
*slot = selected
.iter()
.map(|¢er| spherical_center_dot(&units[i], &units[center]))
.fold(f64::NEG_INFINITY, f64::max);
});
while selected.len() < target {
let cheap_key = (0..n)
.into_par_iter()
.filter(|&i| !chosen[i])
.map(|i| (max_dot[i], dot_to_sum[i]))
.reduce(
|| (f64::INFINITY, f64::INFINITY),
|a, b| {
if b.0.total_cmp(&a.0).then(b.1.total_cmp(&a.1)).is_lt() {
b
} else {
a
}
},
);
let cheap_tied: Vec<usize> = (0..n)
.into_par_iter()
.filter(|&i| {
!chosen[i]
&& max_dot[i].total_cmp(&cheap_key.0).is_eq()
&& dot_to_sum[i].total_cmp(&cheap_key.1).is_eq()
})
.collect();
if cheap_tied.is_empty() {
break;
}
if cheap_key.0 >= 1.0 - SPHERICAL_CENTER_COINCIDENT_TOL {
break;
}
let tied_class = resolve_spherical_profile_tie(&units, &cheap_tied, &mut |built: usize| {
profile_builds += built;
});
let orbit = distinct_spherical_orbit(&units, &tied_class, &selected);
let remaining = target - selected.len();
if orbit.len() > remaining {
crate::bail_invalid_basis!(
"spherical farthest-point tie class has {} distinct directions but only {remaining} of the exact {target}-center budget remain; choose a compatible center count or the harmonic sphere basis",
orbit.len(),
);
}
for &i in &tied_class {
chosen[i] = true;
}
if orbit.is_empty() {
continue;
}
selected.extend(orbit.iter().copied());
let chosen_ref = &chosen;
let orbit_ref = &orbit;
max_dot.par_iter_mut().enumerate().for_each(|(i, slot)| {
if chosen_ref[i] {
return;
}
for ¢er in orbit_ref {
let d = spherical_center_dot(&units[i], &units[center]);
if d > *slot {
*slot = d;
}
}
});
}
if selected.len() < target {
crate::bail_invalid_basis!(
"requested {target} distinct spherical farthest-point centers but the data contain only {} numerically distinct directions",
selected.len()
);
}
if selected.len() < 2 {
return Err(BasisError::InsufficientColumnsForConstraint {
found: selected.len(),
});
}
Ok(SphericalCenterSelection {
rows: selected,
profile_builds,
})
}
#[cfg(test)]
mod spherical_farthest_point_symmetry_tests {
use super::*;
use ndarray::{Array2, array};
fn permute_rows(data: &Array2<f64>, order: &[usize]) -> Array2<f64> {
Array2::from_shape_fn((order.len(), 2), |(row, col)| data[[order[row], col]])
}
fn sorted_center_rows(centers: &Array2<f64>) -> Vec<[f64; 2]> {
let mut rows: Vec<[f64; 2]> = centers.outer_iter().map(|row| [row[0], row[1]]).collect();
rows.sort_by(|a, b| a[0].total_cmp(&b[0]).then(a[1].total_cmp(&b[1])));
rows
}
#[test]
fn symmetric_tie_orbit_is_completed_under_every_row_permutation() {
let data = array![[0.0_f64, 0.0], [0.0, 0.0], [90.0, 0.0], [-90.0, 0.0]];
let permutations = [[0_usize, 1, 2, 3], [0, 1, 3, 2], [2, 0, 3, 1], [3, 1, 2, 0]];
let mut reference: Option<Vec<[f64; 2]>> = None;
for order in permutations {
let permuted = permute_rows(&data, &order);
let centers = select_spherical_farthest_point_centers(permuted.view(), 3, false)
.expect("the complete three-direction symmetry orbit is representable");
assert_eq!(
centers.nrows(),
3,
"the exact three-center target must contain the complete north/south tie class"
);
let center_set = sorted_center_rows(¢ers);
if let Some(expected) = &reference {
assert_eq!(
¢er_set, expected,
"selected physical center set changed under row permutation"
);
} else {
reference = Some(center_set);
}
}
}
#[test]
fn incomplete_nonseed_tie_class_is_refused() {
let data = array![[0.0_f64, 0.0], [0.0, 0.0], [90.0, 0.0], [-90.0, 0.0]];
let error = select_spherical_farthest_point_centers(data.view(), 2, false)
.expect_err("one remaining slot cannot split the north/south tie class");
assert!(
error
.to_string()
.contains("only 1 of the exact 2-center budget remain"),
"unexpected refusal: {error}"
);
}
#[test]
fn symmetry_orbit_larger_than_center_budget_is_refused() {
let antipodal = array![[90.0_f64, 0.0], [-90.0, 0.0]];
let error = select_spherical_farthest_point_centers(antipodal.view(), 1, false)
.expect_err("a two-direction seed orbit cannot fit a one-center budget");
assert!(
error.to_string().contains("symmetry orbit"),
"unexpected refusal: {error}"
);
}
fn latlon_grid(n_lat: usize, n_lon: usize) -> Array2<f64> {
Array2::from_shape_fn((n_lat * n_lon, 2), |(row, col)| {
let (i, j) = (row / n_lon, row % n_lon);
if col == 0 {
-85.0 + (170.0 * i as f64) / (n_lat.saturating_sub(1).max(1) as f64)
} else {
-180.0 + (360.0 * j as f64) / (n_lon.saturating_sub(1).max(1) as f64)
}
})
}
fn latlon_cloud(n: usize) -> Array2<f64> {
let mut state = 0x2545_F491_4F6C_DD1D_u64;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
(state >> 11) as f64 / (1u64 << 53) as f64
};
let draws: Vec<f64> = (0..2 * n).map(|_| next()).collect();
Array2::from_shape_fn((n, 2), |(row, col)| {
if col == 0 {
(1.0 - 2.0 * draws[2 * row]).asin().to_degrees()
} else {
360.0 * draws[2 * row + 1] - 180.0
}
})
}
#[test]
fn spherical_center_selection_costs_no_profile_without_an_exact_tie() {
for n in [2_000_usize, 8_000] {
for m in [40_usize, 200] {
let data = latlon_cloud(n);
let chosen = select_spherical_farthest_point_center_rows(data.view(), m, false)
.expect("an asymmetric cloud admits any center budget below n");
assert_eq!(chosen.rows.len(), m, "exact center budget (n={n}, m={m})");
assert_eq!(
chosen.profile_builds, 0,
"no row can tie the maximin key exactly on an asymmetric cloud, so the \
profile tie-break must never be built (n={n}, m={m})"
);
}
}
}
#[test]
fn spherical_center_profile_cost_does_not_scale_with_the_row_count() {
for (n_lat, n_lon) in [(40_usize, 40_usize), (160, 160)] {
for m in [40_usize, 200] {
let data = latlon_grid(n_lat, n_lon);
let n = data.nrows();
let chosen = select_spherical_farthest_point_center_rows(data.view(), m, false)
.expect("a lat/lon grid admits these center budgets");
assert_eq!(chosen.rows.len(), m, "exact center budget (n={n}, m={m})");
assert!(
chosen.profile_builds < m,
"profile-key builds must stay below one per selected center; got {} at \
n={n} m={m} (the replaced incumbent scan built at least {})",
chosen.profile_builds,
2 * m
);
}
}
}
#[test]
fn spherical_center_profile_key_is_still_built_where_it_decides_an_orbit() {
let data = array![[0.0_f64, 0.0], [0.0, 0.0], [90.0, 0.0], [-90.0, 0.0]];
let chosen = select_spherical_farthest_point_center_rows(data.view(), 3, false)
.expect("the complete three-direction symmetry orbit is representable");
assert!(
chosen.profile_builds > 0,
"the north/south orbit is only provable through the invariant profile key"
);
}
#[test]
fn polar_ring_selection_is_row_order_blind_at_every_budget() {
let ring: Vec<[f64; 2]> = [-30.0_f64, 30.0]
.into_iter()
.flat_map(|lat| (0..6).map(move |j| [lat, -180.0 + 60.0 * j as f64]))
.collect();
let data = Array2::from_shape_fn((ring.len(), 2), |(r, c)| ring[r][c]);
let n = data.nrows();
for budget in 2..=n {
let reference = select_spherical_farthest_point_centers(data.view(), budget, false)
.map(|centers| sorted_center_rows(¢ers));
for order in [
(0..n).rev().collect::<Vec<usize>>(),
(0..n).map(|i| (5 * i + 7) % n).collect::<Vec<usize>>(),
(0..n)
.step_by(5)
.chain((1..n).step_by(5))
.collect::<Vec<usize>>(),
] {
if order.len() != n {
continue;
}
let permuted = permute_rows(&data, &order);
let got = select_spherical_farthest_point_centers(permuted.view(), budget, false)
.map(|centers| sorted_center_rows(¢ers));
match (&reference, &got) {
(Ok(expected), Ok(actual)) => assert_eq!(
actual, expected,
"budget {budget}: selected physical directions changed under a row \
permutation of a symmetric ring"
),
(Err(a), Err(b)) => assert_eq!(
a.to_string(),
b.to_string(),
"budget {budget}: refusal changed under a row permutation"
),
_ => panic!(
"budget {budget}: row order decided whether the request was \
representable ({reference:?} vs {got:?})"
),
}
}
}
}
}
#[cfg(test)]
mod matern_function_metric_tests {
use super::*;
use ndarray::array;
#[test]
fn center_metric_null_ridge_is_covariant_and_targets_only_intercept_function() {
let center_kernel = array![[1.4, 0.3, 0.1], [0.3, 1.2, 0.2], [0.1, 0.2, 1.1]];
let mut embedded = Array2::<f64>::zeros((4, 4));
embedded.slice_mut(s![0..3, 0..3]).assign(¢er_kernel);
let gram =
matern_center_function_gram(&embedded, true, None).expect("raw center function Gram");
let base =
matern_double_penalty_candidates(&embedded, &gram, true).expect("raw candidates");
assert_eq!(base.len(), 2);
let raw_ridge = base[1].matrix.dense() * base[1].normalization_scale;
let intercept = array![[0.0], [0.0], [0.0], [1.0]];
let action_error = (&raw_ridge.dot(&intercept) - &gram.dot(&intercept))
.iter()
.map(|value| value.abs())
.fold(0.0_f64, f64::max);
assert!(
action_error < 2.0e-13,
"ridge must equal G on the structural intercept; error={action_error:.3e}"
);
let transform = array![
[0.2, 0.5, 0.0, 0.0],
[0.0, 3.0, -0.4, 0.0],
[0.0, 0.0, 1.7, 0.0],
[0.0, 0.0, 0.0, 2.5]
];
let primary_t = fast_atb(&transform, &fast_ab(&embedded, &transform));
let gram_t = matern_center_function_gram(&embedded, true, Some(&transform))
.expect("transformed center function Gram");
let transformed = matern_double_penalty_candidates(&primary_t, &gram_t, true)
.expect("transformed candidates");
let ridge_t = transformed[1].matrix.dense() * transformed[1].normalization_scale;
let expected = fast_atb(&transform, &fast_ab(&raw_ridge, &transform));
let covariance_error = (&ridge_t - &expected)
.iter()
.map(|value| value.abs())
.fold(0.0_f64, f64::max);
assert!(
covariance_error < 2.0e-12,
"Matérn function ridge changed under a basis chart; error={covariance_error:.3e}"
);
let no_intercept_gram = matern_center_function_gram(
¢er_kernel,
false,
Some(&transform.slice(s![0..3, 0..3]).to_owned()),
)
.expect("kernel-only Gram");
let kernel_only = matern_double_penalty_candidates(
&fast_atb(
&transform.slice(s![0..3, 0..3]).to_owned(),
&fast_ab(¢er_kernel, &transform.slice(s![0..3, 0..3]).to_owned()),
),
&no_intercept_gram,
false,
)
.expect("kernel-only candidates");
assert_eq!(kernel_only.len(), 1, "an SPD kernel has no null ridge");
}
}
pub fn auto_streaming_chunk_size_for_dense(n_rows: usize, n_basis_cols: usize) -> Option<usize> {
if n_rows == 0 || n_basis_cols == 0 {
return None;
}
const DENSE_THRESHOLD_BYTES: usize = 1024 * 1024 * 1024;
const TARGET_CHUNK_BYTES: usize = 256 * 1024 * 1024;
const MIN_CHUNK_ROWS: usize = 1024;
let dense_bytes = n_rows.saturating_mul(n_basis_cols).saturating_mul(8);
if dense_bytes <= DENSE_THRESHOLD_BYTES {
return None;
}
let row_bytes = n_basis_cols.saturating_mul(8).max(1);
let raw_chunk = TARGET_CHUNK_BYTES / row_bytes;
let clamped = raw_chunk.max(MIN_CHUNK_ROWS).min(n_rows);
Some(clamped)
}