diff --git a/crates/simd-masking-parity/src/lib.rs b/crates/simd-masking-parity/src/lib.rs index 59830e12..9fe7303f 100644 --- a/crates/simd-masking-parity/src/lib.rs +++ b/crates/simd-masking-parity/src/lib.rs @@ -39,7 +39,8 @@ use ndarray::simd::{ lt_u8_to_mask, mask_all, mask_and, mask_and_assign, mask_andnot, mask_andnot_assign, mask_any, mask_gather_u32, mask_not, mask_not_assign, mask_or, mask_or_assign, mask_scatter_or_u32, mask_set_range, mask_shift_morton, mask_ternlog, mask_ternlog_assign, mask_xor, mask_xor_assign, masked_group_sum_i32, masked_group_sum_i32_via, - masked_key_run_count_u32, masked_max_i32, masked_min_i32, masked_strided_group_sum, masked_sum_i32, ne_i32_to_mask, + masked_key_run_count_u32, masked_max_i32, masked_min_i32, masked_strided_group_sum, masked_sum_i32, + masked_sum_wrapping_add_i32, ne_i32_to_mask, ne_i32_to_mask_under, ne_u32_to_mask, ne_u32_to_mask_under, ne_u64_to_mask, ne_u8_to_mask, ternary_match_strided_to_mask, ternary_match_u32_to_mask, ternary_match_u32_to_mask_under, ternary_match_u64_to_mask, ternary_match_u64_to_mask_under, ternlog, I32x16, KeyRunCarry, MortonDir, U32x16, U64x8, @@ -758,6 +759,7 @@ fn check_masked_reductions() -> Result<(), u32> { for &n in &LENS { let nw = words_for(n); let vals = i32_values(n, &mut rng); + let rhs = i32_values(n, &mut rng); let tail_mask = |i: usize| -> u64 { let live = n.saturating_sub(i * 64).min(64); if live == 64 { @@ -784,6 +786,14 @@ fn check_masked_reductions() -> Result<(), u32> { if masked_sum_i32(&vals, m) != want_sum { return Err(0x800 | k); } + let want_wrapping_add = (0..n) + .filter(|&i| (m[i / 64] >> (i % 64)) & 1 == 1) + .fold(0i64, |acc, i| { + acc.wrapping_add(vals[i].wrapping_add(rhs[i]) as i64) + }); + if masked_sum_wrapping_add_i32(&vals, &rhs, m) != want_wrapping_add { + return Err(0x850 | k); + } if masked_min_i32(&vals, m) != selected.iter().copied().min() { return Err(0x810 | k); } diff --git a/src/simd.rs b/src/simd.rs index a55f81ce..737df036 100644 --- a/src/simd.rs +++ b/src/simd.rs @@ -833,6 +833,7 @@ pub use crate::simd_masking_ops::{ masked_min_i32, masked_strided_group_sum, masked_sum_i32, + masked_sum_wrapping_add_i32, ne_i32_to_mask, ne_i32_to_mask_under, ne_u32_to_mask, diff --git a/src/simd_masking_ops.rs b/src/simd_masking_ops.rs index e4f55be4..ace210d6 100644 --- a/src/simd_masking_ops.rs +++ b/src/simd_masking_ops.rs @@ -707,6 +707,70 @@ pub fn masked_sum_i32(values: &[i32], mask_words: &[u64]) -> i64 { acc } +/// Sum `a[i].wrapping_add(b[i])` where mask bit `i` is set, with the +/// wrapped row value widened to `i64` before accumulation. +/// +/// This is the semantics-preserving fused form of a wrapping `i32` row +/// expression followed by [`masked_sum_i32`]. It exists so a consumer can +/// fold `SUM(a + b)` without materialising the derived `i32` lane and +/// without changing the program to `SUM(a) + SUM(b)`, which is not +/// equivalent when any selected row overflows `i32`. +/// +/// Row arithmetic is deliberately **wrapping `i32` first, widening second**: +/// `i32::MAX + 1` contributes `i32::MIN as i64`, not `2^31`. Once each +/// row has been reduced to one `i32`, accumulation has the same `i64` +/// bound and theoretical wrapping behavior as [`masked_sum_i32`]. +/// +/// Mask bits past `a.len()` are ignored. `a` and `b` must have equal +/// length; `mask_words` must cover that length. +/// +/// # Panics +/// +/// Panics if `a.len() != b.len()` or if `mask_words` is too short. +/// +/// # Examples +/// +/// ``` +/// use ndarray::simd::masked_sum_wrapping_add_i32; +/// +/// let a = [i32::MAX, 10]; +/// let b = [1, 20]; +/// // Both rows selected: wrapping MAX+1 = MIN, then +30 in widened i64. +/// assert_eq!( +/// masked_sum_wrapping_add_i32(&a, &b, &[0b11]), +/// i64::from(i32::MIN) + 30 +/// ); +/// ``` +#[inline] +pub fn masked_sum_wrapping_add_i32(a: &[i32], b: &[i32], mask_words: &[u64]) -> i64 { + assert_eq!(a.len(), b.len(), "masked_sum_wrapping_add_i32: a/b length mismatch"); + let n = a.len(); + let words = mask_words_for(n); + assert!( + mask_words.len() >= words, + "masked_sum_wrapping_add_i32: mask_words.len()={} < required {}", + mask_words.len(), + words + ); + + let mut acc: i64 = 0; + for (w, &word) in mask_words.iter().take(words).enumerate() { + let base = w * 64; + let mut bits = word; + let valid = n - base; + if valid < 64 { + bits &= (1u64 << valid) - 1; + } + while bits != 0 { + let lane = bits.trailing_zeros() as usize; + let i = base + lane; + acc = acc.wrapping_add(a[i].wrapping_add(b[i]) as i64); + bits &= bits - 1; + } + } + acc +} + /// Sum a sub-word group field out of a **strided** record, over the records a /// mask selects, widened to `i128` and range-checked into `i64`. ///