Skip to content

Commit 07e044f

Browse files
authored
Merge pull request #83 from AdaWorldAPI/claude/fold-distillation-pr-wave-s57uj7
lgj-abi: status arms for the new mask-risc ExecError variants (Wave 4 base)
2 parents e1f909f + 4cae6da commit 07e044f

3 files changed

Lines changed: 68 additions & 32 deletions

File tree

‎native/lgj-abi/src/exports.rs‎

Lines changed: 48 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -26,8 +26,8 @@ use std::cell::RefCell;
2626
use std::panic::{catch_unwind, AssertUnwindSafe};
2727

2828
use lance_graph_mask_risc::{
29-
execute, reference_execute, reference_scratch, scratch_words_for, ExecError, LaneRef, Planes,
30-
Scratch, Value,
29+
execute_into, reference_execute_into, scratch_words_for, tile_words_for, ExecError, Foreign,
30+
LaneRef, Out, Planes, Scratch, Value,
3131
};
3232

3333
use crate::abi::*;
@@ -1733,8 +1733,15 @@ fn exec_error_to_status(e: ExecError) -> i32 {
17331733
match e {
17341734
ExecError::LaneOutOfRange(_) => LGJ_ERR_INVALID_LANE,
17351735
ExecError::LaneKind { .. } => LGJ_ERR_LANE_KIND_MISMATCH,
1736+
// The lane is the right width but not in key order: a lane-shape
1737+
// precondition of the fold, reported as the same class of mismatch.
1738+
ExecError::LaneNotOrdered { .. } => LGJ_ERR_LANE_KIND_MISMATCH,
17361739
ExecError::LenMismatch { .. } => LGJ_ERR_MASK_LENGTH_MISMATCH,
1737-
ExecError::PlaneOutOfRange(_) | ExecError::PlaneTail(_) => LGJ_ERR_INVALID_HANDLE,
1740+
ExecError::PlaneOutOfRange(_)
1741+
| ExecError::PlaneTail(_)
1742+
| ExecError::ForeignOutOfRange(_) => LGJ_ERR_INVALID_HANDLE,
1743+
ExecError::ForeignLaneOutOfRange(_) => LGJ_ERR_INVALID_LANE,
1744+
ExecError::ForeignLaneKind { .. } => LGJ_ERR_LANE_KIND_MISMATCH,
17381745
ExecError::SumRowBound { .. } => LGJ_ERR_SUM_OVERFLOW,
17391746
ExecError::ScratchTooSmall { .. }
17401747
| ExecError::ScratchBufferTooSmall { .. }
@@ -1743,6 +1750,7 @@ fn exec_error_to_status(e: ExecError) -> i32 {
17431750
| ExecError::ScratchReadBeforeWrite { .. }
17441751
| ExecError::ScratchWords { .. }
17451752
| ExecError::BlendNeedsOut
1753+
| ExecError::TerminalNeedsOut { .. }
17461754
| ExecError::GateAliasesDst { .. }
17471755
| ExecError::RangeOutOfBounds { .. } => LGJ_ERR_ALLOCATION_FAILED,
17481756
}
@@ -1864,33 +1872,53 @@ fn plan_eval_impl(
18641872

18651873
match path {
18661874
Path::Simd => {
1867-
let need = match scratch_words_for(n_words, plan_lower::SLOTS as usize) {
1875+
// TILED: the scratch is `SLOTS × tile_words_for(rows)` words
1876+
// however many rows the resource holds, and the kept mask is
1877+
// written tile by tile into the destination resource's own words
1878+
// — the demanded sink, and the only population-sized write.
1879+
// Validation is total before the first tile, so an error leaves
1880+
// `dst_mask` byte-for-byte as it was, exactly as before.
1881+
let tile = tile_words_for(rows);
1882+
let need = match scratch_words_for(tile, plan_lower::SLOTS as usize) {
18681883
Some(n) => n,
18691884
None => return LGJ_ERR_LENGTH_OVERFLOW,
18701885
};
1886+
let mut g = match mask.write_mask() {
1887+
Some(g) => g,
1888+
None => return LGJ_ERR_WRONG_RESOURCE_KIND,
1889+
};
1890+
if g.words.len() != n_words {
1891+
return LGJ_ERR_MASK_LENGTH_MISMATCH;
1892+
}
18711893
PLAN_SCRATCH.with(|cell| {
18721894
let mut buf = cell.borrow_mut();
18731895
if buf.len() < need {
18741896
buf.resize(need, 0);
18751897
}
18761898
let mut scratch =
1877-
match Scratch::over(&mut buf[..need], n_words, plan_lower::SLOTS as usize) {
1899+
match Scratch::over(&mut buf[..need], tile, plan_lower::SLOTS as usize) {
18781900
Ok(s) => s,
18791901
Err(e) => return exec_error_to_status(e),
18801902
};
1881-
let count = match execute(&program, &planes, &mut scratch, None) {
1882-
Ok(Value::Count(c)) => c as u64,
1903+
match execute_into(
1904+
&program,
1905+
&planes,
1906+
&Foreign::NONE,
1907+
&mut scratch,
1908+
Out::Mask(&mut g.words),
1909+
) {
1910+
Ok(Value::Mask(_)) => {}
18831911
// The lowering emits exactly one terminal and it is
1884-
// `Count`; any other value means this file built a
1912+
// `Keep`; any other value means this file built a
18851913
// program it did not intend to.
18861914
Ok(_) => return LGJ_ERR_ALLOCATION_FAILED,
18871915
Err(e) => return exec_error_to_status(e),
1888-
};
1889-
let words = match scratch.slot(plan_lower::ACC_SLOT) {
1890-
Some(w) => w,
1891-
None => return LGJ_ERR_ALLOCATION_FAILED,
1892-
};
1893-
publish(&mask, words, count, out_count)
1916+
}
1917+
let count = kernels::simd_popcount(&g.words);
1918+
drop(g);
1919+
// SAFETY: non-null (checked by the caller); written only on success.
1920+
unsafe { *out_count = count };
1921+
LGJ_OK
18941922
})
18951923
}
18961924
// Fork A: the scalar symbol runs mask-risc's row-at-a-time oracle, so
@@ -1900,20 +1928,14 @@ fn plan_eval_impl(
19001928
// bit-packing bug — which is why the allocation gate names
19011929
// `lgj_plan_eval` and only it.
19021930
Path::Scalar => {
1903-
let count = match reference_execute(&program, &planes, None) {
1904-
Ok(Value::Count(c)) => c as u64,
1931+
let mut kept = vec![0u64; n_words];
1932+
match reference_execute_into(&program, &planes, &Foreign::NONE, Out::Mask(&mut kept)) {
1933+
Ok(Value::Mask(_)) => {}
19051934
Ok(_) => return LGJ_ERR_ALLOCATION_FAILED,
19061935
Err(e) => return exec_error_to_status(e),
1907-
};
1908-
let slots = match reference_scratch(&program, &planes) {
1909-
Ok(s) => s,
1910-
Err(e) => return exec_error_to_status(e),
1911-
};
1912-
let words = match slots.get(plan_lower::ACC_SLOT as usize) {
1913-
Some(w) => w.as_slice(),
1914-
None => return LGJ_ERR_ALLOCATION_FAILED,
1915-
};
1916-
publish(&mask, words, count, out_count)
1936+
}
1937+
let count = kernels::simd_popcount(&kept);
1938+
publish(&mask, &kept, count, out_count)
19171939
}
19181940
}
19191941
}

‎native/lgj-abi/src/exports/tests/lowering_convergence.rs‎

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -156,12 +156,22 @@ fn cmp_of(o: &LgjOpDesc) -> Cmp {
156156
/// function on both sides is the whole point of comparing LOWERINGS rather
157157
/// than comparing EXECUTIONS.
158158
fn run(p: &Program, planes: &Planes) -> usize {
159-
let words = planes.n_rows.div_ceil(64);
160-
let slots = p.scratch_slots as usize;
161-
let mut buf = vec![0u64; scratch_words_for(words, slots).expect("sized")];
162-
let mut scratch = Scratch::over(&mut buf, words, slots).expect("carves");
163-
match execute(p, planes, &mut scratch, None).expect("runs") {
159+
// The tiled default scratch. This crate's own lowering ends in `Keep`
160+
// (the destination mask is the demanded sink; the count is its
161+
// popcount), the quack lowering in `Count` — both are read here.
162+
let mut scratch = Scratch::for_program(p, planes.n_rows).expect("addressable");
163+
let mut kept = vec![0u64; planes.n_rows.div_ceil(64)];
164+
match execute_into(
165+
p,
166+
planes,
167+
&Foreign::NONE,
168+
&mut scratch,
169+
Out::Mask(&mut kept),
170+
)
171+
.expect("runs")
172+
{
164173
Value::Count(c) => c,
174+
Value::Mask(_) => ndarray::simd::popcount_batch_u64(&kept) as usize,
165175
other => panic!("not a count: {other:?}"),
166176
}
167177
}

‎native/lgj-abi/src/plan_lower.rs‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -96,9 +96,13 @@ pub(crate) fn lower_plan(ops: &[LgjOpDesc]) -> Option<Lowered> {
9696
});
9797
}
9898

99+
// `Keep`, not `Count`: the destination mask IS the demanded result, so
100+
// the executor writes it tile by tile straight into the resource's own
101+
// words (`Out::Mask`) and the count is a popcount over that sink. Under
102+
// tiled execution no population-sized scratch exists to count from.
99103
Some(Lowered::Program(Program::new(
100104
program_ops,
101-
Terminal::Count {
105+
Terminal::Keep {
102106
mask: Operand::Scratch(ACC_SLOT),
103107
},
104108
)))

0 commit comments

Comments
 (0)