From eeb911bde59ac7acf42920fc429e793d2e9cfd5e Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 4 Oct 2026 21:07:28 +0000 Subject: [PATCH 1/2] simd: bounded u8 power-sums folds into 128-bit registers, lossless widen, tiled driver (D-LXC-29) Register layout [n, sum, sum_sq, reserved] as four LE u32 words; the bivariate fold uses two rails, rail0 = the univariate register of x and rail1 = [sum_y, sum_y_sq, sum_xy, reserved], so n is stored once. Tile bound 2^16 rows, compile-checked to fit u32 at u8::MAX; a larger tile is refused before any write. widen_* gives the exact PowerSums / CrossPowerSums, and fold_*_tiles merges tiles with checked_merge, refusing a tile whole if any group's merge overflows. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01X1YcYMRSFvfczXoP748wtB --- src/simd.rs | 12 + src/simd_masking_ops.rs | 738 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 750 insertions(+) diff --git a/src/simd.rs b/src/simd.rs index c4ad7f7f..2268c4b8 100644 --- a/src/simd.rs +++ b/src/simd.rs @@ -793,6 +793,8 @@ pub use crate::simd_masking_ops::{ eq_u32_via_to_mask, eq_u64_to_mask, eq_u8_to_mask, + fold_bounded_cross_power_sums_tiles, + fold_bounded_power_sums_tiles, ge_i32_to_mask, ge_i32_to_mask_under, ge_u64_to_mask, @@ -829,6 +831,12 @@ pub use crate::simd_masking_ops::{ mask_ternlog_popcount, mask_xor, mask_xor_assign, + masked_group_bounded_cross_power_sums_u8, + masked_group_bounded_cross_power_sums_u8_pair, + masked_group_bounded_cross_power_sums_u8_via, + masked_group_bounded_power_sums_u8, + masked_group_bounded_power_sums_u8_pair, + masked_group_bounded_power_sums_u8_via, masked_group_count_u32, masked_group_count_u32_pair, masked_group_count_u32_via, @@ -867,10 +875,14 @@ pub use crate::simd_masking_ops::{ ternary_match_u32_to_mask_under, ternary_match_u64_to_mask, ternary_match_u64_to_mask_under, + widen_bounded_cross_power_sums, + widen_bounded_power_sums, + BoundedFoldError, CrossPowerSums, KeyRunCarry, MortonDir, PowerSums, + BOUNDED_TILE_ROWS, SYM_EMPTY_I64, }; // The popcount that closes the loop on the masks above: `mask_count` in ABI diff --git a/src/simd_masking_ops.rs b/src/simd_masking_ops.rs index b8737873..68aa37f3 100644 --- a/src/simd_masking_ops.rs +++ b/src/simd_masking_ops.rs @@ -2391,6 +2391,461 @@ fn cross_power_sums_fold<'a>( } } +/// Rows one bounded tile may hold: `2^16 = 65,536`. +/// +/// The bound at which a group's power sums of `u8` values fit `u32` words +/// exactly, even when every row of the tile lands in one group with value +/// `u8::MAX`: `n ≤ 2^16`, `Σx ≤ 255·2^16`, `Σx² ≤ 255²·2^16 < 2^32`, and the +/// cross term `Σxy` has the same bound as `Σx²`. Checked below at compile time. +pub const BOUNDED_TILE_ROWS: usize = 1 << 16; + +const _: () = { + let rows = BOUNDED_TILE_ROWS as u64; + let max = u8::MAX as u64; + assert!(rows <= u32::MAX as u64); + assert!(max * rows <= u32::MAX as u64); + assert!(max * max * rows <= u32::MAX as u64); +}; + +/// Why a bounded power-sums fold refused. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum BoundedFoldError { + /// The tile has more than [`BOUNDED_TILE_ROWS`] rows. Refused before any + /// register is written: past the bound a `u32` word could wrap. + TileTooLarge { + /// Rows in the refused tile. + rows: usize, + }, + /// Merging a widened tile into the wide sink overflowed for `group`. + /// Nothing of that tile was merged into any group. + MergeOverflow { + /// The first group whose merge overflowed. + group: usize, + }, +} + +/// One little-endian `u32` word of a 128-bit register. +#[inline(always)] +fn register_word(r: &[u8; 16], w: usize) -> u32 { + u32::from_le_bytes([r[4 * w], r[4 * w + 1], r[4 * w + 2], r[4 * w + 3]]) +} + +/// Set one little-endian `u32` word of a 128-bit register. +#[inline(always)] +fn set_register_word(r: &mut [u8; 16], w: usize, v: u32) { + r[4 * w..4 * w + 4].copy_from_slice(&v.to_le_bytes()); +} + +/// Refuse a tile past the bound, before anything is written. +#[inline(always)] +fn check_bounded_tile(rows: usize) -> Result<(), BoundedFoldError> { + if rows > BOUNDED_TILE_ROWS { + return Err(BoundedFoldError::TileTooLarge { rows }); + } + Ok(()) +} + +/// Zero the three accumulator words of every register. Word 3 is reserved +/// and left exactly as the caller holds it. +#[inline(always)] +fn seed_registers(regs: &mut [[u8; 16]]) { + for r in regs { + r[..12].fill(0); + } +} + +/// Keyed power sums of a `u8` lane over ONE bounded tile, into 128-bit +/// registers: for every row `i` selected by `mask_words`, folds `values[i]` +/// into `regs[keys[i]]`, provided `keys[i] < regs.len()`. +/// +/// # Register layout +/// +/// Four little-endian `u32` words: `[n, Σx, Σx², reserved]`. Word 3 is never +/// read or written. [`widen_bounded_power_sums`] turns a register into the +/// exact [`PowerSums`] that [`masked_group_power_sums_i32`] folds from the +/// same values. +/// +/// # Contract +/// +/// - The tile is at most [`BOUNDED_TILE_ROWS`] rows, else +/// [`BoundedFoldError::TileTooLarge`] and nothing is written. Within the +/// bound no word can overflow, whatever the keys and values. +/// - The registers are seeded by this call (words 0..3 zeroed): one call is +/// one tile. Larger populations go tile by tile through +/// [`fold_bounded_power_sums_tiles`]. +/// - Keys, drop rules and the final-word clamp are exactly those of +/// [`masked_group_power_sums_i32`]. +/// +/// # Panics +/// +/// Panics if `keys.len() != values.len()`, or if `mask_words.len() < +/// values.len().div_ceil(64)`. +/// +/// # Examples +/// +/// ``` +/// use ndarray::simd::{masked_group_bounded_power_sums_u8, widen_bounded_power_sums, PowerSums}; +/// +/// let mask = [0b1011u64]; // rows 0, 1, 3 +/// let keys = [0u32, 1, 0, 1]; +/// let values = [3u8, 2, 100, 255]; +/// let mut regs = [[0u8; 16]; 2]; +/// masked_group_bounded_power_sums_u8(&mask, &keys, &values, &mut regs).unwrap(); +/// assert_eq!(widen_bounded_power_sums(®s[0]), PowerSums { n: 1, sum: 3, sum_sq: 9 }); +/// assert_eq!(widen_bounded_power_sums(®s[1]), PowerSums { n: 2, sum: 257, sum_sq: 65029 }); +/// ``` +pub fn masked_group_bounded_power_sums_u8( + mask_words: &[u64], keys: &[u32], values: &[u8], regs: &mut [[u8; 16]], +) -> Result<(), BoundedFoldError> { + assert_eq!(keys.len(), values.len(), "masked_group_bounded_power_sums_u8: keys/values length mismatch"); + check_bounded_tile(values.len())?; + seed_registers(regs); + let groups = regs.len(); + group_walk( + "masked_group_bounded_power_sums_u8", + mask_words, + values.len(), + GroupKeyAddr::Resident(keys), + groups, + bounded_power_sums_fold(values, regs), + ); + Ok(()) +} + +/// [`masked_group_bounded_power_sums_u8`] with the group read through an +/// index lane, `regs[table[index[i]]]`, zero-fallback exactly as +/// [`masked_group_power_sums_i32_via`]. +/// +/// # Panics +/// +/// Panics if `index.len() != values.len()`, or if `mask_words.len() < +/// values.len().div_ceil(64)`. +pub fn masked_group_bounded_power_sums_u8_via( + mask_words: &[u64], index: &[u32], table: &[u32], values: &[u8], regs: &mut [[u8; 16]], +) -> Result<(), BoundedFoldError> { + assert_eq!(index.len(), values.len(), "masked_group_bounded_power_sums_u8_via: index/values length mismatch"); + check_bounded_tile(values.len())?; + seed_registers(regs); + let groups = regs.len(); + group_walk( + "masked_group_bounded_power_sums_u8_via", + mask_words, + values.len(), + GroupKeyAddr::Via { index, table }, + groups, + bounded_power_sums_fold(values, regs), + ); + Ok(()) +} + +/// [`masked_group_bounded_power_sums_u8`] with the composite group +/// `hi[i] * stride + lo[i]`, zero-fallback exactly as +/// [`masked_group_power_sums_i32_pair`]. +/// +/// # Panics +/// +/// Panics if `hi`, `lo` and `values` differ in length, or if +/// `mask_words.len() < values.len().div_ceil(64)`. +pub fn masked_group_bounded_power_sums_u8_pair( + mask_words: &[u64], hi: &[u32], lo: &[u32], stride: u32, values: &[u8], regs: &mut [[u8; 16]], +) -> Result<(), BoundedFoldError> { + assert_eq!(hi.len(), lo.len(), "masked_group_bounded_power_sums_u8_pair: hi.len() != lo.len()"); + assert_eq!(hi.len(), values.len(), "masked_group_bounded_power_sums_u8_pair: hi/values length mismatch"); + check_bounded_tile(values.len())?; + seed_registers(regs); + let groups = regs.len(); + group_walk( + "masked_group_bounded_power_sums_u8_pair", + mask_words, + values.len(), + GroupKeyAddr::Pair { hi, lo, stride }, + groups, + bounded_power_sums_fold(values, regs), + ); + Ok(()) +} + +/// The one bounded univariate fold, shared by all three key addresses. +#[inline(always)] +fn bounded_power_sums_fold<'a>(values: &'a [u8], regs: &'a mut [[u8; 16]]) -> impl FnMut(usize, usize) + 'a { + move |k, i| { + let x = u32::from(values[i]); + let r = &mut regs[k]; + set_register_word(r, 0, register_word(r, 0).wrapping_add(1)); + set_register_word(r, 1, register_word(r, 1).wrapping_add(x)); + set_register_word(r, 2, register_word(r, 2).wrapping_add(x * x)); + } +} + +/// Keyed joint power sums of two `u8` lanes over ONE bounded tile, into two +/// 128-bit registers per group. +/// +/// # Register layout +/// +/// - `rail0[k]` = `[n, Σx, Σx², reserved]` — exactly the register +/// [`masked_group_bounded_power_sums_u8`] folds from `xs` alone. +/// - `rail1[k]` = `[Σy, Σy², Σxy, reserved]`. +/// +/// `n` is stored once. Word 3 of each rail is never read or written. +/// [`widen_bounded_cross_power_sums`] turns the pair into the exact +/// [`CrossPowerSums`] that [`masked_group_cross_power_sums_i32`] folds. +/// +/// Same contract as [`masked_group_bounded_power_sums_u8`] otherwise. +/// +/// # Panics +/// +/// Panics if `keys`, `xs` and `ys` differ in length, if `rail0.len() != +/// rail1.len()`, or if `mask_words.len() < xs.len().div_ceil(64)`. +/// +/// # Examples +/// +/// ``` +/// use ndarray::simd::{masked_group_bounded_cross_power_sums_u8, widen_bounded_cross_power_sums}; +/// +/// let mask = [0b111u64]; +/// let keys = [0u32, 0, 1]; +/// let (xs, ys) = ([2u8, 255, 1], [3u8, 255, 0]); +/// let (mut r0, mut r1) = ([[0u8; 16]; 2], [[0u8; 16]; 2]); +/// masked_group_bounded_cross_power_sums_u8(&mask, &keys, &xs, &ys, &mut r0, &mut r1).unwrap(); +/// let c = widen_bounded_cross_power_sums(&r0[0], &r1[0]); +/// assert_eq!((c.n, c.sum_x, c.sum_y, c.sum_xy), (2, 257, 258, 6 + 255 * 255)); +/// ``` +pub fn masked_group_bounded_cross_power_sums_u8( + mask_words: &[u64], keys: &[u32], xs: &[u8], ys: &[u8], rail0: &mut [[u8; 16]], rail1: &mut [[u8; 16]], +) -> Result<(), BoundedFoldError> { + assert_eq!(keys.len(), xs.len(), "masked_group_bounded_cross_power_sums_u8: keys/xs length mismatch"); + assert_eq!(xs.len(), ys.len(), "masked_group_bounded_cross_power_sums_u8: xs/ys length mismatch"); + assert_eq!(rail0.len(), rail1.len(), "masked_group_bounded_cross_power_sums_u8: rail0/rail1 length mismatch"); + check_bounded_tile(xs.len())?; + seed_registers(rail0); + seed_registers(rail1); + let groups = rail0.len(); + group_walk( + "masked_group_bounded_cross_power_sums_u8", + mask_words, + xs.len(), + GroupKeyAddr::Resident(keys), + groups, + bounded_cross_power_sums_fold(xs, ys, rail0, rail1), + ); + Ok(()) +} + +/// [`masked_group_bounded_cross_power_sums_u8`] with the group read through +/// an index lane, zero-fallback exactly as +/// [`masked_group_cross_power_sums_i32_via`]. +/// +/// # Panics +/// +/// Panics if `index`, `xs` and `ys` differ in length, if `rail0.len() != +/// rail1.len()`, or if `mask_words.len() < xs.len().div_ceil(64)`. +pub fn masked_group_bounded_cross_power_sums_u8_via( + mask_words: &[u64], index: &[u32], table: &[u32], xs: &[u8], ys: &[u8], rail0: &mut [[u8; 16]], + rail1: &mut [[u8; 16]], +) -> Result<(), BoundedFoldError> { + assert_eq!(index.len(), xs.len(), "masked_group_bounded_cross_power_sums_u8_via: index/xs length mismatch"); + assert_eq!(xs.len(), ys.len(), "masked_group_bounded_cross_power_sums_u8_via: xs/ys length mismatch"); + assert_eq!( + rail0.len(), + rail1.len(), + "masked_group_bounded_cross_power_sums_u8_via: rail0/rail1 length mismatch" + ); + check_bounded_tile(xs.len())?; + seed_registers(rail0); + seed_registers(rail1); + let groups = rail0.len(); + group_walk( + "masked_group_bounded_cross_power_sums_u8_via", + mask_words, + xs.len(), + GroupKeyAddr::Via { index, table }, + groups, + bounded_cross_power_sums_fold(xs, ys, rail0, rail1), + ); + Ok(()) +} + +/// [`masked_group_bounded_cross_power_sums_u8`] with the composite group +/// `hi[i] * stride + lo[i]`, zero-fallback exactly as +/// [`masked_group_cross_power_sums_i32_pair`]. +/// +/// # Panics +/// +/// Panics if `hi`, `lo`, `xs` and `ys` differ in length, if `rail0.len() != +/// rail1.len()`, or if `mask_words.len() < xs.len().div_ceil(64)`. +#[allow(clippy::too_many_arguments)] +pub fn masked_group_bounded_cross_power_sums_u8_pair( + mask_words: &[u64], hi: &[u32], lo: &[u32], stride: u32, xs: &[u8], ys: &[u8], rail0: &mut [[u8; 16]], + rail1: &mut [[u8; 16]], +) -> Result<(), BoundedFoldError> { + assert_eq!(hi.len(), lo.len(), "masked_group_bounded_cross_power_sums_u8_pair: hi.len() != lo.len()"); + assert_eq!(hi.len(), xs.len(), "masked_group_bounded_cross_power_sums_u8_pair: hi/xs length mismatch"); + assert_eq!(xs.len(), ys.len(), "masked_group_bounded_cross_power_sums_u8_pair: xs/ys length mismatch"); + assert_eq!( + rail0.len(), + rail1.len(), + "masked_group_bounded_cross_power_sums_u8_pair: rail0/rail1 length mismatch" + ); + check_bounded_tile(xs.len())?; + seed_registers(rail0); + seed_registers(rail1); + let groups = rail0.len(); + group_walk( + "masked_group_bounded_cross_power_sums_u8_pair", + mask_words, + xs.len(), + GroupKeyAddr::Pair { hi, lo, stride }, + groups, + bounded_cross_power_sums_fold(xs, ys, rail0, rail1), + ); + Ok(()) +} + +/// The one bounded bivariate fold, shared by all three key addresses. +#[inline(always)] +fn bounded_cross_power_sums_fold<'a>( + xs: &'a [u8], ys: &'a [u8], rail0: &'a mut [[u8; 16]], rail1: &'a mut [[u8; 16]], +) -> impl FnMut(usize, usize) + 'a { + move |k, i| { + let (x, y) = (u32::from(xs[i]), u32::from(ys[i])); + let r0 = &mut rail0[k]; + set_register_word(r0, 0, register_word(r0, 0).wrapping_add(1)); + set_register_word(r0, 1, register_word(r0, 1).wrapping_add(x)); + set_register_word(r0, 2, register_word(r0, 2).wrapping_add(x * x)); + let r1 = &mut rail1[k]; + set_register_word(r1, 0, register_word(r1, 0).wrapping_add(y)); + set_register_word(r1, 1, register_word(r1, 1).wrapping_add(y * y)); + set_register_word(r1, 2, register_word(r1, 2).wrapping_add(x * y)); + } +} + +/// Widen a bounded register `[n, Σx, Σx², reserved]` to the exact +/// [`PowerSums`] record. Lossless: every `u32` word fits the wider field. +#[inline] +#[must_use] +pub fn widen_bounded_power_sums(reg: &[u8; 16]) -> PowerSums { + PowerSums { + n: u64::from(register_word(reg, 0)), + sum: i64::from(register_word(reg, 1)), + sum_sq: u128::from(register_word(reg, 2)), + } +} + +/// Widen a bounded register pair (`rail0` = `[n, Σx, Σx², reserved]`, +/// `rail1` = `[Σy, Σy², Σxy, reserved]`) to the exact [`CrossPowerSums`]. +#[inline] +#[must_use] +pub fn widen_bounded_cross_power_sums(rail0: &[u8; 16], rail1: &[u8; 16]) -> CrossPowerSums { + CrossPowerSums { + n: u64::from(register_word(rail0, 0)), + sum_x: i64::from(register_word(rail0, 1)), + sum_y: i64::from(register_word(rail1, 0)), + sum_x_sq: u128::from(register_word(rail0, 2)), + sum_y_sq: u128::from(register_word(rail1, 1)), + sum_xy: i128::from(register_word(rail1, 2)), + } +} + +/// The tiles a population of `n_rows` is cut into. Each tile starts at a +/// multiple of [`BOUNDED_TILE_ROWS`], so its mask window starts at word +/// `start / 64`. +#[inline] +fn bounded_tiles(n_rows: usize) -> impl Iterator> { + (0..n_rows) + .step_by(BOUNDED_TILE_ROWS) + .map(move |s| s..(s + BOUNDED_TILE_ROWS).min(n_rows)) +} + +/// Fold a population of any size through bounded tiles: for each tile of at +/// most [`BOUNDED_TILE_ROWS`] rows, `tile(range, regs)` folds the tile into +/// the registers (one of the `masked_group_bounded_power_sums_u8*` kernels +/// over the tile's slices), then every register is widened and merged into +/// `out` with [`PowerSums::checked_merge`]. +/// +/// A population larger than the tile bound is therefore not an error: only +/// the per-tile accumulator is narrow. `out` is accumulated into (start it at +/// [`PowerSums::default`]). A tile whose merge would overflow any group is +/// refused whole, so `out` only ever holds fully merged tiles. +/// +/// # Errors +/// +/// Whatever `tile` returns, or [`BoundedFoldError::MergeOverflow`]. +/// +/// # Panics +/// +/// Panics if `regs.len() != out.len()`. +/// +/// # Examples +/// +/// ``` +/// use ndarray::simd::{ +/// fold_bounded_power_sums_tiles, masked_group_bounded_power_sums_u8, PowerSums, BOUNDED_TILE_ROWS, +/// }; +/// +/// let n = BOUNDED_TILE_ROWS + 3; +/// let mask = vec![u64::MAX; n.div_ceil(64)]; +/// let keys = vec![0u32; n]; +/// let values = vec![255u8; n]; +/// let (mut regs, mut out) = ([[0u8; 16]; 1], [PowerSums::default(); 1]); +/// fold_bounded_power_sums_tiles(n, &mut regs, &mut out, |t, regs| { +/// masked_group_bounded_power_sums_u8(&mask[t.start / 64..], &keys[t.clone()], &values[t], regs) +/// }) +/// .unwrap(); +/// assert_eq!(out[0].n, n as u64); +/// assert_eq!(out[0].sum_sq, 255 * 255 * n as u128); +/// ``` +pub fn fold_bounded_power_sums_tiles( + n_rows: usize, regs: &mut [[u8; 16]], out: &mut [PowerSums], + mut tile: impl FnMut(core::ops::Range, &mut [[u8; 16]]) -> Result<(), BoundedFoldError>, +) -> Result<(), BoundedFoldError> { + assert_eq!(regs.len(), out.len(), "fold_bounded_power_sums_tiles: regs/out length mismatch"); + for range in bounded_tiles(n_rows) { + tile(range, regs)?; + for (group, (r, o)) in regs.iter().zip(out.iter()).enumerate() { + o.checked_merge(widen_bounded_power_sums(r)) + .ok_or(BoundedFoldError::MergeOverflow { group })?; + } + for (r, o) in regs.iter().zip(out.iter_mut()) { + *o = o + .checked_merge(widen_bounded_power_sums(r)) + .expect("checked above"); + } + } + Ok(()) +} + +/// [`fold_bounded_power_sums_tiles`] for the bivariate fold: each tile folds +/// into the register pair, which is widened to [`CrossPowerSums`] and merged +/// into `out` with [`CrossPowerSums::checked_merge`]. +/// +/// # Errors +/// +/// Whatever `tile` returns, or [`BoundedFoldError::MergeOverflow`]. +/// +/// # Panics +/// +/// Panics if `rail0`, `rail1` and `out` differ in length. +pub fn fold_bounded_cross_power_sums_tiles( + n_rows: usize, rail0: &mut [[u8; 16]], rail1: &mut [[u8; 16]], out: &mut [CrossPowerSums], + mut tile: impl FnMut(core::ops::Range, &mut [[u8; 16]], &mut [[u8; 16]]) -> Result<(), BoundedFoldError>, +) -> Result<(), BoundedFoldError> { + assert_eq!(rail0.len(), out.len(), "fold_bounded_cross_power_sums_tiles: rail0/out length mismatch"); + assert_eq!(rail1.len(), out.len(), "fold_bounded_cross_power_sums_tiles: rail1/out length mismatch"); + for range in bounded_tiles(n_rows) { + tile(range, rail0, rail1)?; + for (group, ((a, b), o)) in rail0.iter().zip(rail1.iter()).zip(out.iter()).enumerate() { + o.checked_merge(widen_bounded_cross_power_sums(a, b)) + .ok_or(BoundedFoldError::MergeOverflow { group })?; + } + for ((a, b), o) in rail0.iter().zip(rail1.iter()).zip(out.iter_mut()) { + *o = o + .checked_merge(widen_bounded_cross_power_sums(a, b)) + .expect("checked above"); + } + } + Ok(()) +} + /// 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 @@ -8464,3 +8919,286 @@ mod group_family_tests { masked_group_cross_power_sums_i32(&[1], &[0], &[1], &[], &mut [CrossPowerSums::default()]); } } + +#[cfg(test)] +mod bounded_register_tests { + use super::*; + + /// Deterministic `u8` lane with every value 0..=255 hit, including both ends. + fn lane(n: usize, seed: u64) -> Vec { + let mut s = seed; + (0..n) + .map(|_| { + s = s + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + (s >> 56) as u8 + }) + .collect() + } + + /// A mask with a hole every 7th row, so the selection is not "all". + fn holey_mask(n: usize) -> Vec { + let mut m = vec![0u64; n.div_ceil(64)]; + for i in (0..n).filter(|i| i % 7 != 3) { + m[i / 64] |= 1 << (i % 64); + } + m + } + + fn widen_i32(v: &[u8]) -> Vec { + v.iter().map(|&x| i32::from(x)).collect() + } + + /// FAILS IF: a tile of exactly the bound, every row in one group at + /// `u8::MAX`, is refused or loses a bit — the worst case the bound covers. + #[test] + fn the_full_bound_at_u8_max_fits_exactly() { + let n = BOUNDED_TILE_ROWS; + let mask = vec![u64::MAX; n / 64]; + let keys = vec![0u32; n]; + let (xs, ys) = (vec![u8::MAX; n], vec![u8::MAX; n]); + let mut regs = [[0u8; 16]; 1]; + masked_group_bounded_power_sums_u8(&mask, &keys, &xs, &mut regs).unwrap(); + let p = widen_bounded_power_sums(®s[0]); + assert_eq!(p.n, n as u64); + assert_eq!(p.sum, 255 * n as i64); + assert_eq!(p.sum_sq, 255 * 255 * n as u128); + assert!(p.sum_sq < 1u128 << 32, "worst case stays inside u32"); + + let (mut r0, mut r1) = ([[0u8; 16]; 1], [[0u8; 16]; 1]); + masked_group_bounded_cross_power_sums_u8(&mask, &keys, &xs, &ys, &mut r0, &mut r1).unwrap(); + let c = widen_bounded_cross_power_sums(&r0[0], &r1[0]); + assert_eq!(c.sum_xy, 255 * 255 * n as i128); + assert_eq!(c.sum_y_sq, 255 * 255 * n as u128); + } + + /// FAILS IF: one row past the bound is folded instead of refused, or the + /// refusal writes anything. The registers are pre-filled so a partial + /// write is visible. + #[test] + fn one_row_past_the_bound_is_refused_before_any_write() { + let n = BOUNDED_TILE_ROWS + 1; + let mask = vec![u64::MAX; n.div_ceil(64)]; + let keys = vec![0u32; n]; + let xs = vec![u8::MAX; n]; + let sentinel = [0xAB; 16]; + let mut regs = [sentinel; 1]; + assert_eq!( + masked_group_bounded_power_sums_u8(&mask, &keys, &xs, &mut regs), + Err(BoundedFoldError::TileTooLarge { rows: n }) + ); + assert_eq!(regs[0], sentinel, "a refused tile wrote nothing"); + let (mut r0, mut r1) = ([sentinel; 1], [sentinel; 1]); + assert_eq!( + masked_group_bounded_cross_power_sums_u8(&mask, &keys, &xs, &xs, &mut r0, &mut r1), + Err(BoundedFoldError::TileTooLarge { rows: n }) + ); + assert_eq!((r0[0], r1[0]), (sentinel, sentinel)); + } + + /// FAILS IF: word 3 is touched (it is reserved, carries no semantics). + #[test] + fn the_reserved_word_is_never_written() { + let mut regs = [[0xCD; 16]; 2]; + masked_group_bounded_power_sums_u8(&[0b11], &[0, 1], &[9, 200], &mut regs).unwrap(); + for r in ®s { + assert_eq!(r[12..], [0xCD; 4]); + } + let (mut r0, mut r1) = ([[0xCD; 16]; 1], [[0xEF; 16]; 1]); + masked_group_bounded_cross_power_sums_u8(&[1], &[0], &[3], &[4], &mut r0, &mut r1).unwrap(); + assert_eq!(r0[0][12..], [0xCD; 4]); + assert_eq!(r1[0][12..], [0xEF; 4]); + } + + /// FAILS IF: narrow → widen differs from the wide `i32` kernel for any of + /// the three key addresses (lane / via / pair), drop rules included. + #[test] + fn narrow_then_widen_equals_the_wide_kernel_for_every_address() { + let n = 5_000; + let groups = 9; + let mask = holey_mask(n); + let xs = lane(n, 1); + let wide = widen_i32(&xs); + // keys deliberately include out-of-range values (dropped by both paths) + let keys: Vec = (0..n as u32) + .map(|i| (i.wrapping_mul(2_654_435_761) >> 13) % (groups as u32 + 2)) + .collect(); + + let mut regs = vec![[0u8; 16]; groups]; + let mut want = vec![PowerSums::default(); groups]; + masked_group_bounded_power_sums_u8(&mask, &keys, &xs, &mut regs).unwrap(); + masked_group_power_sums_i32(&mask, &keys, &wide, &mut want); + let got: Vec<_> = regs.iter().map(widen_bounded_power_sums).collect(); + assert_eq!(got, want, "lane"); + assert!(want.iter().all(|p| p.n > 0), "every group populated: {want:?}"); + + let table: Vec = (0..13u32).map(|t| (t * 5) % (groups as u32 + 1)).collect(); + let index: Vec = (0..n as u32).map(|i| (i * 11) % 15).collect(); + let mut regs = vec![[0u8; 16]; groups]; + let mut want = vec![PowerSums::default(); groups]; + masked_group_bounded_power_sums_u8_via(&mask, &index, &table, &xs, &mut regs).unwrap(); + masked_group_power_sums_i32_via(&mask, &index, &table, &wide, &mut want); + assert_eq!( + regs.iter() + .map(widen_bounded_power_sums) + .collect::>(), + want, + "via" + ); + + let hi: Vec = (0..n as u32).map(|i| i % 4).collect(); + let lo: Vec = (0..n as u32).map(|i| (i / 4) % 3).collect(); + let mut regs = vec![[0u8; 16]; groups]; + let mut want = vec![PowerSums::default(); groups]; + masked_group_bounded_power_sums_u8_pair(&mask, &hi, &lo, 3, &xs, &mut regs).unwrap(); + masked_group_power_sums_i32_pair(&mask, &hi, &lo, 3, &wide, &mut want); + assert_eq!( + regs.iter() + .map(widen_bounded_power_sums) + .collect::>(), + want, + "pair" + ); + } + + /// FAILS IF: the bivariate narrow → widen differs from the wide cross + /// kernel for any address, or rail 0 differs from the univariate register. + #[test] + fn bivariate_narrow_then_widen_equals_the_wide_cross_kernel() { + let n = 4_321; + let groups = 6; + let mask = holey_mask(n); + let (xs, ys) = (lane(n, 2), lane(n, 3)); + let (wx, wy) = (widen_i32(&xs), widen_i32(&ys)); + let keys: Vec = (0..n as u32) + .map(|i| (i * 13) % (groups as u32 + 1)) + .collect(); + + let (mut r0, mut r1) = (vec![[0u8; 16]; groups], vec![[0u8; 16]; groups]); + let mut want = vec![CrossPowerSums::default(); groups]; + masked_group_bounded_cross_power_sums_u8(&mask, &keys, &xs, &ys, &mut r0, &mut r1).unwrap(); + masked_group_cross_power_sums_i32(&mask, &keys, &wx, &wy, &mut want); + let got: Vec<_> = r0 + .iter() + .zip(&r1) + .map(|(a, b)| widen_bounded_cross_power_sums(a, b)) + .collect(); + assert_eq!(got, want, "lane"); + let mut uni = vec![[0u8; 16]; groups]; + masked_group_bounded_power_sums_u8(&mask, &keys, &xs, &mut uni).unwrap(); + for (a, b) in r0.iter().zip(&uni) { + assert_eq!(a[..12], b[..12], "rail 0 is the univariate register of x"); + } + + let table: Vec = vec![0, 1, 2, 3, 4, 5, 9]; + let index: Vec = (0..n as u32).map(|i| (i * 3) % 8).collect(); + let (mut r0, mut r1) = (vec![[0u8; 16]; groups], vec![[0u8; 16]; groups]); + let mut want = vec![CrossPowerSums::default(); groups]; + masked_group_bounded_cross_power_sums_u8_via(&mask, &index, &table, &xs, &ys, &mut r0, &mut r1).unwrap(); + masked_group_cross_power_sums_i32_via(&mask, &index, &table, &wx, &wy, &mut want); + let got: Vec<_> = r0 + .iter() + .zip(&r1) + .map(|(a, b)| widen_bounded_cross_power_sums(a, b)) + .collect(); + assert_eq!(got, want, "via"); + + let hi: Vec = (0..n as u32).map(|i| i % 3).collect(); + let lo: Vec = (0..n as u32).map(|i| (i / 3) % 2).collect(); + let (mut r0, mut r1) = (vec![[0u8; 16]; groups], vec![[0u8; 16]; groups]); + let mut want = vec![CrossPowerSums::default(); groups]; + masked_group_bounded_cross_power_sums_u8_pair(&mask, &hi, &lo, 2, &xs, &ys, &mut r0, &mut r1).unwrap(); + masked_group_cross_power_sums_i32_pair(&mask, &hi, &lo, 2, &wx, &wy, &mut want); + let got: Vec<_> = r0 + .iter() + .zip(&r1) + .map(|(a, b)| widen_bounded_cross_power_sums(a, b)) + .collect(); + assert_eq!(got, want, "pair"); + } + + /// FAILS IF: tiling + `checked_merge` across disjoint tiles differs from + /// one wide fold over the whole population. Population spans three tiles, + /// the last partial and not 64-aligned at its end. + #[test] + fn partitioned_tiles_merge_to_the_whole() { + let n = 2 * BOUNDED_TILE_ROWS + 1_234; + let groups = 5; + let mask = holey_mask(n); + let (xs, ys) = (lane(n, 4), lane(n, 5)); + let (wx, wy) = (widen_i32(&xs), widen_i32(&ys)); + let keys: Vec = (0..n as u32) + .map(|i| (i ^ (i >> 5)) % groups as u32) + .collect(); + + let mut regs = vec![[0u8; 16]; groups]; + let mut out = vec![PowerSums::default(); groups]; + let mut tiles = 0; + fold_bounded_power_sums_tiles(n, &mut regs, &mut out, |t, regs| { + tiles += 1; + masked_group_bounded_power_sums_u8(&mask[t.start / 64..], &keys[t.clone()], &xs[t], regs) + }) + .unwrap(); + assert_eq!(tiles, 3); + let mut want = vec![PowerSums::default(); groups]; + masked_group_power_sums_i32(&mask, &keys, &wx, &mut want); + assert_eq!(out, want); + assert!( + out.iter().map(|p| p.n).sum::() > BOUNDED_TILE_ROWS as u64, + "the population really spans more than one tile" + ); + + let (mut r0, mut r1) = (vec![[0u8; 16]; groups], vec![[0u8; 16]; groups]); + let mut out = vec![CrossPowerSums::default(); groups]; + fold_bounded_cross_power_sums_tiles(n, &mut r0, &mut r1, &mut out, |t, a, b| { + masked_group_bounded_cross_power_sums_u8( + &mask[t.start / 64..], + &keys[t.clone()], + &xs[t.clone()], + &ys[t], + a, + b, + ) + }) + .unwrap(); + let mut want = vec![CrossPowerSums::default(); groups]; + masked_group_cross_power_sums_i32(&mask, &keys, &wx, &wy, &mut want); + assert_eq!(out, want); + } + + /// FAILS IF: a tile whose merge overflows is partially committed. Group 1 + /// is seeded at the `n` ceiling so its merge overflows while group 0's + /// would succeed; group 0 must be left untouched. + #[test] + fn an_overflowing_merge_commits_nothing_of_the_tile() { + let mut regs = [[0u8; 16]; 2]; + let mut out = [ + PowerSums::default(), + PowerSums { + n: u64::MAX, + sum: 0, + sum_sq: 0, + }, + ]; + let before = out; + let err = fold_bounded_power_sums_tiles(2, &mut regs, &mut out, |_, regs| { + masked_group_bounded_power_sums_u8(&[0b11], &[0, 1], &[5, 6], regs) + }); + assert_eq!(err, Err(BoundedFoldError::MergeOverflow { group: 1 })); + assert_eq!(out, before); + } + + /// FAILS IF: a tile error from the caller's closure is swallowed. + #[test] + fn a_refused_tile_aborts_the_tiled_fold() { + let n = BOUNDED_TILE_ROWS + 5; + let mut regs = [[0u8; 16]; 1]; + let mut out = [PowerSums::default()]; + let err = fold_bounded_power_sums_tiles(n, &mut regs, &mut out, |_, _| { + Err(BoundedFoldError::TileTooLarge { rows: n }) + }); + assert_eq!(err, Err(BoundedFoldError::TileTooLarge { rows: n })); + assert_eq!(out[0], PowerSums::default()); + } +} From 4df12bc5ec96e2e2dfdfd2533c3f7f7648cdf3c1 Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 4 Oct 2026 21:12:16 +0000 Subject: [PATCH 2/2] examples: bounded vs wide power-sums benchmark (D-LXC-29) Asserts exactness, then times the u8 register folds against the wide i32 kernels per tile size and over a 16-tile population, and prints the realized SIMD tier. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01X1YcYMRSFvfczXoP748wtB --- Cargo.toml | 4 + examples/bounded_power_sums_bench.rs | 179 +++++++++++++++++++++++++++ 2 files changed, 183 insertions(+) create mode 100644 examples/bounded_power_sums_bench.rs diff --git a/Cargo.toml b/Cargo.toml index b5e1097a..50a0c9a9 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -111,6 +111,10 @@ required-features = ["std"] name = "amx_probe" required-features = ["std"] +[[example]] +name = "bounded_power_sums_bench" +required-features = ["std"] + [[example]] name = "amx_rb_probe" required-features = ["std"] diff --git a/examples/bounded_power_sums_bench.rs b/examples/bounded_power_sums_bench.rs new file mode 100644 index 00000000..bd7824b0 --- /dev/null +++ b/examples/bounded_power_sums_bench.rs @@ -0,0 +1,179 @@ +//! D-LXC-29 measurement: bounded `u8` power sums into 128-bit registers vs the +//! wide `i32` kernels, univariate and bivariate, per tile size and over a +//! multi-tile population. Exactness is asserted on every run before timing is +//! reported. Prints the realized SIMD tier; a timing without it is an anecdote. +//! +//! ```sh +//! env -u RUSTFLAGS cargo run --release --example bounded_power_sums_bench +//! ``` + +use ndarray::simd::{ + fold_bounded_cross_power_sums_tiles, fold_bounded_power_sums_tiles, masked_group_bounded_cross_power_sums_u8, + masked_group_bounded_power_sums_u8, masked_group_cross_power_sums_i32, masked_group_power_sums_i32, + widen_bounded_cross_power_sums, widen_bounded_power_sums, CrossPowerSums, PowerSums, BOUNDED_TILE_ROWS, +}; +use std::hint::black_box; +use std::time::Instant; + +const GROUPS: usize = 16; + +fn data(n: usize) -> (Vec, Vec, Vec, Vec) { + let mut s = 0x9E37_79B9_7F4A_7C15u64; + let mut next = || { + s = s + .wrapping_mul(6364136223846793005) + .wrapping_add(1442695040888963407); + s + }; + let mut mask = vec![0u64; n.div_ceil(64)]; + for (i, _) in (0..n).enumerate().filter(|(i, _)| i % 9 != 2) { + mask[i / 64] |= 1 << (i % 64); + } + let keys = (0..n) + .map(|_| (next() >> 40) as u32 % GROUPS as u32) + .collect(); + let xs = (0..n).map(|_| (next() >> 56) as u8).collect(); + let ys = (0..n).map(|_| (next() >> 56) as u8).collect(); + (mask, keys, xs, ys) +} + +/// Median ns/row over `reps` runs. +fn time(rows: usize, reps: usize, mut f: impl FnMut()) -> f64 { + let mut v: Vec = (0..reps) + .map(|_| { + let t = Instant::now(); + f(); + t.elapsed().as_nanos() as f64 / rows as f64 + }) + .collect(); + v.sort_by(f64::total_cmp); + v[reps / 2] +} + +fn main() { + #[cfg(target_arch = "x86_64")] + println!("tier: avx512f={} avx2={}", cfg!(target_feature = "avx512f"), cfg!(target_feature = "avx2")); + #[cfg(not(target_arch = "x86_64"))] + println!("tier: arch={}", std::env::consts::ARCH); + println!("groups={GROUPS} median of 31 runs, ns/row (lower is better)\n"); + println!("{:<34}{:>10}{:>10}{:>9}", "case", "wide i32", "bounded", "ratio"); + + for &tile in &[4_096usize, 16_384, BOUNDED_TILE_ROWS] { + let (mask, keys, xs, ys) = data(tile); + let wx: Vec = xs.iter().map(|&x| i32::from(x)).collect(); + let wy: Vec = ys.iter().map(|&y| i32::from(y)).collect(); + + // exactness first + let mut regs = [[0u8; 16]; GROUPS]; + let mut want = [PowerSums::default(); GROUPS]; + masked_group_bounded_power_sums_u8(&mask, &keys, &xs, &mut regs).unwrap(); + masked_group_power_sums_i32(&mask, &keys, &wx, &mut want); + assert!(regs + .iter() + .map(widen_bounded_power_sums) + .eq(want.iter().copied())); + + let w = time(tile, 31, || { + let mut o = [PowerSums::default(); GROUPS]; + masked_group_power_sums_i32(black_box(&mask), black_box(&keys), black_box(&wx), &mut o); + black_box(o); + }); + let b = time(tile, 31, || { + let mut r = [[0u8; 16]; GROUPS]; + masked_group_bounded_power_sums_u8(black_box(&mask), black_box(&keys), black_box(&xs), &mut r).unwrap(); + black_box(r); + }); + println!("{:<34}{w:>10.3}{b:>10.3}{:>9.2}", format!("univariate tile={tile}"), w / b); + + let (mut r0, mut r1) = ([[0u8; 16]; GROUPS], [[0u8; 16]; GROUPS]); + let mut want = [CrossPowerSums::default(); GROUPS]; + masked_group_bounded_cross_power_sums_u8(&mask, &keys, &xs, &ys, &mut r0, &mut r1).unwrap(); + masked_group_cross_power_sums_i32(&mask, &keys, &wx, &wy, &mut want); + assert!(r0 + .iter() + .zip(&r1) + .map(|(a, b)| widen_bounded_cross_power_sums(a, b)) + .eq(want.iter().copied())); + + let w = time(tile, 31, || { + let mut o = [CrossPowerSums::default(); GROUPS]; + masked_group_cross_power_sums_i32( + black_box(&mask), + black_box(&keys), + black_box(&wx), + black_box(&wy), + &mut o, + ); + black_box(o); + }); + let b = time(tile, 31, || { + let (mut a, mut c) = ([[0u8; 16]; GROUPS], [[0u8; 16]; GROUPS]); + masked_group_bounded_cross_power_sums_u8( + black_box(&mask), + black_box(&keys), + black_box(&xs), + black_box(&ys), + &mut a, + &mut c, + ) + .unwrap(); + black_box((a, c)); + }); + println!("{:<34}{w:>10.3}{b:>10.3}{:>9.2}", format!("bivariate tile={tile}"), w / b); + } + + // Multi-tile population through the tiled drivers (widen + merge included). + let n = 16 * BOUNDED_TILE_ROWS + 123; + let (mask, keys, xs, ys) = data(n); + let wx: Vec = xs.iter().map(|&x| i32::from(x)).collect(); + let wy: Vec = ys.iter().map(|&y| i32::from(y)).collect(); + let tiled_uni = || { + let (mut regs, mut out) = ([[0u8; 16]; GROUPS], [PowerSums::default(); GROUPS]); + fold_bounded_power_sums_tiles(n, &mut regs, &mut out, |t, regs| { + masked_group_bounded_power_sums_u8(&mask[t.start / 64..], &keys[t.clone()], &xs[t], regs) + }) + .unwrap(); + out + }; + let mut want = [PowerSums::default(); GROUPS]; + masked_group_power_sums_i32(&mask, &keys, &wx, &mut want); + assert_eq!(tiled_uni(), want); + let w = time(n, 31, || { + let mut o = [PowerSums::default(); GROUPS]; + masked_group_power_sums_i32(&mask, &keys, black_box(&wx), &mut o); + black_box(o); + }); + let b = time(n, 31, || { + black_box(tiled_uni()); + }); + println!("{:<34}{w:>10.3}{b:>10.3}{:>9.2}", format!("univariate tiled n={n}"), w / b); + + let tiled_bi = || { + let (mut a, mut c, mut out) = ([[0u8; 16]; GROUPS], [[0u8; 16]; GROUPS], [CrossPowerSums::default(); GROUPS]); + fold_bounded_cross_power_sums_tiles(n, &mut a, &mut c, &mut out, |t, a, c| { + masked_group_bounded_cross_power_sums_u8( + &mask[t.start / 64..], + &keys[t.clone()], + &xs[t.clone()], + &ys[t], + a, + c, + ) + }) + .unwrap(); + out + }; + let mut want = [CrossPowerSums::default(); GROUPS]; + masked_group_cross_power_sums_i32(&mask, &keys, &wx, &wy, &mut want); + assert_eq!(tiled_bi(), want); + let w = time(n, 31, || { + let mut o = [CrossPowerSums::default(); GROUPS]; + masked_group_cross_power_sums_i32(&mask, &keys, black_box(&wx), black_box(&wy), &mut o); + black_box(o); + }); + let b = time(n, 31, || { + black_box(tiled_bi()); + }); + println!("{:<34}{w:>10.3}{b:>10.3}{:>9.2}", format!("bivariate tiled n={n}"), w / b); + println!("\ninput bytes/row: wide i32 = 4 per lane, bounded u8 = 1 per lane"); +}