#[derive(Debug, Clone, Copy)]
pub struct KernelParityReport {
pub cosine: f64,
pub max_abs_diff: f32,
pub max_abs_idx: usize,
pub n_diff_bits: usize,
pub total: usize,
}
impl KernelParityReport {
pub fn passes(&self, cos_min: f64, max_abs_max: f32) -> bool {
self.cosine.is_finite() && self.cosine >= cos_min && self.max_abs_diff <= max_abs_max
}
}
impl std::fmt::Display for KernelParityReport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"cosine={:.6}, max_abs={:.3e} at index {}, {}/{} elements bit-different",
self.cosine, self.max_abs_diff, self.max_abs_idx, self.n_diff_bits, self.total,
)
}
}
pub fn kernel_parity_report(a: &[f32], b: &[f32]) -> KernelParityReport {
assert_eq!(
a.len(),
b.len(),
"kernel_parity_report: length mismatch ({} vs {})",
a.len(),
b.len(),
);
let total = a.len();
let mut dot = 0.0_f64;
let mut na = 0.0_f64;
let mut nb = 0.0_f64;
let mut max_abs_diff = 0.0_f32;
let mut max_abs_idx = 0_usize;
let mut n_diff_bits = 0_usize;
for (i, (&x, &y)) in a.iter().zip(b.iter()).enumerate() {
let xf = x as f64;
let yf = y as f64;
dot += xf * yf;
na += xf * xf;
nb += yf * yf;
let d = (x - y).abs();
if d > max_abs_diff {
max_abs_diff = d;
max_abs_idx = i;
}
if x.to_bits() != y.to_bits() {
n_diff_bits += 1;
}
}
let cosine = if na > 0.0 && nb > 0.0 {
dot / (na.sqrt() * nb.sqrt())
} else {
f64::NAN
};
KernelParityReport {
cosine,
max_abs_diff,
max_abs_idx,
n_diff_bits,
total,
}
}
pub fn max_abs_diff(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len());
a.iter()
.zip(b.iter())
.map(|(&x, &y)| (x - y).abs())
.fold(0.0_f32, f32::max)
}
#[track_caller]
pub fn assert_kernel_equivalence(
a: &[f32],
b: &[f32],
cos_min: f64,
max_abs_max: f32,
label: &str,
) {
let report = kernel_parity_report(a, b);
assert!(
report.passes(cos_min, max_abs_max),
"kernel equivalence FAILED for `{label}`: {report} \
(required: cosine >= {cos_min}, max_abs_diff <= {max_abs_max:.3e})"
);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn identical_vectors_have_perfect_parity() {
let a = vec![1.0_f32, 2.0, -3.0, 4.5, 0.0];
let report = kernel_parity_report(&a, &a);
assert_eq!(report.cosine, 1.0);
assert_eq!(report.max_abs_diff, 0.0);
assert_eq!(report.n_diff_bits, 0);
assert_eq!(report.total, 5);
assert!(report.passes(0.9999, 1e-6));
}
#[test]
fn ulp_level_diff_passes_strict_threshold() {
let a = vec![1.0_f32, 2.0, 3.0];
let b = vec![1.0 + 1e-6, 2.0, 3.0 - 1e-6];
let report = kernel_parity_report(&a, &b);
assert!(report.cosine >= 0.9999);
assert!(report.max_abs_diff < 1e-4);
assert!(report.passes(0.9999, 1e-4));
assert!(report.n_diff_bits > 0);
}
#[test]
fn observed_real_drift_at_iter83_passes_tolerance() {
let n = 33700_usize;
let a: Vec<f32> = (0..n).map(|i| ((i as f32) * 1e-3).sin()).collect();
let mut b = a.clone();
for j in (0..407).map(|j| j * 80) {
if j < n {
let bits = b[j].to_bits().wrapping_add(((j % 7 + 1) * 4) as u32);
b[j] = f32::from_bits(bits);
}
}
let report = kernel_parity_report(&a, &b);
assert!(
report.cosine >= 0.9999,
"cosine {} should be ≥ 0.9999",
report.cosine
);
assert!(
report.max_abs_diff <= 1e-4,
"max_abs_diff {:e} should be ≤ 1e-4 for ULP-class perturbations",
report.max_abs_diff,
);
assert!(
report.passes(0.9999, 1e-4),
"iter83 ULP profile should pass tolerance: {report}"
);
}
#[test]
fn structural_divergence_fails_tolerance() {
let n = 262144_usize;
let a: Vec<f32> = (0..n).map(|i| ((i as f32) * 1e-3).sin() * 0.5).collect();
let b = a.clone();
let mut a_mut = a.clone();
for i in (0..n).step_by(5) {
a_mut[i] = 0.0; }
let report = kernel_parity_report(&a_mut, &b);
assert!(
!report.passes(0.9999, 1e-4),
"structural divergence (19% zeros vs finite) must fail tolerance: {report}"
);
}
#[test]
fn zero_norm_vector_returns_nan_cosine() {
let a = vec![0.0_f32; 10];
let b = vec![1.0_f32; 10];
let report = kernel_parity_report(&a, &b);
assert!(report.cosine.is_nan());
assert!(!report.passes(0.9999, 1e-4));
}
#[test]
fn max_abs_diff_standalone() {
let a = vec![1.0_f32, 2.0, 3.0];
let b = vec![1.0, 2.5, 2.5];
assert!((max_abs_diff(&a, &b) - 0.5).abs() < 1e-9);
}
#[test]
#[should_panic(expected = "length mismatch")]
fn length_mismatch_panics() {
let a = vec![1.0_f32; 10];
let b = vec![1.0_f32; 11];
let _ = kernel_parity_report(&a, &b);
}
}