Skip to content

Commit ad8bfa6

Browse files
committed
simd: scatter-or and group-sum kernels accumulate into out, never zero it
A fold kernel adds or ORs into the caller's demanded sink; the caller zeroes that sink once. Whole-buffer zeroing inside the kernel made every call population-sized in writes and broke tiled execution, where the same sink receives one call per tile. mask_gather_u32 is unchanged (its destination is its own output tile). Tests: 6 two-sided accumulation tests (each red with the zeroing restored, verified); parity group 13 gains three from-nonzero checks. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01GXUahz73MZxtxWcfpHp9dG
1 parent 13ef875 commit ad8bfa6

3 files changed

Lines changed: 165 additions & 35 deletions

File tree

‎.claude/blackboard.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ All three are deliberately scalar bit-walks — permutations/scatters indexed by
66
Parity: `check_gather_scatter_group` (0xDxx) in `crates/simd-masking-parity`, against naive per-element references, disable-verified red-then-green.
77
`masked_group_sum_i32_via(mask, index, remap, values, out)` — the same one-pass keyed sum with the key read through a foreign-key hop (`SUM(line.amount) GROUP BY partner.country`); the indirection is fused so no remapped key lane of N is ever materialised. Fourth arm of the same parity group (0xD3x).
88
Consumer: lance-graph-mask-risc `Gather`/`ScatterOr`/`GroupSum` (landing next).
9+
⊘ CORRECTED (operator ruling, same day): `mask_scatter_or_u32`, `masked_group_sum_i32`, and `masked_group_sum_i32_via` no longer zero `out`/`out_words` on entry — a fold kernel accumulates into the caller's demanded sink, and the caller zeroes it once. `mask_gather_u32` is unaffected (its destination is its own output tile). Each function's doc comment now states "accumulates into `out`; the caller zeroes `out` once before the first call"; the surplus/unreferenced-slot assertions in the tail tests were reworded to "untouched, stays as the caller left it" rather than "cleared". Every unit test whose `out` buffer relied on the old zeroing was given an explicit zeroed prefill in its arrange step, and each of the three functions got two new two-sided tests (a preloaded slot/bit that survives untouched alongside the call's own contribution landing correctly; a second call with a different mask summing/unioning on top of the first) — disable-verified red-then-green by temporarily restoring the whole-buffer zero in each function in turn. The parity crate's `check_gather_scatter_group` (0xDxx) now zeroes `out`/`out2` explicitly before its from-zero reference checks and adds one accumulation check per function (0xD11/0xD22/0xD32) that preloads `out` and asserts the call adds on top.
910

1011
## 2026-09-17 (19) — G8 named: a tree-depth column (`lzcnt(bswap(x)) >> 2`) is the missing primitive for basin-local ranking; popcount is only its tie-break
1112

‎crates/simd-masking-parity/src/lib.rs‎

Lines changed: 63 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1163,6 +1163,10 @@ fn check_gather_scatter_group() -> Result<(), u32> {
11631163
}
11641164

11651165
// ── mask_scatter_or_u32 ──────────────────────────────────────────
1166+
// `out2` starts explicitly zeroed (the caller's job now, not the
1167+
// primitive's): the accumulate contract means a dirty prefill would
1168+
// stay dirty rather than being cleared, so parity against a
1169+
// from-zero reference needs a from-zero `out2`.
11661170
let out_rows = 50usize;
11671171
let out_words_count = words_for(out_rows);
11681172
let src_bits: Vec<u64> = (0..nw).map(|_| rng.next()).collect();
@@ -1176,8 +1180,8 @@ fn check_gather_scatter_group() -> Result<(), u32> {
11761180
}
11771181
})
11781182
.collect();
1179-
let out_len2 = out_words_count + 1; // dirty, over-long
1180-
let mut out2 = vec![u64::MAX; out_len2];
1183+
let out_len2 = out_words_count + 1; // over-long, but zeroed not dirty
1184+
let mut out2 = vec![0u64; out_len2];
11811185
mask_scatter_or_u32(&src_bits, &idx2, &mut out2, out_rows);
11821186
let mut want2 = vec![false; out_rows];
11831187
for i in 0..n {
@@ -1193,6 +1197,23 @@ fn check_gather_scatter_group() -> Result<(), u32> {
11931197
return Err(0xD10);
11941198
}
11951199

1200+
// Accumulation check: preload a bit, prove it survives the call
1201+
// unioned with the scatter result — the reference is the union
1202+
// regardless of whether the preloaded bit also happens to be a
1203+
// scatter target, which is exactly what accumulation must produce.
1204+
{
1205+
let preset_bit = out_rows - 1;
1206+
let mut out2_acc = vec![0u64; out_len2];
1207+
out2_acc[preset_bit / 64] |= 1u64 << (preset_bit % 64);
1208+
mask_scatter_or_u32(&src_bits, &idx2, &mut out2_acc, out_rows);
1209+
let mut want2_acc = want2.clone();
1210+
want2_acc[preset_bit] = true;
1211+
let want2_acc_words = reference_mask(out_rows, out_len2, |t| want2_acc[t]);
1212+
if out2_acc != want2_acc_words {
1213+
return Err(0xD11);
1214+
}
1215+
}
1216+
11961217
// ── masked_group_sum_i32 ─────────────────────────────────────────
11971218
let n_groups = 12usize;
11981219
let mask_bits: Vec<u64> = (0..nw).map(|_| rng.next()).collect();
@@ -1207,7 +1228,9 @@ fn check_gather_scatter_group() -> Result<(), u32> {
12071228
})
12081229
.collect();
12091230
let values = i32_values(n, &mut rng);
1210-
let mut group_out = vec![-1i64; n_groups + 1]; // garbage + one unreferenced slot
1231+
// `out` starts explicitly zeroed: the caller's job now, not the
1232+
// primitive's.
1233+
let mut group_out = vec![0i64; n_groups + 1]; // one unreferenced slot
12111234
masked_group_sum_i32(&mask_bits, &keys, &values, &mut group_out);
12121235
let mut want_group = vec![0i64; n_groups];
12131236
for i in 0..n {
@@ -1221,10 +1244,26 @@ fn check_gather_scatter_group() -> Result<(), u32> {
12211244
if group_out[..n_groups] != want_group[..] {
12221245
return Err(0xD20);
12231246
}
1224-
// The unreferenced slot must be zeroed, not left as garbage.
1247+
// The unreferenced slot must stay exactly as the caller left it
1248+
// (zero here), never touched.
12251249
if group_out[n_groups] != 0 {
12261250
return Err(0xD21);
12271251
}
1252+
// Accumulation check: preload every group slot with a known value,
1253+
// prove the call adds its contribution on top rather than resetting.
1254+
{
1255+
let preload = 1_000_000i64;
1256+
let mut group_out_acc = vec![preload; n_groups + 1];
1257+
masked_group_sum_i32(&mask_bits, &keys, &values, &mut group_out_acc);
1258+
for k in 0..n_groups {
1259+
if group_out_acc[k] != preload.wrapping_add(want_group[k]) {
1260+
return Err(0xD22);
1261+
}
1262+
}
1263+
if group_out_acc[n_groups] != preload {
1264+
return Err(0xD22);
1265+
}
1266+
}
12281267

12291268
// ── masked_group_sum_i32_via ─────────────────────────────────────
12301269
// Same n_groups/mask_bits/values as above, but the key is reached
@@ -1251,7 +1290,9 @@ fn check_gather_scatter_group() -> Result<(), u32> {
12511290
}
12521291
})
12531292
.collect();
1254-
let mut via_out = vec![-1i64; n_groups + 1]; // garbage + one unreferenced slot
1293+
// `out` starts explicitly zeroed: the caller's job now, not the
1294+
// primitive's.
1295+
let mut via_out = vec![0i64; n_groups + 1]; // one unreferenced slot
12551296
masked_group_sum_i32_via(&mask_bits, &index, &remap, &values, &mut via_out);
12561297
let mut want_via = vec![0i64; n_groups];
12571298
for i in 0..n {
@@ -1270,9 +1311,26 @@ fn check_gather_scatter_group() -> Result<(), u32> {
12701311
if via_out[..n_groups] != want_via[..] {
12711312
return Err(0xD30);
12721313
}
1314+
// The unreferenced slot must stay exactly as the caller left it
1315+
// (zero here), never touched.
12731316
if via_out[n_groups] != 0 {
12741317
return Err(0xD31);
12751318
}
1319+
// Accumulation check: preload every group slot, prove the call adds
1320+
// its contribution on top rather than resetting.
1321+
{
1322+
let preload = 2_000_000i64;
1323+
let mut via_out_acc = vec![preload; n_groups + 1];
1324+
masked_group_sum_i32_via(&mask_bits, &index, &remap, &values, &mut via_out_acc);
1325+
for k in 0..n_groups {
1326+
if via_out_acc[k] != preload.wrapping_add(want_via[k]) {
1327+
return Err(0xD32);
1328+
}
1329+
}
1330+
if via_out_acc[n_groups] != preload {
1331+
return Err(0xD32);
1332+
}
1333+
}
12761334
}
12771335
Ok(())
12781336
}

‎src/simd_masking_ops.rs‎

Lines changed: 101 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -935,11 +935,12 @@ pub fn mask_gather_u32(src: &[u64], src_rows: usize, index: &[u32], out_words: &
935935
/// (union), so repeats are harmless and order-independent — the same
936936
/// reason `vsa_bundle`-shaped accumulation is safe under reordering.
937937
///
938-
/// `out_words[..out_rows.div_ceil(64)]` is **fully overwritten**, not
939-
/// OR-ed into an existing result: it is zeroed first, then every scattered
940-
/// bit is set. A caller composing this into an accumulating pipeline
941-
/// combines the *result* with `mask_or`/`mask_or_assign`, not by pre-seeding
942-
/// `out_words`.
938+
/// `out_words[..out_rows.div_ceil(64)]` is **accumulated into, not
939+
/// overwritten**: every scattered bit is OR-ed into whatever `out_words`
940+
/// already holds, and a bit set before the call that this call does not
941+
/// itself scatter to stays set. The caller zeroes `out_words` once before
942+
/// the first call in a sequence; repeated calls (e.g. one per source batch)
943+
/// compose as a running union without re-zeroing between them.
943944
///
944945
/// **An out-of-range target (`index[i] >= out_rows`) is silently dropped**,
945946
/// not an error — the same zero-fallback contract as [`mask_gather_u32`]'s
@@ -993,11 +994,8 @@ pub fn mask_scatter_or_u32(src: &[u64], index: &[u32], out_words: &mut [u64], ou
993994
out_word_count
994995
);
995996

996-
// Zero first, whole buffer — same full-overwrite convention as every
997-
// other writer in this module.
998-
for w in out_words.iter_mut() {
999-
*w = 0;
1000-
}
997+
// Accumulate: OR scattered bits into whatever the caller already has in
998+
// `out_words`. The caller zeroes once before the first call.
1001999
for (w, &word) in src.iter().take(src_words).enumerate() {
10021000
let base = w * 64;
10031001
let mut bits = word;
@@ -1037,8 +1035,10 @@ pub fn mask_scatter_or_u32(src: &[u64], index: &[u32], out_words: &mut [u64], ou
10371035
/// selected population and has no notion of a register. Same word
10381036
/// "group", two unrelated shapes; do not conflate them.
10391037
///
1040-
/// `out` is **fully overwritten**, not accumulated into an existing
1041-
/// result: it is zeroed first. **A key at or past `out.len()` is dropped,
1038+
/// `out` is **accumulated into, not overwritten**: contributions are
1039+
/// added to whatever `out` already holds, so the caller zeroes `out` once
1040+
/// before the first call in a sequence rather than this function doing it.
1041+
/// **A key at or past `out.len()` is dropped,
10421042
/// not an error** — the zero-fallback contract shared by
10431043
/// [`mask_gather_u32`]/[`mask_scatter_or_u32`]: a key naming no group in
10441044
/// `out` is not a group, the same way an unminted classid is not a class.
@@ -1100,9 +1100,8 @@ pub fn masked_group_sum_i32(mask_words: &[u64], keys: &[u32], values: &[i32], ou
11001100
words
11011101
);
11021102

1103-
for o in out.iter_mut() {
1104-
*o = 0;
1105-
}
1103+
// Accumulate: add into whatever `out` already holds. The caller zeroes
1104+
// once before the first call.
11061105
for (w, &word) in mask_words.iter().take(words).enumerate() {
11071106
let base = w * 64;
11081107
let mut bits = word;
@@ -1145,9 +1144,10 @@ pub fn masked_group_sum_i32(mask_words: &[u64], keys: &[u32], values: &[i32], ou
11451144
/// to avoid — one fused scan over the selected rows costs no more than the
11461145
/// naive two-hop lookup per selected row, with no second array in between.
11471146
///
1148-
/// `out` is fully overwritten (zeroed first); overflow wraps the same way
1149-
/// as [`masked_group_sum_i32`] (widened to `i64`, `wrapping_add`), and the
1150-
/// mask tail is clamped identically.
1147+
/// `out` is accumulated into (not zeroed by this function) — the caller
1148+
/// zeroes `out` once before the first call, same as [`masked_group_sum_i32`];
1149+
/// overflow wraps the same way as [`masked_group_sum_i32`] (widened to
1150+
/// `i64`, `wrapping_add`), and the mask tail is clamped identically.
11511151
///
11521152
/// # Panics
11531153
///
@@ -1181,9 +1181,8 @@ pub fn masked_group_sum_i32_via(mask_words: &[u64], index: &[u32], remap: &[u32]
11811181
words
11821182
);
11831183

1184-
for o in out.iter_mut() {
1185-
*o = 0;
1186-
}
1184+
// Accumulate: add into whatever `out` already holds. The caller zeroes
1185+
// once before the first call.
11871186
for (w, &word) in mask_words.iter().take(words).enumerate() {
11881187
let base = w * 64;
11891188
let mut bits = word;
@@ -4208,9 +4207,9 @@ mod tests {
42084207
fn mask_scatter_or_u32_empty_index_writes_nothing() {
42094208
let src: [u64; 0] = [];
42104209
let index: [u32; 0] = [];
4211-
let mut out = [0xFFFF_FFFF_FFFF_FFFFu64; 1];
4210+
let mut out = [0u64; 1]; // caller zeroes before the first call
42124211
mask_scatter_or_u32(&src, &index, &mut out, 10);
4213-
assert_eq!(out[0], 0, "no source rows selected ⇒ output cleared, nothing set");
4212+
assert_eq!(out[0], 0, "no source rows selected ⇒ output unchanged, nothing set");
42144213
}
42154214

42164215
#[test]
@@ -4232,10 +4231,12 @@ mod tests {
42324231
.collect();
42334232
let want = bits_to_words(&naive_scatter(&src_bits, &index, out_rows));
42344233
let out_words = out_rows.div_ceil(64);
4235-
let mut out = vec![0xFFFF_FFFF_FFFF_FFFFu64; out_words + 1]; // dirty, over-long
4234+
// Caller zeroes once before the first call; over-long by one word
4235+
// to prove the surplus word is left untouched, not cleared by us.
4236+
let mut out = vec![0u64; out_words + 1];
42364237
mask_scatter_or_u32(&src, &index, &mut out, out_rows);
42374238
assert_eq!(&out[..want.len()], &want[..], "scatter mismatch at n={n}");
4238-
assert_eq!(out[out_words], 0, "surplus word must be cleared at n={n}");
4239+
assert_eq!(out[out_words], 0, "surplus word is untouched (stays as the caller left it) at n={n}");
42394240
}
42404241
}
42414242

@@ -4266,6 +4267,26 @@ mod tests {
42664267
assert_eq!(out[0], 1u64 << 3);
42674268
}
42684269

4270+
#[test]
4271+
fn mask_scatter_or_u32_accumulates_into_a_preloaded_out_buffer() {
4272+
// A bit the call never scatters to must survive; the call's own
4273+
// contribution must land alongside it. This must FAIL if the old
4274+
// whole-buffer zeroing is restored (bit 7 would be cleared).
4275+
let src = [0b1u64]; // row 0 selected
4276+
let index = [3u32]; // scatters to bit 3
4277+
let mut out = [1u64 << 7]; // pre-set bit 7, not touched by this call
4278+
mask_scatter_or_u32(&src, &index, &mut out, 8);
4279+
assert_eq!(out[0], (1u64 << 3) | (1u64 << 7), "pre-existing bit 7 must survive alongside the new bit 3");
4280+
}
4281+
4282+
#[test]
4283+
fn mask_scatter_or_u32_a_second_call_with_a_different_mask_unions_on_top() {
4284+
let mut out = [0u64; 1];
4285+
mask_scatter_or_u32(&[0b1u64], &[2u32], &mut out, 8);
4286+
mask_scatter_or_u32(&[0b1u64], &[5u32], &mut out, 8);
4287+
assert_eq!(out[0], (1u64 << 2) | (1u64 << 5), "two calls union, the second does not erase the first");
4288+
}
4289+
42694290
#[test]
42704291
#[should_panic(expected = "out_words.len()")]
42714292
fn mask_scatter_or_u32_rejects_short_out_buffer() {
@@ -4302,7 +4323,7 @@ mod tests {
43024323
let mask: [u64; 0] = [];
43034324
let keys: [u32; 0] = [];
43044325
let values: [i32; 0] = [];
4305-
let mut out = [123i64; 4]; // garbage, must be cleared
4326+
let mut out = [0i64; 4]; // caller zeroes before the first call
43064327
masked_group_sum_i32(&mask, &keys, &values, &mut out);
43074328
assert_eq!(out, [0, 0, 0, 0]);
43084329
}
@@ -4327,10 +4348,15 @@ mod tests {
43274348
// Signed values spanning both sides of zero.
43284349
let values: Vec<i32> = (0..n).map(|_| (splitmix(&mut seed) as i32) / 2).collect();
43294350
let want = naive_group_sum(&mask_bits, &keys, &values, n_groups);
4330-
let mut out = vec![-999i64; n_groups + 1]; // garbage, and one extra slot no key ever hits
4351+
// Caller zeroes once before the first call; one extra slot no
4352+
// key ever hits, to prove it is left untouched, not cleared.
4353+
let mut out = vec![0i64; n_groups + 1];
43314354
masked_group_sum_i32(&mask, &keys, &values, &mut out);
43324355
assert_eq!(&out[..n_groups], &want[..], "group sum mismatch at n={n}");
4333-
assert_eq!(out[n_groups], 0, "an unreferenced group slot must be zero, not garbage, at n={n}");
4356+
assert_eq!(
4357+
out[n_groups], 0,
4358+
"an unreferenced group slot is untouched (stays as the caller left it) at n={n}"
4359+
);
43344360
}
43354361
}
43364362

@@ -4365,6 +4391,28 @@ mod tests {
43654391
assert_eq!(out, [10, 20], "the out-of-range key contributes nothing");
43664392
}
43674393

4394+
#[test]
4395+
fn masked_group_sum_i32_accumulates_into_a_preloaded_out_buffer() {
4396+
// Slot 1 starts pre-loaded and is never touched by this call; slot 0
4397+
// starts pre-loaded and IS the call's target. This must FAIL if the
4398+
// old whole-buffer zeroing is restored (both would reset to 0 first).
4399+
let mask = [0b1u64]; // row 0 selected
4400+
let keys = [0u32]; // routes to slot 0
4401+
let values = [7i32];
4402+
let mut out = [5i64, 42i64]; // slot 0 preloaded 5, slot 1 preloaded 42
4403+
masked_group_sum_i32(&mask, &keys, &values, &mut out);
4404+
assert_eq!(out[0], 12, "slot 0: preload 5 + contribution 7 = 12");
4405+
assert_eq!(out[1], 42, "slot 1 is untouched, must survive exactly as preloaded");
4406+
}
4407+
4408+
#[test]
4409+
fn masked_group_sum_i32_a_second_call_with_a_different_mask_sums_on_top() {
4410+
let mut out = [0i64; 2];
4411+
masked_group_sum_i32(&[0b1u64], &[0u32], &[10i32], &mut out);
4412+
masked_group_sum_i32(&[0b1u64], &[0u32], &[3i32], &mut out);
4413+
assert_eq!(out, [13, 0], "two calls sum, the second does not erase the first's contribution");
4414+
}
4415+
43684416
#[test]
43694417
#[should_panic(expected = "keys/values length mismatch")]
43704418
fn masked_group_sum_i32_rejects_mismatched_keys_and_values() {
@@ -4408,7 +4456,7 @@ mod tests {
44084456
let mut want = vec![0i64; n_groups];
44094457
masked_group_sum_i32(&mask, &keys, &values, &mut want);
44104458

4411-
let mut got = vec![-1i64; n_groups];
4459+
let mut got = vec![0i64; n_groups]; // caller zeroes before the first call
44124460
masked_group_sum_i32_via(&mask, &index, &remap, &values, &mut got);
44134461
assert_eq!(got, want, "via mismatched the plain two-hop-materialised form at n={n}");
44144462
}
@@ -4484,11 +4532,34 @@ mod tests {
44844532
want[k] = want[k].wrapping_add(values[i] as i64);
44854533
}
44864534
}
4487-
let mut got = vec![-1i64; n_groups];
4535+
let mut got = vec![0i64; n_groups]; // caller zeroes before the first call
44884536
masked_group_sum_i32_via(&mask, &index, &remap, &values, &mut got);
44894537
assert_eq!(got, want);
44904538
}
44914539

4540+
#[test]
4541+
fn masked_group_sum_i32_via_accumulates_into_a_preloaded_out_buffer() {
4542+
// Slot 1 is preloaded and never targeted; slot 0 is preloaded and IS
4543+
// the resolved target. Must FAIL if whole-buffer zeroing returns.
4544+
let mask = [0b1u64]; // row 0 selected
4545+
let index = [0u32]; // partner 0
4546+
let remap = [0u32]; // partner 0 -> group 0
4547+
let values = [7i32];
4548+
let mut out = [5i64, 42i64];
4549+
masked_group_sum_i32_via(&mask, &index, &remap, &values, &mut out);
4550+
assert_eq!(out[0], 12, "group 0: preload 5 + contribution 7 = 12");
4551+
assert_eq!(out[1], 42, "group 1 is untouched, must survive exactly as preloaded");
4552+
}
4553+
4554+
#[test]
4555+
fn masked_group_sum_i32_via_a_second_call_with_a_different_mask_sums_on_top() {
4556+
let remap = [0u32];
4557+
let mut out = [0i64; 1];
4558+
masked_group_sum_i32_via(&[0b1u64], &[0u32], &remap, &[10i32], &mut out);
4559+
masked_group_sum_i32_via(&[0b1u64], &[0u32], &remap, &[3i32], &mut out);
4560+
assert_eq!(out, [13], "two calls sum, the second does not erase the first's contribution");
4561+
}
4562+
44924563
#[test]
44934564
#[should_panic(expected = "index/values length mismatch")]
44944565
fn masked_group_sum_i32_via_rejects_mismatched_index_and_values() {

0 commit comments

Comments
 (0)