Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion crates/simd-masking-parity/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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 {
Expand All @@ -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);
}
Expand Down
1 change: 1 addition & 0 deletions src/simd.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
64 changes: 64 additions & 0 deletions src/simd_masking_ops.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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`.
///
Expand Down
Loading