use super::*;
use crate::assignment::{AssignmentMode, SaeAssignment};
use approx::assert_abs_diff_eq;
use gam_solve::inference::residual_factor::{ResidualFactorInput, StructuredResidualModel};
use gam_solve::rho_optimizer::{FixedPointCoordinateCertificate, OuterObjective};
use gam_terms::latent::LatentManifold;
use ndarray::{Array1, Array2};
use super::tests::{
PlantedCircleAssignmentMode, TestPeriodicEvaluator, periodic_basis, planted_circle_embedded,
planted_circle_seed_term, small_two_atom_periodic_term,
};
use std::sync::Arc;
#[test]
fn streaming_cache_outer_gradient_matches_dense_cache() {
let target = planted_circle_embedded(32, 4, 0.02);
let mut term0 = planted_circle_seed_term(target.view(), PlantedCircleAssignmentMode::Softmax).0;
term0.atoms[0].basis_second_jet = Some(Arc::new(
PeriodicHarmonicEvaluator::new(3).expect("periodic evaluator"),
));
let rho = SaeManifoldRho::new(0.0, 0.05_f64.ln(), vec![Array1::<f64>::zeros(1)]);
let inner_max_iter = 40;
let learning_rate = 1.0;
let ridge = 1.0e-6;
let mut dense = term0.clone();
let mut streaming = term0;
let (dense_cost, dense_loss, dense_cache) = dense
.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
inner_max_iter,
learning_rate,
ridge,
ridge,
)
.expect("dense cache criterion");
let (stream_cost, stream_loss, stream_cache) = streaming
.penalized_quasi_laplace_criterion_streaming_exact_with_cache(
target.view(),
&rho,
None,
inner_max_iter,
learning_rate,
ridge,
ridge,
)
.expect("streaming cache criterion");
assert_abs_diff_eq!(stream_cost, dense_cost, epsilon = 1.0e-8);
let smooth = rho.lambda_smooth_vec().unwrap();
let dense_solver = dense
.outer_gradient_arrow_solver(&dense_cache, &smooth)
.expect("dense outer-gradient solver");
let dense_grad = dense
.analytic_outer_rho_gradient_components(
target.view(),
&rho,
&dense_loss,
&dense_cache,
&dense_solver,
)
.expect("dense outer-gradient components")
.gradient();
let stream_solver = streaming
.outer_gradient_arrow_solver(&stream_cache, &smooth)
.expect("streaming outer-gradient solver");
let stream_grad = streaming
.analytic_outer_rho_gradient_components(
target.view(),
&rho,
&stream_loss,
&stream_cache,
&stream_solver,
)
.expect("streaming outer-gradient components")
.gradient();
assert_eq!(
dense_grad.len(),
stream_grad.len(),
"streaming outer gradient has a different ρ dimension than the dense one"
);
for (i, (d, s)) in dense_grad.iter().zip(stream_grad.iter()).enumerate() {
assert!(
d.is_finite() && s.is_finite(),
"outer-gradient component {i} must be finite (dense={d}, streaming={s})"
);
assert_abs_diff_eq!(d, s, epsilon = 1.0e-8);
}
let g2: f64 = dense_grad.iter().map(|v| v * v).sum();
assert!(
g2 > 0.0 && g2.is_finite(),
"the dense outer gradient must be non-trivial to make the parity check meaningful; ‖g‖²={g2}"
);
assert_abs_diff_eq!(stream_loss.total(), dense_loss.total(), epsilon = 1.0e-8);
}
fn lcg_uniform(s: &mut u64) -> f64 {
*s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((*s >> 11) as f64) / ((1u64 << 53) as f64)
}
fn lcg_normal(s: &mut u64) -> f64 {
let u1 = lcg_uniform(s).max(1e-12);
let u2 = lcg_uniform(s);
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
fn build_softmax_term(n: usize, p: usize, k: usize) -> SaeManifoldTerm {
let coord_cols: Vec<Array2<f64>> = (0..k)
.map(|i| {
Array2::<f64>::from_shape_fn((n, 1), |(r, _)| {
(0.03 + 0.11 * i as f64 + 0.017 * (i + 1) as f64 * r as f64).rem_euclid(1.0)
})
})
.collect();
let atoms: Vec<SaeManifoldAtom> = (0..k)
.map(|i| {
let (phi, jet) = periodic_basis(&coord_cols[i]);
let f = (i as f64) + 1.0;
let decoder = Array2::<f64>::from_shape_fn((3, p), |(m, c)| {
0.1 * f * ((m + 1) as f64) - 0.05 * (c as f64) + 0.02 * f
});
SaeManifoldAtom::new_with_provided_function_gram(
format!("atom{i}"),
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.expect("the fixture's basis, decoder and Gram blocks agree in dimension")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator))
})
.collect();
let manifolds = vec![LatentManifold::Circle { period: 1.0 }; k];
let logits =
Array2::<f64>::from_shape_fn((n, k), |(r, c)| 0.3 * (c as f64) - 0.1 * (r as f64) + 0.2);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coord_cols,
manifolds,
AssignmentMode::softmax(0.8),
)
.expect("the fixture's logits, coordinate blocks and manifolds agree in length");
SaeManifoldTerm::new(atoms, assignment)
.expect("the fixture's atoms and assignment describe the same latent blocks")
}
fn fit_structured_metric(n: usize, p: usize) -> gam_problem::RowMetric {
let lam = [1.0_f64, -0.7, 0.4, 0.9, -0.5];
let dscale = [0.10_f64, 0.55, 0.95, 0.30, 0.70];
let mut seed = 0x2026_00D5_1234_ABCDu64;
let mut residuals = Array2::<f64>::zeros((n, p));
let mut activity = Array1::<f64>::zeros(n);
for row in 0..n {
let common = lcg_normal(&mut seed);
activity[row] = 0.25 + (row as f64) / (n as f64);
let amp = activity[row].sqrt();
for i in 0..p {
residuals[[row, i]] = amp * lam[i % lam.len()] * common
+ dscale[i % dscale.len()] * lcg_normal(&mut seed);
}
}
let model = StructuredResidualModel::fit(ResidualFactorInput {
residuals: residuals.view(),
activity: activity.view(),
max_factor_rank: 2,
})
.expect("StructuredResidualModel::fit");
model.row_metric(n).expect("row_metric")
}
#[test]
fn wide_border_routes_to_streaming_with_complete_analytic_gradient_certificate() {
let (n, p, k, d_max) = (500usize, 128usize, 32usize, 1usize);
let total_basis = 2 * k; let border_dim = total_basis * p;
let budget = 2 * 1024 * 1024 * 1024usize; let host_available = 8 * 1024 * 1024 * 1024usize;
let chunk_window = SAE_CPU_L2_CACHE_BYTES * SAE_CHUNK_CACHE_MULTIPLE;
let plan = sae_streaming_plan_from_budget(
n,
total_basis,
k,
d_max,
border_dim,
budget,
chunk_window,
host_available,
);
assert!(
!plan.direct_admitted,
"the dense direct evidence peak ({} bytes) must exceed the 2 GiB budget so the \
criterion routes to streaming",
plan.estimated_direct_peak_bytes
);
assert!(
plan.matrix_free_admitted,
"the matrix-free plan ({} bytes) must be admitted so the fit has a route",
plan.estimated_matrix_free_peak_bytes
);
assert!(
plan.streaming,
"a non-direct-admitted plan must select streaming"
);
assert_eq!(
sae_outer_gradient_capability(),
Derivative::Analytic,
"matrix-free SAE must advertise the complete rational-value/single-adjoint gradient"
);
let dense_plan = sae_streaming_plan_from_budget(
n,
total_basis,
k,
d_max,
border_dim,
usize::MAX,
chunk_window,
usize::MAX,
);
assert!(dense_plan.direct_admitted);
assert_eq!(
sae_outer_gradient_capability(),
Derivative::Analytic,
"dense SAE retains its exact joint-Hessian IFT gradient"
);
let (_representative_term, _, representative_rho) = small_two_atom_periodic_term();
assert_eq!(
assignment_strength_gradient_coordinate(&representative_rho),
representative_rho.sparse_flat_index(),
"every active assignment strength must enter Hybrid-EFS's \
exact-gradient block; the outer-plan crossover decides whether that block \
is consumed, not whether the coordinate has an analytic root"
);
plan.admitted_or_error(n, border_dim, k)
.expect("matrix-free-admitted plan must not hard-error at the admission gate");
}
#[test]
fn production_objective_forced_streaming_value_gradient_matches_dense() {
let target = planted_circle_embedded(32, 4, 0.02);
let mut term = planted_circle_seed_term(target.view(), PlantedCircleAssignmentMode::Softmax).0;
term.atoms[0].basis_second_jet = Some(Arc::new(
PeriodicHarmonicEvaluator::new(3).expect("periodic evaluator"),
));
let seed_rho = SaeManifoldRho::new(0.0, 0.05_f64.ln(), vec![Array1::<f64>::zeros(1)]);
let mut dense = SaeManifoldOuterObjective::new(
term.clone(),
target.clone(),
None,
seed_rho.clone(),
40,
1.0,
1.0e-6,
1.0e-6,
);
let mut streaming =
SaeManifoldOuterObjective::new(term, target, None, seed_rho, 40, 1.0, 1.0e-6, 1.0e-6);
let rho_flat = dense.baseline_rho.to_flat();
let rho = streaming
.baseline_rho
.from_flat(rho_flat.view())
.expect("dense and streaming objectives must own the same typed rho layout");
assert_eq!(
rho_flat.len(),
2,
"K=1 Softmax has no assignment-strength coordinate"
);
let dense_eval =
OuterObjective::eval(&mut dense, &rho_flat).expect("dense production value+gradient");
let streaming_artifact = streaming
.evaluate_outer_criterion_route(&rho, false, false)
.expect("forced streaming production artifact");
let streaming_gradient = streaming
.analytic_gradient_for_outer_evaluation(&rho, &streaming_artifact)
.expect("forced streaming production gradient");
let streaming_eval = OuterEval {
cost: streaming_artifact.cost,
gradient: streaming_gradient,
hessian: HessianValue::Unavailable,
inner_beta_hint: Some(streaming.term.flatten_beta()),
};
assert!(dense_eval.cost.is_finite() && streaming_eval.cost.is_finite());
assert_eq!(dense_eval.gradient.len(), streaming_eval.gradient.len());
let dense_norm_sq = dense_eval.gradient.dot(&dense_eval.gradient);
assert!(
dense_norm_sq.is_finite() && dense_norm_sq > 1.0e-12,
"route parity must exercise a nonzero analytic gradient; norm^2={dense_norm_sq}"
);
assert_abs_diff_eq!(streaming_eval.cost, dense_eval.cost, epsilon = 1.0e-7);
for (coordinate, (&streamed, &direct)) in streaming_eval
.gradient
.iter()
.zip(dense_eval.gradient.iter())
.enumerate()
{
assert_abs_diff_eq!(streamed, direct, epsilon = 1.0e-6);
assert!(
streamed.is_finite(),
"streaming gradient coordinate {coordinate} is non-finite"
);
}
}
#[test]
fn evidence_assembly_row_fingerprint_sources_2515() {
let target = planted_circle_embedded(32, 4, 0.02);
let mut term = planted_circle_seed_term(target.view(), PlantedCircleAssignmentMode::Softmax).0;
term.atoms[0].basis_second_jet = Some(Arc::new(
PeriodicHarmonicEvaluator::new(3).expect("periodic evaluator"),
));
let rho = SaeManifoldRho::new(0.0, 0.05_f64.ln(), vec![Array1::<f64>::zeros(1)]);
term.refresh_decoder_repulsion_gate();
term.refresh_barrier_coactivation_gate();
term.refresh_amplitude_barrier_gate();
term.streaming_gates_frozen = true;
let direct = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("direct arrow-Schur assembly");
let (chunked, _chunk_term) = term
.assemble_full_matrix_free_evidence_system(target.view(), &rho, None, None)
.expect("matrix-free evidence assembly");
let direct_fp = direct.current_row_hessian_fingerprint();
let chunked_fp = chunked.current_row_hessian_fingerprint();
println!(
"[#2515 B3-SOURCE] gates frozen throughout: direct_assembly_row_fp={direct_fp} \
chunked_evidence_row_fp={chunked_fp} equal={}",
direct_fp == chunked_fp
);
let (frozen, _) = term
.assemble_full_matrix_free_evidence_system(target.view(), &rho, None, None)
.expect("frozen-gate evidence system");
term.streaming_gates_frozen = false;
let (refreshed, _) = term
.assemble_full_matrix_free_evidence_system(target.view(), &rho, None, None)
.expect("refreshed-gate evidence system");
println!(
"[#2515 B3-SOURCE] same assembler, gate state only: frozen_row_fp={} \
refreshed_row_fp={} equal={}",
frozen.row_hessian_fingerprint,
refreshed.row_hessian_fingerprint,
frozen.row_hessian_fingerprint == refreshed.row_hessian_fingerprint
);
assert_eq!(
direct.manifold_mode_fingerprint, chunked.manifold_mode_fingerprint,
"#2515: the two assemblers must agree on the MANIFOLD fingerprint — the \
stale-pair guard reports it alongside the row fingerprint precisely so a \
row mismatch can be read as an operator difference and not as a changed \
dictionary"
);
assert_eq!(
frozen.manifold_mode_fingerprint, refreshed.manifold_mode_fingerprint,
"#2515: the collapse-prevention gate state must not move the MANIFOLD \
fingerprint; it is a property of the atoms, not of the penalty gates"
);
assert_ne!(direct_fp, 0, "#2515: the direct assembly must publish a real row fingerprint");
assert_ne!(chunked_fp, 0, "#2515: the chunked assembly must publish a real row fingerprint");
}
#[test]
fn row_hessian_fingerprint_is_a_function_of_the_operator_2515() {
use super::kronecker::SaeKroneckerRows;
let target = planted_circle_embedded(32, 4, 0.02);
let mut term = planted_circle_seed_term(target.view(), PlantedCircleAssignmentMode::Softmax).0;
term.atoms[0].basis_second_jet = Some(Arc::new(
PeriodicHarmonicEvaluator::new(3).expect("periodic evaluator"),
));
let rho = SaeManifoldRho::new(0.0, 0.05_f64.ln(), vec![Array1::<f64>::zeros(1)]);
term.refresh_decoder_repulsion_gate();
term.refresh_barrier_coactivation_gate();
term.refresh_amplitude_barrier_gate();
term.streaming_gates_frozen = true;
let direct = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("direct arrow-Schur assembly");
let (chunked, _chunk_term) = term
.assemble_full_matrix_free_evidence_system(target.view(), &rho, None, None)
.expect("matrix-free evidence assembly");
assert!(
direct.htbeta_matvec.is_some() && chunked.htbeta_matvec.is_some(),
"#2515: both assemblies must install the matrix-free cross-block operator, \
or this test does not exercise the path the fingerprint defect lived on"
);
assert_eq!(
direct.current_row_hessian_fingerprint(),
chunked.current_row_hessian_fingerprint(),
"#2515: two assemblies of the same state must produce the same row-Hessian \
fingerprint. A guard that cannot be satisfied by an unchanged operator is \
not strict, it is inoperative — that was the defect, and the pointer \
address is what made it unsatisfiable."
);
let p = 4usize;
let a_phi: Arc<[Vec<(usize, f64)>]> =
Arc::from(vec![vec![(0usize, 1.0_f64), (3, -0.5)], vec![(1, 0.25)]].into_boxed_slice());
let local_jac: Arc<[Vec<f64>]> =
Arc::from(vec![vec![0.5_f64; p], vec![-0.25_f64; p]].into_boxed_slice());
let base = SaeKroneckerRows::new(p, Arc::clone(&a_phi), Arc::clone(&local_jac));
let base_fp = base.content_fingerprint();
let mut a_phi_moved = a_phi.to_vec();
a_phi_moved[0][0].1 += 1.0e-9;
let changed_a_phi = SaeKroneckerRows::new(
p,
Arc::from(a_phi_moved.into_boxed_slice()),
Arc::clone(&local_jac),
);
assert_ne!(
base_fp,
changed_a_phi.content_fingerprint(),
"#2515: a change to the sparse support weights must move the operator \
fingerprint (perturbation 1e-9 on one weight)"
);
let mut jac_moved = local_jac.to_vec();
jac_moved[1][0] += 1.0e-9;
let changed_jac = SaeKroneckerRows::new(
p,
Arc::clone(&a_phi),
Arc::from(jac_moved.into_boxed_slice()),
);
assert_ne!(
base_fp,
changed_jac.content_fingerprint(),
"#2515: a change to the local Jacobian must move the operator fingerprint"
);
let wider_jac: Arc<[Vec<f64>]> =
Arc::from(vec![vec![0.5_f64; p + 1], vec![-0.25_f64; p + 1]].into_boxed_slice());
let changed_p = SaeKroneckerRows::new(p + 1, Arc::clone(&a_phi), wider_jac);
assert_ne!(
base_fp,
changed_p.content_fingerprint(),
"#2515: a change to the decoder output dimension must move the operator \
fingerprint"
);
let identity_metric = gam_problem::RowMetric::euclidean(p, p)
.expect("a p-dimensional Euclidean row metric is constructible");
let changed_metric =
SaeKroneckerRows::new(p, a_phi, local_jac).with_output_metric(Some(identity_metric));
assert_ne!(
base_fp,
changed_metric.content_fingerprint(),
"#2515: installing an output metric must move the operator fingerprint"
);
}
#[test]
fn fixed_point_certificate_covers_non_ordered_beta_bernoulli_exact_gradient() {
let make_objective = || {
let (term, target, rho) = small_two_atom_periodic_term();
let rho_flat = rho.to_flat();
(
SaeManifoldOuterObjective::new(term, target, None, rho, 2, 0.25, 1.0e-4, 1.0e-4),
rho_flat,
)
};
let (mut iteration_objective, rho) = make_objective();
let iteration = iteration_objective
.eval_efs(&rho)
.expect("non-ordered Beta--Bernoulli EFS startup evaluation");
let gradient = iteration
.psi_gradient
.as_ref()
.expect("assignment strength must be the Hybrid-EFS gradient block")[0];
assert_eq!(
iteration.psi_indices.as_deref(),
Some(&[0][..]),
"the Hybrid-EFS gradient must map back to log_lambda_sparse"
);
assert!(gradient.is_finite(), "assignment gradient must be finite");
assert_abs_diff_eq!(
iteration.steps[0],
-gradient / gradient.abs().max(1.0),
epsilon = 1.0e-12
);
let (mut proof_objective, proof_rho) = make_objective();
let proof = proof_objective
.eval_fixed_point_certificate(&proof_rho)
.expect("fixed-point proof hook must evaluate");
let (mut exact_objective, exact_rho) = make_objective();
let exact = exact_objective
.eval(&exact_rho)
.expect("authoritative analytic gradient");
assert_eq!(proof.coordinates.len(), proof_rho.len());
match &proof.coordinates[0] {
FixedPointCoordinateCertificate::Covered { update, scale } => {
assert_abs_diff_eq!(*update, -exact.gradient[0], epsilon = 1.0e-12);
assert_eq!(*scale, 1.0);
}
FixedPointCoordinateCertificate::Uncovered { reason } => panic!(
"the exact assignment-strength derivative must certify this coordinate: {reason}"
),
}
}
#[test]
fn fixed_point_certificate_covers_ordered_beta_bernoulli_complete_gradient() {
let make_objective = || {
let (mut term, target, mut rho) = small_two_atom_periodic_term();
term.assignment.mode = AssignmentMode::ordered_beta_bernoulli(0.8, 1.0, true);
rho.log_lambda_sparse = 0.7_f64.ln();
let rho_flat = rho.to_flat();
(
SaeManifoldOuterObjective::new(term, target, None, rho, 2, 0.25, 1.0e-4, 1.0e-4),
rho_flat,
)
};
let (mut iteration_objective, rho) = make_objective();
let iteration = iteration_objective
.eval_efs(&rho)
.expect("ordered Beta--Bernoulli EFS startup evaluation");
assert!(
iteration.cost.is_finite(),
"the evaluation was refused before any gradient block was reached \
(cost={}); this is an infeasibility whose reason string was dropped by \
`infeasible_evaluation`, NOT a missing learnable-concentration gradient",
iteration.cost,
);
let gradient = iteration
.psi_gradient
.as_ref()
.expect("learnable concentration must use the complete gradient block")[0];
assert!(gradient.is_finite());
assert_eq!(iteration.psi_indices.as_deref(), Some(&[0][..]));
assert_abs_diff_eq!(
iteration.steps[0],
-gradient / gradient.abs().max(1.0),
epsilon = 1.0e-12
);
let (mut proof_objective, proof_rho) = make_objective();
let proof = proof_objective
.eval_fixed_point_certificate(&proof_rho)
.expect("ordered Beta--Bernoulli fixed-point proof hook must evaluate");
let (mut exact_objective, exact_rho) = make_objective();
let exact = exact_objective
.eval(&exact_rho)
.expect("authoritative analytic gradient");
match &proof.coordinates[0] {
FixedPointCoordinateCertificate::Covered { update, scale } => {
assert_abs_diff_eq!(*update, -exact.gradient[0], epsilon = 1.0e-12);
assert_eq!(*scale, 1.0);
}
FixedPointCoordinateCertificate::Uncovered { reason } => panic!(
"the complete ordered Beta--Bernoulli concentration derivative must certify this coordinate: {reason}"
),
}
}
#[test]
fn assignment_strength_trace_from_probes_matches_dense_softmax() {
let (n, p, k) = (24usize, 2usize, 2usize);
let term = build_softmax_term(n, p, k);
let rho = SaeManifoldRho::new(
0.7_f64.ln(),
0.8_f64.ln(),
vec![Array1::from_elem(1, 1.2_f64.ln()); k],
);
let fitted = term
.try_fitted_for_rho(&rho)
.expect("softmax positive-rank fixture reconstruction");
let target = Array2::<f64>::from_shape_fn((n, p), |(row, col)| {
fitted[[row, col]] + 1.0e-3 * ((row + 2 * col) as f64 * 0.17).sin()
});
let (system, _chunk_term) = term
.assemble_full_matrix_free_evidence_system(target.view(), &rho, None, None)
.expect("softmax matrix-free evidence system");
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_, _, cache) = solve_arrow_newton_step_with_options(&system, 0.0, 0.0, &options)
.expect("direct factorization");
assert!(
cache.deflated_row_directions.iter().all(Vec::is_empty),
"the probe identity is defined on the plain undeflated fixture"
);
let solver = DeflatedArrowSolver::plain(&cache);
let dense = term
.assignment_log_strength_hessian_trace(&rho, &cache, &solver)
.expect("dense assignment-strength trace");
let border_dim = cache.k;
let sqrt_dim = (border_dim as f64).sqrt();
let probes = (0..border_dim)
.map(|column| {
let mut probe = Array1::<f64>::zeros(border_dim);
probe[column] = sqrt_dim;
probe
})
.collect::<Vec<_>>();
let inverse_probes = probes
.iter()
.map(|probe| {
cache
.schur_inverse_apply(probe.view())
.expect("exact reduced-Schur inverse probe")
})
.collect::<Vec<_>>();
let matrix_free = term
.assignment_log_strength_hessian_trace_from_probes(
&rho,
&cache,
&probes,
&inverse_probes,
EvidenceOperator::Majorizer,
)
.expect("matrix-free assignment-strength trace");
assert!(
dense.abs() > 1.0e-12,
"fixture must excite a nonzero assignment-strength trace"
);
assert_abs_diff_eq!(matrix_free, dense, epsilon = 1.0e-9);
}
#[test]
fn complete_matrix_free_outer_gradient_matches_dense_softmax() {
let (n, p, k) = (24usize, 2usize, 2usize);
let term = build_softmax_term(n, p, k);
let rho = SaeManifoldRho::new(
0.7_f64.ln(),
0.8_f64.ln(),
vec![Array1::from_elem(1, 1.2_f64.ln()); k],
);
let fitted = term
.try_fitted_for_rho(&rho)
.expect("softmax positive-rank fixture reconstruction");
let target = Array2::<f64>::from_shape_fn((n, p), |(row, col)| {
fitted[[row, col]] + 1.0e-3 * ((row + 2 * col) as f64 * 0.17).sin()
});
let (system, _chunk_term) = term
.assemble_full_matrix_free_evidence_system(target.view(), &rho, None, None)
.expect("softmax matrix-free evidence system");
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_, _, cache) = solve_arrow_newton_step_with_options(&system, 0.0, 0.0, &options)
.expect("direct factorization");
assert!(
cache.deflated_row_directions.iter().all(Vec::is_empty),
"the probe identity is defined on the plain undeflated fixture"
);
let border_dim = cache.k;
let sqrt_dim = (border_dim as f64).sqrt();
let probes = (0..border_dim)
.map(|column| {
let mut probe = Array1::<f64>::zeros(border_dim);
probe[column] = sqrt_dim;
probe
})
.collect::<Vec<_>>();
let inverse_probes = probes
.iter()
.map(|probe| {
cache
.schur_inverse_apply(probe.view())
.expect("exact reduced-Schur inverse probe")
})
.collect::<Vec<_>>();
let plain_solver = DeflatedArrowSolver::plain(&cache);
let loss = term.loss(target.view(), &rho).expect("softmax loss");
let dense = term
.analytic_outer_rho_gradient_components(target.view(), &rho, &loss, &cache, &plain_solver)
.expect("dense complete outer gradient")
.gradient();
let matrix_free = term
.analytic_outer_rho_gradient_components_with_bundle(
target.view(),
&rho,
&loss,
&cache,
&plain_solver,
Some(BundleEvidenceGeometry {
operator: EvidenceOperator::Majorizer,
cache: &cache,
probes: &probes,
sinv: &inverse_probes,
}),
Some(&system),
)
.expect("matrix-free complete outer gradient")
.gradient();
assert_eq!(
dense.len(),
matrix_free.len(),
"matrix-free gradient has a different ρ dimension than the dense one"
);
let g2: f64 = dense.iter().map(|v| v * v).sum();
assert!(
g2 > 1.0e-10 && g2.is_finite(),
"the dense complete gradient must be non-trivial to make parity meaningful; ‖g‖²={g2}"
);
let mut max_abs = 0.0_f64;
for (i, (d, m)) in dense.iter().zip(matrix_free.iter()).enumerate() {
assert!(
d.is_finite() && m.is_finite(),
"gradient component {i} must be finite (dense={d}, matrix_free={m})"
);
max_abs = max_abs.max((d - m).abs());
assert_abs_diff_eq!(d, m, epsilon = 1.0e-8);
}
eprintln!(
"[complete_matrix_free_outer_gradient] max|dense-matrix_free| over {} coords = {:.3e}",
dense.len(),
max_abs
);
}
#[test]
fn whitened_streaming_criterion_completes() {
let (n, p, k) = (128usize, 16usize, 8usize);
let mut term = build_softmax_term(n, p, k);
let metric = fit_structured_metric(n, p);
assert!(
metric.whitens_likelihood(),
"the fitted structured-residual metric must whiten the likelihood"
);
term.set_row_metric(metric).unwrap();
let target = Array2::<f64>::from_shape_fn((n, p), |(r, c)| {
0.4 - 0.15 * (r as f64 / n as f64)
+ 0.25 * (c as f64 / p as f64)
+ 0.05 * (((r + c) % 7) as f64)
});
let rho = SaeManifoldRho::new(
-1.0_f64,
0.7_f64.ln(),
vec![Array1::<f64>::from_elem(1, 0.0); k],
);
let (cost, loss) = term
.penalized_quasi_laplace_criterion_streaming_exact(
target.view(),
&rho,
None,
2,
0.25,
1.0e-4,
1.0e-4,
)
.expect("whitened streaming criterion must complete, not hard-error");
assert!(
cost.is_finite(),
"streaming penalized quasi-Laplace criterion must be finite; got {cost}"
);
assert!(
loss.total().is_finite() && loss.data_fit.is_finite(),
"whitened loss components must be finite (data_fit={}, total={})",
loss.data_fit,
loss.total()
);
}