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
16 changes: 15 additions & 1 deletion crates/simd-masking-parity/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ use ndarray::simd::{
eq_u32_via_to_mask, eq_u64_to_mask, eq_u8_to_mask, ge_i32_to_mask, ge_i32_to_mask_under, ge_u64_to_mask,
ge_u8_to_mask, gt_i32_to_mask, gt_i32_to_mask_under, gt_u64_to_mask, gt_u8_to_mask, le_i32_to_mask,
le_i32_to_mask_under, le_u64_to_mask, le_u8_to_mask, lt_i32_to_mask, lt_i32_to_mask_under, lt_u64_to_mask,
lt_u8_to_mask, mask_all, mask_and, mask_and_assign, mask_andnot, mask_andnot_assign, mask_any, mask_gather_u32,
lt_u8_to_mask, mask_all, mask_and, mask_and_assign, mask_andnot, mask_andnot_assign, mask_any, mask_gather_u32, mask_gather_u32_under,
mask_not, mask_not_assign, mask_or, mask_or_assign, mask_scatter_or_u32, mask_set_range, mask_shift_morton,
mask_ternlog, mask_ternlog_any, mask_ternlog_assign, mask_ternlog_popcount, 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,
Expand Down Expand Up @@ -1269,6 +1269,20 @@ fn check_gather_scatter_group() -> Result<(), u32> {
return Err(0xD00);
}

// ── mask_gather_u32_under ────────────────────────────────────────
// Gate tail bits past `n` are left dirty on purpose: a kernel that
// used them would address `index` out of bounds.
let gate: Vec<u64> = (0..out_len).map(|_| rng.next()).collect();
let mut out = vec![u64::MAX; out_len];
mask_gather_u32_under(&src, src_rows, &index, &gate, &mut out);
let want = reference_mask(n, out_len, |i| {
let idx = index[i] as usize;
(gate[i / 64] >> (i % 64)) & 1 == 1 && idx < src_rows && (src[idx / 64] >> (idx % 64)) & 1 == 1
});
if out != want {
return Err(0xD01);
}

// ── mask_scatter_or_u32 ──────────────────────────────────────────
// `out2` starts explicitly zeroed (the caller's job now, not the
// primitive's): the accumulate contract means a dirty prefill would
Expand Down
1 change: 1 addition & 0 deletions src/simd.rs
Original file line number Diff line number Diff line change
Expand Up @@ -818,6 +818,7 @@ pub use crate::simd_masking_ops::{
mask_andnot_assign,
mask_any,
mask_gather_u32,
mask_gather_u32_under,
mask_not,
mask_not_assign,
mask_or,
Expand Down
196 changes: 192 additions & 4 deletions src/simd_masking_ops.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1076,6 +1076,20 @@ pub fn masked_strided_group_sum(
/// no input mask to skip zero words against — every entry of `index` must
/// be consulted regardless of what `src` contains.
///
/// **Branch-free per row.** Each row's bit is shifted into the output word
/// unconditionally; the out-of-range test becomes a select on the address
/// plus an AND on the bit (see `gather_bit`). The earlier form branched on
/// `idx < src_rows && bit == 1`. Whether that compiles to a branch depends
/// on the call site. Measured on a 1M-row lane with an unpredictable source:
/// called through lance-graph's mask-risc `Gather`, it ran 5.3 ms and the
/// branch-free form runs 1.5 ms, level with a hand-written branch-free loop
/// (`D-GATED-GATHER-0`). Inlined into a plain loop, the compiler had already
/// removed the branch, and the two forms time the same. The branch-free form
/// removes the dependence on that choice.
///
/// When the caller already holds the rows it wants — a conjunct's survivors —
/// use [`mask_gather_u32_under`], which reads `index` only at those rows.
///
/// # Panics
///
/// Panics if `out_words.len() < index.len().div_ceil(64)`, or if
Expand Down Expand Up @@ -1113,20 +1127,116 @@ pub fn mask_gather_u32(src: &[u64], src_rows: usize, index: &[u32], out_words: &
for w in out_words.iter_mut() {
*w = 0;
}
// With no source rows every index is out of range, and `src` may be
// empty, so `gather_bit`'s fallback read of `src[0]` would not exist.
if src_rows == 0 {
return;
}
for (w, out_word) in out_words.iter_mut().enumerate().take(words) {
let base = w * 64;
let live = (n - base).min(64);
let mut acc = 0u64;
for lane in 0..live {
let idx = index[base + lane] as usize;
if idx < src_rows && (src[idx / 64] >> (idx % 64)) & 1 == 1 {
acc |= 1u64 << lane;
}
acc |= gather_bit(src, src_rows, index[base + lane] as usize) << lane;
}
*out_word = acc;
}
}

/// Bit `idx` of `src`, or 0 when `idx >= src_rows` — without a branch on the
/// bit's value. An out-of-range address is redirected to row 0 by a select
/// and its result is cleared by `& ok`, so the read is always in bounds.
///
/// Requires `src_rows >= 1` (so `src[0]` exists); both gathers return early
/// on `src_rows == 0` before calling this.
#[inline(always)]
fn gather_bit(src: &[u64], src_rows: usize, idx: usize) -> u64 {
let ok = (idx < src_rows) as u64;
let j = if ok == 1 { idx } else { 0 };
((src[j / 64] >> (j % 64)) & 1) & ok
}

/// Gated gather: bit `i` of the result is `under[i] && src[index[i]]`, with
/// `index[i]` read **only** for rows whose `under` bit is set.
///
/// Same addressing and the same out-of-range-is-false contract as
/// [`mask_gather_u32`]; the result equals `mask_gather_u32(..) & under`. The
/// difference is the schedule: a zero `under` word costs one test, and inside
/// a live word only the set bits are visited (`trailing_zeros`), so the work
/// is proportional to `popcount(under)`, not to `index.len()`. This is the
/// semijoin a conjunction needs — `gate ∧ gather(fk, foreign)` — without
/// gathering the rows the gate has already rejected.
///
/// Measured (lance-graph `D-GATED-GATHER-0`, 1M rows): fastest or within
/// noise of fastest at every gate density from 0.01 % to 100 %, and at 100 %
/// no slower than the full branch-free gather. A word-granular gate (gather
/// all 64 lanes of every live word) is not a substitute: on scattered 1 %
/// survivors it issues ~47× the loads.
///
/// `under` bits at or past `index.len()` are ignored, never used to address
/// `index`. `out_words` is **fully overwritten**: every bit past
/// `index.len()` and every surplus word is written `0`.
///
/// # Panics
///
/// Panics if `out_words.len()` or `under.len()` is below
/// `index.len().div_ceil(64)`, or if `src.len() < src_rows.div_ceil(64)`.
///
/// # Examples
///
/// ```
/// use ndarray::simd::{mask_gather_u32, mask_gather_u32_under};
///
/// let src = [0b101u64]; // rows 0 and 2 of a 3-row source are set
/// let index = [2u32, 1, 0, 5];
/// let mut full = [0u64; 1];
/// mask_gather_u32(&src, 3, &index, &mut full);
/// assert_eq!(full[0], 0b0101);
/// // keep only rows 2 and 3 of `index`: row 2 reads src row 0 (set),
/// // row 3 is out of range (false)
/// let mut gated = [u64::MAX; 1];
/// mask_gather_u32_under(&src, 3, &index, &[0b1100], &mut gated);
/// assert_eq!(gated[0], 0b0100);
/// ```
#[inline]
pub fn mask_gather_u32_under(src: &[u64], src_rows: usize, index: &[u32], under: &[u64], out_words: &mut [u64]) {
let n = index.len();
let words = mask_words_for(n);
assert!(
out_words.len() >= words,
"mask_gather_u32_under: out_words.len()={} < required {}",
out_words.len(),
words
);
assert!(under.len() >= words, "mask_gather_u32_under: under.len()={} < required {}", under.len(), words);
let src_words = mask_words_for(src_rows);
assert!(src.len() >= src_words, "mask_gather_u32_under: src.len()={} < required {}", src.len(), src_words);

for w in out_words.iter_mut() {
*w = 0;
}
if src_rows == 0 {
return;
}
for w in 0..words {
let base = w * 64;
let mut bits = under[w];
// A gate bit past `index.len()` names no row: drop it before it can
// address `index` out of bounds.
let valid = n - base;
if valid < 64 {
bits &= (1u64 << valid) - 1;
}
let mut acc = 0u64;
while bits != 0 {
let lane = bits.trailing_zeros() as usize;
bits &= bits - 1;
acc |= gather_bit(src, src_rows, index[base + lane] as usize) << lane;
}
out_words[w] = acc;
}
}

/// Scatter-OR by row index: for every element `i` selected by `src` (bit `i`
/// set, `i < index.len()`), sets bit `index[i]` of the result — provided
/// `index[i] < out_rows`.
Expand Down Expand Up @@ -6246,6 +6356,84 @@ mod tests {
mask_gather_u32(&src, 65, &index, &mut out); // claims 65 rows
}

#[test]
fn mask_gather_u32_zero_src_rows_is_all_false_even_with_empty_src() {
// `src_rows == 0` needs no source words, so `src` may be empty: the
// branch-free fallback read of `src[0]` must not happen.
let index = [0u32, 1, 7];
let mut out = [u64::MAX; 1];
mask_gather_u32(&[], 0, &index, &mut out);
assert_eq!(out[0], 0);
let mut gated = [u64::MAX; 1];
mask_gather_u32_under(&[], 0, &index, &[u64::MAX], &mut gated);
assert_eq!(gated[0], 0);
}

#[test]
fn mask_gather_u32_under_matches_gather_and_gate_across_the_tail() {
for &n in &[0usize, 1, 63, 64, 65, 67, 130] {
let mut seed = 0x2468_ACE0_1357_9BDFu64;
let src_rows = 40usize;
let src: Vec<u64> = (0..src_rows.div_ceil(64))
.map(|_| splitmix(&mut seed))
.collect();
let index: Vec<u32> = (0..n)
.map(|i| {
if i % 3 == 0 {
(src_rows as u64 + 5 + i as u64) as u32
} else {
(splitmix(&mut seed) % src_rows as u64) as u32
}
})
.collect();
let out_words = n.div_ceil(64).max(1);
// Gates: empty, full, random — each with its tail bits past `n`
// deliberately set, so a missing clamp reads `index` out of range.
let gates: [Vec<u64>; 3] = [
vec![0u64; out_words],
vec![u64::MAX; out_words],
(0..out_words).map(|_| splitmix(&mut seed)).collect(),
];
for gate in &gates {
let want: Vec<bool> = naive_gather(&src, src_rows, &index)
.iter()
.enumerate()
.map(|(i, &b)| b && (gate[i / 64] >> (i % 64)) & 1 == 1)
.collect();
let want = bits_to_words(&want);
let mut out = vec![u64::MAX; out_words + 1];
mask_gather_u32_under(&src, src_rows, &index, gate, &mut out);
assert_eq!(&out[..want.len()], &want[..], "gated gather mismatch at n={n}");
assert_eq!(out[out_words], 0, "surplus word must be cleared at n={n}");
}
}
}

#[test]
fn mask_gather_u32_under_keeps_only_gated_rows() {
// Every row's address selects a set source bit, so the gate alone
// decides the result. (Whether ungated rows are READ is not
// observable through this API; the dirty-tail gates in the parity
// test above are what prove a gate bit past `index.len()` never
// addresses `index`.)
let src = [u64::MAX]; // 64 rows, all set
let index = [0u32, 1, 2, 3];
let mut out = [0u64; 1];
mask_gather_u32_under(&src, 64, &index, &[0b0100], &mut out);
assert_eq!(out[0], 0b0100, "only the gated row survives");
mask_gather_u32_under(&src, 64, &index, &[0], &mut out);
assert_eq!(out[0], 0, "an empty gate gathers nothing");
}

#[test]
#[should_panic(expected = "under.len()")]
fn mask_gather_u32_under_rejects_short_gate() {
let src = [0u64; 1];
let index = vec![0u32; 65];
let mut out = [0u64; 2];
mask_gather_u32_under(&src, 1, &index, &[u64::MAX], &mut out);
}

fn naive_scatter(src_bits: &[bool], index: &[u32], out_rows: usize) -> Vec<bool> {
let mut out = vec![false; out_rows];
for (i, &selected) in src_bits.iter().enumerate() {
Expand Down
Loading