From 1d876f2904e698617d92694d004c4ecd9ccbf980 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 22:42:35 +0000 Subject: [PATCH 1/3] =?UTF-8?q?simd:=20PowerSums::checked=5Fmerge=20?= =?UTF-8?q?=E2=80=94=20combine=20per-chunk=20power=20sums=20without=20wrap?= =?UTF-8?q?ping?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The keyed fold wraps, as the whole family does. A caller that tiles a population and combines the per-tile results needs the opposite: exact addition that refuses past the documented 2^32-rows-per-group bound. Each field is checked independently. Tests: two-chunk merge equals the one-pass fold at word, mid-word and end splits in both orders; three-way associativity; per-field refusal. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01X1YcYMRSFvfczXoP748wtB --- src/simd_masking_ops.rs | 82 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 82 insertions(+) diff --git a/src/simd_masking_ops.rs b/src/simd_masking_ops.rs index eeaabd5f..be9b6d71 100644 --- a/src/simd_masking_ops.rs +++ b/src/simd_masking_ops.rs @@ -2157,6 +2157,39 @@ fn power_sums_fold<'a>(values: &'a [i32], out: &'a mut [PowerSums]) -> impl FnMu } } +impl PowerSums { + /// Power sums of the union of two disjoint row sets — exact integer + /// addition, hence associative and commutative: chunks of a population + /// folded in any order and any grouping combine to the one-pass result. + /// + /// `None` if any field overflows (for `sum`, that is the union passing + /// the `2^32`-rows-per-group bound documented on [`PowerSums`]). The fold + /// itself wraps, as the whole keyed family does; a caller that combines + /// per-chunk results — a tiled executor, say — uses this to refuse + /// instead of wrapping. + /// + /// # Examples + /// + /// ``` + /// use ndarray::simd::PowerSums; + /// + /// let a = PowerSums { n: 2, sum: 5, sum_sq: 13 }; + /// let b = PowerSums { n: 1, sum: -4, sum_sq: 16 }; + /// assert_eq!(a.checked_merge(b), Some(PowerSums { n: 3, sum: 1, sum_sq: 29 })); + /// let full = PowerSums { n: 1, sum: i64::MAX, sum_sq: 0 }; + /// assert_eq!(full.checked_merge(PowerSums { n: 1, sum: 1, sum_sq: 0 }), None); + /// ``` + #[inline] + #[must_use] + pub fn checked_merge(self, other: Self) -> Option { + Some(Self { + n: self.n.checked_add(other.n)?, + sum: self.sum.checked_add(other.sum)?, + sum_sq: self.sum_sq.checked_add(other.sum_sq)?, + }) + } +} + /// Packs `index[i] < table.len() && table[index[i]] == v` into `out_words`, /// one bit per row `i < index.len()`, LSB-first — an equality predicate /// evaluated **through an index lane**, with no gathered mask and no @@ -8020,4 +8053,53 @@ mod group_family_tests { fn power_sums_refuse_a_short_mask_with_the_callers_name() { masked_group_power_sums_i32(&[], &[0], &[1], &mut [PowerSums::default()]); } + + /// Two chunks of one population, folded separately and merged, equal the + /// one-pass fold — at word boundaries, mid-word, and at both ends, in + /// both orders; three-way merges are associative. + #[test] + fn power_sums_chunks_merge_to_the_one_pass_result() { + let n = 1000; + let fx = fixture(n, 0xc4a1); + let mut whole = vec![PowerSums::default(); GROUPS]; + masked_group_power_sums_i32(&fx.mask, &fx.keys, &fx.values, &mut whole); + for split in [0usize, 1, 63, 64, 129, 640, 999, 1000] { + let (mut lo, mut hi) = (fx.mask.clone(), fx.mask.clone()); + for i in 0..n { + let (w, b) = (i / 64, i % 64); + if i < split { + hi[w] &= !(1u64 << b); + } else { + lo[w] &= !(1u64 << b); + } + } + let mut a = vec![PowerSums::default(); GROUPS]; + masked_group_power_sums_i32(&lo, &fx.keys, &fx.values, &mut a); + let mut b = vec![PowerSums::default(); GROUPS]; + masked_group_power_sums_i32(&hi, &fx.keys, &fx.values, &mut b); + for g in 0..GROUPS { + assert_eq!(a[g].checked_merge(b[g]), Some(whole[g]), "split={split} g={g}"); + assert_eq!(b[g].checked_merge(a[g]), Some(whole[g]), "commuted split={split} g={g}"); + } + } + let (x, y, z) = (whole[0], whole[1], whole[2]); + assert!(x.n > 0 && y.n > 0 && z.n > 0, "fixture leaves a group empty"); + assert_eq!( + x.checked_merge(y).and_then(|xy| xy.checked_merge(z)), + y.checked_merge(z).and_then(|yz| x.checked_merge(yz)), + ); + } + + /// The merge refuses — each field independently — rather than wrapping. + #[test] + fn power_sums_merge_refuses_to_wrap() { + let z = PowerSums::default(); + let n = PowerSums { n: u64::MAX, ..z }; + let s = PowerSums { sum: i64::MIN, ..z }; + let q = PowerSums { sum_sq: u128::MAX, ..z }; + assert_eq!(n.checked_merge(PowerSums { n: 1, ..z }), None); + assert_eq!(s.checked_merge(PowerSums { sum: -1, ..z }), None); + assert_eq!(q.checked_merge(PowerSums { sum_sq: 1, ..z }), None); + assert_eq!(n.checked_merge(z), Some(n), "identity merge must not refuse"); + } } From c56f9c2154a68db3d3d4612219e2b7b453d52845 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 22:45:07 +0000 Subject: [PATCH 2/3] =?UTF-8?q?simd:=20joint=20power=20sums=20fold=20(n,?= =?UTF-8?q?=20=CE=A3x,=20=CE=A3y,=20=CE=A3x=C2=B2,=20=CE=A3y=C2=B2,=20?= =?UTF-8?q?=CE=A3xy)=20on=20the=20visitor=20walker?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit CrossPowerSums is the bivariate sibling of PowerSums: same widths for the marginals (u128 square sums), plus a signed i128 Σxy. x()/y() return the marginals exactly as the univariate fold produces them; checked_merge combines chunks without wrapping. masked_group_cross_power_sums_i32 with _via / _pair share one fold, as the power-sum family does. Tests: i128/u128 longhand oracle on resident and VIA keys at three mask densities, marginals equal the univariate fold (resident and pair), chunk merges in both orders, exactness at the i32 extremes, per-field merge refusal. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01X1YcYMRSFvfczXoP748wtB --- src/simd.rs | 4 + src/simd_masking_ops.rs | 361 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 365 insertions(+) diff --git a/src/simd.rs b/src/simd.rs index c0a6e470..c4ad7f7f 100644 --- a/src/simd.rs +++ b/src/simd.rs @@ -832,6 +832,9 @@ pub use crate::simd_masking_ops::{ masked_group_count_u32, masked_group_count_u32_pair, masked_group_count_u32_via, + masked_group_cross_power_sums_i32, + masked_group_cross_power_sums_i32_pair, + masked_group_cross_power_sums_i32_via, masked_group_max_i32, masked_group_max_i32_pair, masked_group_max_i32_via, @@ -864,6 +867,7 @@ pub use crate::simd_masking_ops::{ ternary_match_u32_to_mask_under, ternary_match_u64_to_mask, ternary_match_u64_to_mask_under, + CrossPowerSums, KeyRunCarry, MortonDir, PowerSums, diff --git a/src/simd_masking_ops.rs b/src/simd_masking_ops.rs index be9b6d71..b8737873 100644 --- a/src/simd_masking_ops.rs +++ b/src/simd_masking_ops.rs @@ -2190,6 +2190,207 @@ impl PowerSums { } } +/// Joint power sums of a pair of `i32` lanes, per group: `n`, `Σx`, `Σy`, +/// `Σx²`, `Σy²`, `Σxy` — the degree-≤2 sufficient statistics that +/// covariance, Pearson's r and a simple least-squares line are projected +/// from. The bivariate sibling of [`PowerSums`]; each marginal is exactly +/// the [`PowerSums`] the univariate fold would produce ([`x`](Self::x), +/// [`y`](Self::y)). +/// +/// # Widths +/// +/// `n`, `sum_x`, `sum_y`, `sum_x_sq`, `sum_y_sq` carry the [`PowerSums`] +/// widths and bound. `sum_xy: i128` is the one signed second-order sum: a +/// product lies in `[i32::MIN·i32::MAX, i32::MIN²] = [-(2^62 - 2^31), 2^62]`, +/// so it stays exact below `2^65` rows. Folded with wrapping adds, like the +/// rest of the keyed family; combine chunks with +/// [`checked_merge`](Self::checked_merge). +/// +/// # Examples +/// +/// ``` +/// use ndarray::simd::{CrossPowerSums, PowerSums}; +/// +/// let c = CrossPowerSums { n: 2, sum_x: 3, sum_y: 1, sum_x_sq: 5, sum_y_sq: 5, sum_xy: 0 }; +/// assert_eq!(c.x(), PowerSums { n: 2, sum: 3, sum_sq: 5 }); +/// ``` +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash)] +pub struct CrossPowerSums { + /// `Σ1` — the number of row pairs folded in. + pub n: u64, + /// `Σx`. + pub sum_x: i64, + /// `Σy`. + pub sum_y: i64, + /// `Σx²`. + pub sum_x_sq: u128, + /// `Σy²`. + pub sum_y_sq: u128, + /// `Σxy`. + pub sum_xy: i128, +} + +impl CrossPowerSums { + /// The `x` marginal — exactly what [`masked_group_power_sums_i32`] folds + /// from the `x` lane over the same rows. + #[inline] + #[must_use] + pub const fn x(&self) -> PowerSums { + PowerSums { + n: self.n, + sum: self.sum_x, + sum_sq: self.sum_x_sq, + } + } + + /// The `y` marginal. + #[inline] + #[must_use] + pub const fn y(&self) -> PowerSums { + PowerSums { + n: self.n, + sum: self.sum_y, + sum_sq: self.sum_y_sq, + } + } + + /// Joint power sums of the union of two disjoint row sets; `None` if any + /// field overflows. Same contract as [`PowerSums::checked_merge`]. + /// + /// # Examples + /// + /// ``` + /// use ndarray::simd::CrossPowerSums; + /// + /// let a = CrossPowerSums { n: 1, sum_x: 2, sum_y: -1, sum_x_sq: 4, sum_y_sq: 1, sum_xy: -2 }; + /// let m = a.checked_merge(a).unwrap(); + /// assert_eq!((m.n, m.sum_xy), (2, -4)); + /// ``` + #[inline] + #[must_use] + pub fn checked_merge(self, other: Self) -> Option { + Some(Self { + n: self.n.checked_add(other.n)?, + sum_x: self.sum_x.checked_add(other.sum_x)?, + sum_y: self.sum_y.checked_add(other.sum_y)?, + sum_x_sq: self.sum_x_sq.checked_add(other.sum_x_sq)?, + sum_y_sq: self.sum_y_sq.checked_add(other.sum_y_sq)?, + sum_xy: self.sum_xy.checked_add(other.sum_xy)?, + }) + } +} + +/// Keyed joint power sums: for every row `i` selected by `mask_words` +/// (`i < xs.len()`), folds the pair `(xs[i], ys[i])` into `out[keys[i]]` — +/// provided `keys[i] < out.len()`. +/// +/// `SELECT key, COUNT(*), SUM(x), SUM(y), SUM(x*x), SUM(y*y), SUM(x*y) … +/// GROUP BY key` in one pass. Same contract as +/// [`masked_group_power_sums_i32`] in every other respect: `out` is +/// accumulated into (start it at [`CrossPowerSums::default`]), a key past +/// `out.len()` is dropped, and the final mask word is clamped to `xs.len()`. +/// +/// # Panics +/// +/// Panics if `keys`, `xs` and `ys` differ in length, or if `mask_words.len() +/// < xs.len().div_ceil(64)`. +/// +/// # Examples +/// +/// ``` +/// use ndarray::simd::{masked_group_cross_power_sums_i32, CrossPowerSums}; +/// +/// let mask = [0b1011u64]; // rows 0, 1, 3 +/// let keys = [0u32, 1, 0, 0]; +/// let xs = [2i32, 7, 100, -1]; +/// let ys = [3i32, 1, 100, 4]; +/// let mut out = [CrossPowerSums::default(); 2]; +/// masked_group_cross_power_sums_i32(&mask, &keys, &xs, &ys, &mut out); +/// assert_eq!((out[0].n, out[0].sum_xy), (2, 2 * 3 + -4)); +/// ``` +#[inline] +pub fn masked_group_cross_power_sums_i32( + mask_words: &[u64], keys: &[u32], xs: &[i32], ys: &[i32], out: &mut [CrossPowerSums], +) { + assert_eq!(keys.len(), xs.len(), "masked_group_cross_power_sums_i32: keys/xs length mismatch"); + assert_eq!(xs.len(), ys.len(), "masked_group_cross_power_sums_i32: xs/ys length mismatch"); + group_walk( + "masked_group_cross_power_sums_i32", + mask_words, + xs.len(), + GroupKeyAddr::Resident(keys), + out.len(), + cross_power_sums_fold(xs, ys, out), + ); +} + +/// [`masked_group_cross_power_sums_i32`] with the group read through an +/// index lane: row `i` folds into `out[table[index[i]]]`, dropped when +/// `index[i]` is past `table` or the group is past `out`. +/// +/// # Panics +/// +/// Panics if `index`, `xs` and `ys` differ in length, or if +/// `mask_words.len() < xs.len().div_ceil(64)`. +#[inline] +pub fn masked_group_cross_power_sums_i32_via( + mask_words: &[u64], index: &[u32], table: &[u32], xs: &[i32], ys: &[i32], out: &mut [CrossPowerSums], +) { + assert_eq!(index.len(), xs.len(), "masked_group_cross_power_sums_i32_via: index/xs length mismatch"); + assert_eq!(xs.len(), ys.len(), "masked_group_cross_power_sums_i32_via: xs/ys length mismatch"); + group_walk( + "masked_group_cross_power_sums_i32_via", + mask_words, + xs.len(), + GroupKeyAddr::Via { index, table }, + out.len(), + cross_power_sums_fold(xs, ys, out), + ); +} + +/// [`masked_group_cross_power_sums_i32`] with a composite key +/// `hi[i] * stride + lo[i]`, under the same drop rules as the rest of the +/// `_pair` family. +/// +/// # Panics +/// +/// Panics if `hi`, `lo`, `xs` and `ys` differ in length, or if +/// `mask_words.len() < xs.len().div_ceil(64)`. +#[inline] +pub fn masked_group_cross_power_sums_i32_pair( + mask_words: &[u64], hi: &[u32], lo: &[u32], stride: u32, xs: &[i32], ys: &[i32], out: &mut [CrossPowerSums], +) { + assert_eq!(hi.len(), lo.len(), "masked_group_cross_power_sums_i32_pair: hi.len() != lo.len()"); + assert_eq!(hi.len(), xs.len(), "masked_group_cross_power_sums_i32_pair: hi/xs length mismatch"); + assert_eq!(xs.len(), ys.len(), "masked_group_cross_power_sums_i32_pair: xs/ys length mismatch"); + group_walk( + "masked_group_cross_power_sums_i32_pair", + mask_words, + hi.len(), + GroupKeyAddr::Pair { hi, lo, stride }, + out.len(), + cross_power_sums_fold(xs, ys, out), + ); +} + +/// The one joint power-sum fold, shared by all three key addresses. +#[inline(always)] +fn cross_power_sums_fold<'a>( + xs: &'a [i32], ys: &'a [i32], out: &'a mut [CrossPowerSums], +) -> impl FnMut(usize, usize) + 'a { + move |k, i| { + let (x, y) = (xs[i] as i64, ys[i] as i64); + let p = &mut out[k]; + p.n = p.n.wrapping_add(1); + p.sum_x = p.sum_x.wrapping_add(x); + p.sum_y = p.sum_y.wrapping_add(y); + // Each of these products fits an i64: |x·y| <= 2^62. + p.sum_x_sq = p.sum_x_sq.wrapping_add((x * x) as u128); + p.sum_y_sq = p.sum_y_sq.wrapping_add((y * y) as u128); + p.sum_xy = p.sum_xy.wrapping_add((x * y) as i128); + } +} + /// Packs `index[i] < table.len() && table[index[i]] == v` into `out_words`, /// one bit per row `i < index.len()`, LSB-first — an equality predicate /// evaluated **through an index lane**, with no gathered mask and no @@ -8102,4 +8303,164 @@ mod group_family_tests { assert_eq!(q.checked_merge(PowerSums { sum_sq: 1, ..z }), None); assert_eq!(n.checked_merge(z), Some(n), "identity merge must not refuse"); } + + // ── the joint (cross) power-sum fold ────────────────────────────────── + + /// A second value lane, independent of the fixture's, carrying the i32 + /// extremes on a different period so MIN·MIN, MIN·MAX and MAX·MIN pairs + /// all occur. + fn second_lane(n: usize, seed: u64) -> Vec { + let mut s = seed; + (0..n) + .map(|i| match i % 13 { + 5 => i32::MIN, + 6 => i32::MAX, + _ => (lcg(&mut s) % 4001) as i32 - 2000, + }) + .collect() + } + + type Wide6 = (u64, i128, i128, u128, u128, i128); + + /// The independent joint oracle: every row longhand in i128/u128. + fn cross_reference( + mask: &[u64], xs: &[i32], ys: &[i32], groups: usize, key: impl Fn(usize) -> Option, + ) -> Vec { + let mut out = vec![(0u64, 0i128, 0i128, 0u128, 0u128, 0i128); groups]; + for i in 0..xs.len() { + if mask[i / 64] >> (i % 64) & 1 == 1 { + if let Some(k) = key(i) { + let (x, y) = (xs[i] as i128, ys[i] as i128); + let o = &mut out[k]; + o.0 += 1; + o.1 += x; + o.2 += y; + o.3 += (x * x) as u128; + o.4 += (y * y) as u128; + o.5 += x * y; + } + } + } + out + } + + fn cross_wide(got: &[CrossPowerSums]) -> Vec { + got.iter() + .map(|c| (c.n, c.sum_x as i128, c.sum_y as i128, c.sum_x_sq, c.sum_y_sq, c.sum_xy)) + .collect() + } + + /// Resident and VIA joint sums equal the oracle at every length and three + /// densities, and each marginal equals the univariate fold. + #[test] + fn cross_power_sums_match_the_wide_oracle_and_their_marginals() { + for &n in LENS { + let fx = fixture(n, 0xc055 ^ n as u64); + let ys = second_lane(n, 0x9e37 ^ n as u64); + for (density, mask) in masks(&fx) { + let mut got = vec![CrossPowerSums::default(); GROUPS]; + masked_group_cross_power_sums_i32(&mask, &fx.keys, &fx.values, &ys, &mut got); + let want = cross_reference(&mask, &fx.values, &ys, GROUPS, |i| key_resident(&fx, i)); + assert_eq!(cross_wide(&got), want, "resident {density} n={n}"); + + let mut via = vec![CrossPowerSums::default(); GROUPS]; + masked_group_cross_power_sums_i32_via(&mask, &fx.index, &fx.table, &fx.values, &ys, &mut via); + let want = cross_reference(&mask, &fx.values, &ys, GROUPS, |i| key_via(&fx, i)); + assert_eq!(cross_wide(&via), want, "via {density} n={n}"); + + let mut mx = vec![PowerSums::default(); GROUPS]; + masked_group_power_sums_i32(&mask, &fx.keys, &fx.values, &mut mx); + let mut my = vec![PowerSums::default(); GROUPS]; + masked_group_power_sums_i32(&mask, &fx.keys, &ys, &mut my); + for g in 0..GROUPS { + assert_eq!((got[g].x(), got[g].y()), (mx[g], my[g]), "marginals {density} n={n} g={g}"); + } + } + } + // Anti-vacuity: an extreme product must actually reach a group. + let fx = fixture(1000, 0xc055 ^ 1000); + let ys = second_lane(1000, 0x9e37 ^ 1000); + let extreme = + |i: usize| fx.values[i].unsigned_abs() >= i32::MAX as u32 && ys[i].unsigned_abs() >= i32::MAX as u32; + assert!((0..1000).any(|i| selected(&fx, i) && key_resident(&fx, i).is_some() && extreme(i))); + } + + /// The pair address: each marginal equals the univariate pair fold. + #[test] + fn cross_power_sums_pair_agree_with_the_univariate_pair_fold() { + let fx = fixture(1000, 0xba1e); + let ys = second_lane(1000, 0xba1f); + let hi: Vec = fx.keys.iter().map(|k| k % 3).collect(); + let lo: Vec = fx.index.iter().map(|x| x % 4).collect(); + let mut c = [CrossPowerSums::default(); 9]; + masked_group_cross_power_sums_i32_pair(&fx.mask, &hi, &lo, 3, &fx.values, &ys, &mut c); + let mut mx = [PowerSums::default(); 9]; + masked_group_power_sums_i32_pair(&fx.mask, &hi, &lo, 3, &fx.values, &mut mx); + let mut my = [PowerSums::default(); 9]; + masked_group_power_sums_i32_pair(&fx.mask, &hi, &lo, 3, &ys, &mut my); + assert!(c.iter().any(|g| g.n > 1), "fixture reaches no pair group twice"); + for g in 0..9 { + assert_eq!((c[g].x(), c[g].y()), (mx[g], my[g]), "g={g}"); + } + } + + /// Chunks merge to the one-pass result in both orders. + #[test] + fn cross_power_sums_chunks_merge_to_the_one_pass_result() { + let n = 1000; + let fx = fixture(n, 0xc4a2); + let ys = second_lane(n, 0xc4a3); + let mut whole = vec![CrossPowerSums::default(); GROUPS]; + masked_group_cross_power_sums_i32(&fx.mask, &fx.keys, &fx.values, &ys, &mut whole); + for split in [0usize, 1, 63, 64, 129, 640, 999, 1000] { + let (mut lo, mut hi) = (fx.mask.clone(), fx.mask.clone()); + for i in 0..n { + let (w, b) = (i / 64, i % 64); + if i < split { + hi[w] &= !(1u64 << b); + } else { + lo[w] &= !(1u64 << b); + } + } + let mut a = vec![CrossPowerSums::default(); GROUPS]; + masked_group_cross_power_sums_i32(&lo, &fx.keys, &fx.values, &ys, &mut a); + let mut b = vec![CrossPowerSums::default(); GROUPS]; + masked_group_cross_power_sums_i32(&hi, &fx.keys, &fx.values, &ys, &mut b); + for g in 0..GROUPS { + assert_eq!(a[g].checked_merge(b[g]), Some(whole[g]), "split={split} g={g}"); + assert_eq!(b[g].checked_merge(a[g]), Some(whole[g]), "commuted split={split} g={g}"); + } + } + } + + /// Exact at the i32 extremes, and the merge refuses per field. + #[test] + fn cross_power_sums_are_exact_at_the_extremes_and_merge_refuses_to_wrap() { + let mut m = [CrossPowerSums::default(); 1]; + let xs = [i32::MIN, i32::MIN, i32::MAX]; + let ys = [i32::MIN, i32::MAX, i32::MIN]; + masked_group_cross_power_sums_i32(&[0b111], &[0, 0, 0], &xs, &ys, &mut m); + let (lo, hi) = (i32::MIN as i128, i32::MAX as i128); + assert_eq!(m[0].sum_xy, lo * lo + lo * hi + hi * lo); + assert_eq!(m[0].sum_x_sq, (lo * lo + lo * lo + hi * hi) as u128); + assert_eq!(m[0].sum_y_sq, (lo * lo + hi * hi + lo * lo) as u128); + + let z = CrossPowerSums::default(); + assert_eq!(CrossPowerSums { sum_xy: i128::MAX, ..z }.checked_merge(CrossPowerSums { sum_xy: 1, ..z }), None); + assert_eq!(CrossPowerSums { sum_y: i64::MIN, ..z }.checked_merge(CrossPowerSums { sum_y: -1, ..z }), None); + assert_eq!( + CrossPowerSums { + sum_y_sq: u128::MAX, + ..z + } + .checked_merge(CrossPowerSums { sum_y_sq: 1, ..z }), + None + ); + } + + #[test] + #[should_panic(expected = "masked_group_cross_power_sums_i32: xs/ys length mismatch")] + fn cross_power_sums_refuse_mismatched_lanes() { + masked_group_cross_power_sums_i32(&[1], &[0], &[1], &[], &mut [CrossPowerSums::default()]); + } } From a62add431e30347e2f37425f10b176669b62eb51 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 29 Sep 2026 20:50:32 +0000 Subject: [PATCH 3/3] pillar/cov_high_d: from_symmetric_fn + pub get; test sandwich with non-identity M MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit CovHighD had no public way to be built from data (only identity() and the private zero()), so a consumer could not push a real covariance through sandwich without re-spelling the packed lower-triangle index. from_symmetric_fn reads the lower triangle once; get is made public. Every existing sandwich test used M = I, where a transposed or mis-indexed product is invisible. Adds a dense M·Σ·M triple-product check for a non-identity symmetric M. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01X1YcYMRSFvfczXoP748wtB --- src/hpc/pillar/cov_high_d.rs | 70 +++++++++++++++++++++++++++++++++++- 1 file changed, 69 insertions(+), 1 deletion(-) diff --git a/src/hpc/pillar/cov_high_d.rs b/src/hpc/pillar/cov_high_d.rs index 8f018c85..6d565ba3 100644 --- a/src/hpc/pillar/cov_high_d.rs +++ b/src/hpc/pillar/cov_high_d.rs @@ -69,8 +69,17 @@ impl CovHighD { } /// Read entry (i, j) — exploits symmetry so (i, j) and (j, i) both work. + /// + /// # Example + /// + /// ```rust + /// use ndarray::hpc::pillar::cov_high_d::CovHighD; + /// let s = CovHighD::<3>::from_symmetric_fn(|i, j| (10 * i + j) as f32); + /// assert_eq!(s.get(2, 1), 21.0); + /// assert_eq!(s.get(1, 2), 21.0); // mirrored from the lower triangle + /// ``` #[inline] - fn get(&self, i: usize, j: usize) -> f32 { + pub fn get(&self, i: usize, j: usize) -> f32 { if i >= j { self.lt[Self::idx(i, j)] } else { @@ -99,6 +108,30 @@ impl CovHighD { Self { lt } } + /// Construct a symmetric N×N matrix from an entry function. + /// + /// Only the lower triangle is read: `f(i, j)` is called once per `i ≥ j`, + /// and `(j, i)` is defined by symmetry. A caller holding a dense matrix + /// that is only approximately symmetric therefore gets its lower triangle, + /// not an average — symmetrise first if that matters. + /// + /// # Example + /// + /// ```rust + /// use ndarray::hpc::pillar::cov_high_d::CovHighD; + /// let d = CovHighD::<4>::from_symmetric_fn(|i, j| if i == j { 2.0 } else { 0.0 }); + /// assert!((d.frobenius_sq() - 16.0_f32).abs() < 1e-5); + /// ``` + pub fn from_symmetric_fn(mut f: impl FnMut(usize, usize) -> f32) -> Self { + let mut lt = Vec::with_capacity(N * (N + 1) / 2); + for i in 0..N { + for j in 0..=i { + lt.push(f(i, j)); + } + } + Self { lt } + } + /// Construct a zero N×N matrix. fn zero() -> Self { let size = N * (N + 1) / 2; @@ -423,6 +456,41 @@ mod tests { } } + #[test] + fn sandwich_matches_dense_m_sigma_m_for_non_identity_m() { + // Every other sandwich test uses M = I, where a transposed or + // mis-indexed product is invisible. Check a dense triple product. + const N: usize = 5; + let m = CovHighD::::from_symmetric_fn(|i, j| 0.3 * (i + 2 * j) as f32 - 1.0); + let s = CovHighD::::from_symmetric_fn(|i, j| if i == j { 1.0 + i as f32 } else { 0.1 * (i * j) as f32 }); + let r = s.sandwich(&m); + for i in 0..N { + for l in 0..N { + let mut want = 0.0_f64; + for j in 0..N { + for k in 0..N { + want += m.get(i, j) as f64 * s.get(j, k) as f64 * m.get(k, l) as f64; + } + } + let got = r.get(i, l) as f64; + assert!((got - want).abs() <= 1e-4 * want.abs().max(1.0), "[{i}][{l}] got {got} want {want}"); + } + } + } + + #[test] + fn from_symmetric_fn_reads_only_the_lower_triangle() { + let mut calls = Vec::new(); + let s = CovHighD::<3>::from_symmetric_fn(|i, j| { + calls.push((i, j)); + (i * 3 + j) as f32 + }); + assert!(calls.iter().all(|&(i, j)| i >= j)); + assert_eq!(calls.len(), 6); + assert_eq!(s.get(0, 2), s.get(2, 0)); + assert_eq!(s.get(2, 0), 6.0); + } + #[test] fn sandwich_scales_frobenius() { // If M = s·I then M Σ Mᵀ = s²·Σ, so ‖result‖_F = s² ‖Σ‖_F