Store record batches in caller-owned frame buffers instead of aliased stack pointers

This commit is contained in:
Dennis Kobert
2026-08-16 11:29:20 +00:00
parent 9055a5f3d5
commit ac40f1838f
6 changed files with 292 additions and 289 deletions
+136 -269
View File
@@ -6,153 +6,148 @@ use std::mem::MaybeUninit;
use std::ops::Range; use std::ops::Range;
#[derive(Debug)] #[derive(Debug)]
pub enum BatchStatus<'a, T> { pub enum BatchStatus<'a> {
Lent(RecordBatch<'a, T>, Finality), /// Producer-resident lanes, shared: read-only for the caller.
Filled(RecordBatch<'a, T>, Finality), Lent(RecordBatch<'a>, Finality),
/// The caller's scratch, filled: the caller is the exclusive owner and may
/// mutate the lanes or reclaim the buffer for in-place reuse.
Filled(RecordBatchMut<'a>, Finality),
/// No batch implementation behind this edge; a driver answers with the
/// per-lane eval and copy-out loop ([`crate::record::fill_frames`]).
Unbatched,
Pending, Pending,
Error(GraphError), Error(GraphError),
NeedBuffer, NeedBuffer,
InvalidRange, InvalidRange,
} }
/// Owns the initialized prefix of a caller-supplied scratch buffer, dropping every /// A shared view over a batch of records in one flat frame buffer: lane `i`
/// lane unless [`FilledBatch::into_values`] hands the obligation back to the caller. /// starts at `frames + i * stride` with `stride = layout.lane_stride()`.
#[derive(Debug)] /// Frame bytes carry no drop glue (droppable elements ride parked,
pub struct FilledBatch<'a, T> { /// arena-owned), so the view has no drop obligation; `'a` covers the frames
values: &'a mut [T], /// and the layout.
} #[derive(Clone, Copy, Debug)]
pub struct RecordBatch<'a> {
impl<'a, T> FilledBatch<'a, T> { frames: *const u8,
/// # Safety stride: usize,
/// len: usize,
/// The first `len` elements of `scratch` must be initialized, and `len` must not exceed `scratch.len()`.
pub unsafe fn new(scratch: &'a mut [MaybeUninit<T>], len: usize) -> Self {
Self {
values: unsafe { assume_init_prefix_mut(scratch, len) },
}
}
pub fn values(&self) -> &[T] {
self.values
}
pub fn into_values(self) -> &'a mut [T] {
let mut guard = std::mem::ManuallyDrop::new(self);
std::mem::take(&mut guard.values)
}
}
impl<T> Drop for FilledBatch<'_, T> {
fn drop(&mut self) {
// SAFETY: every lane was initialized when the guard was built and none has
// been moved out, since `into_values` consumes the guard instead.
unsafe { std::ptr::drop_in_place(self.values as *mut [T]) }
}
}
/// # Safety
///
/// The first `len` elements of `scratch` must be initialized, and `len` must not exceed `scratch.len()`.
pub unsafe fn assume_init_prefix_mut<T>(scratch: &mut [MaybeUninit<T>], len: usize) -> &mut [T] {
debug_assert!(len <= scratch.len());
unsafe { std::slice::from_raw_parts_mut(scratch.as_mut_ptr().cast::<T>(), len) }
}
/// A borrow-for-scope view over a batch of records whose element type is `T`,
/// paired with their shared [`Layout`](crate::record::Layout). Row-major backed
/// today; the interface (`len`/`layout`/`get`/`for_each`) is storage-agnostic so
/// a columnar backing can replace it without touching consumers.
#[derive(Debug)]
pub struct RecordBatch<'a, T> {
lanes: LaneStore<'a, T>,
layout: &'a crate::record::Layout, layout: &'a crate::record::Layout,
_lifetime: PhantomData<&'a [u8]>,
} }
#[derive(Debug)] impl<'a> RecordBatch<'a> {
enum LaneStore<'a, T> { /// # Safety
/// Borrows resident storage (the `Lent` status): no drop obligation. /// `frames` must hold `len` initialized records of `layout`, packed at
Borrowed(&'a [T]), /// `layout.lane_stride()` stride and valid for `'a`.
/// Owns the caller scratch's initialized prefix (the `Filled` status). pub unsafe fn new(frames: *const u8, len: usize, layout: &'a crate::record::Layout) -> Self {
Owned(FilledBatch<'a, T>), Self {
} frames,
stride: layout.lane_stride(),
impl<'a, T> RecordBatch<'a, T> { len,
pub fn lent(values: &'a [T], layout: &'a crate::record::Layout) -> Self { layout,
Self { lanes: LaneStore::Borrowed(values), layout } _lifetime: PhantomData,
}
pub fn filled(filled: FilledBatch<'a, T>, layout: &'a crate::record::Layout) -> Self {
Self { lanes: LaneStore::Owned(filled), layout }
}
fn lanes(&self) -> &[T] {
match &self.lanes {
LaneStore::Borrowed(values) => values,
LaneStore::Owned(filled) => filled.values(),
} }
} }
pub fn len(&self) -> usize { pub fn len(&self) -> usize {
self.lanes().len() self.len
} }
pub fn is_empty(&self) -> bool { pub fn is_empty(&self) -> bool {
self.len() == 0 self.len == 0
} }
pub fn layout(&self) -> &crate::record::Layout { pub fn layout(&self) -> &'a crate::record::Layout {
self.layout self.layout
} }
/// Lends lane `lane`'s record to `f` for the callback's scope only. pub fn get(&self, lane: usize) -> RecordLane<'a> {
pub fn get<R>(&self, lane: usize, f: impl FnOnce(RecordLane<'_, T>) -> R) -> R { assert!(lane < self.len, "lane {lane} out of bounds for a batch of {}", self.len);
f(RecordLane { value: &self.lanes()[lane], layout: self.layout }) RecordLane {
} // SAFETY: in-bounds by the assert against the constructor's contract.
rec: unsafe { crate::record::Rec::new(self.frames.add(lane * self.stride)) },
/// Lends every lane's record in order, each for its callback's scope only. layout: self.layout,
pub fn for_each(&self, mut f: impl FnMut(usize, RecordLane<'_, T>)) {
for (lane, value) in self.lanes().iter().enumerate() {
f(lane, RecordLane { value, layout: self.layout });
} }
} }
/// Hands the owned scratch prefix back to the caller, cancelling the drop pub fn for_each(&self, mut f: impl FnMut(usize, RecordLane<'a>)) {
/// obligation. Panics on a lent batch, which owns nothing to return. for lane in 0..self.len {
pub fn into_values(self) -> &'a mut [T] { f(lane, self.get(lane));
match self.lanes {
LaneStore::Owned(filled) => filled.into_values(),
LaneStore::Borrowed(_) => panic!("into_values on a lent batch"),
} }
} }
} }
/// One lane's record, lent for a callback scope. Derefs to the raw lane value; /// The exclusive view over caller-owned frames (the `Filled` status): while it
/// for record elements, [`rec`](RecordLane::rec) and [`attr`](RecordLane::attr) /// lives, the borrow of the caller's scratch guarantees nobody else can read
/// read the record through its layout. /// the lanes, so mutating them or reclaiming the buffer is sound.
#[derive(Debug)] #[derive(Debug)]
pub struct RecordLane<'r, T> { pub struct RecordBatchMut<'a> {
value: &'r T, scratch: &'a mut [MaybeUninit<u64>],
layout: &'r crate::record::Layout, len: usize,
layout: &'a crate::record::Layout,
} }
impl<T> std::ops::Deref for RecordLane<'_, T> { impl<'a> RecordBatchMut<'a> {
type Target = T; /// # Safety
/// `scratch` must start with `len` initialized records of `layout`, packed
fn deref(&self) -> &T { /// at `layout.lane_stride()` stride.
self.value pub unsafe fn new(scratch: &'a mut [MaybeUninit<u64>], len: usize, layout: &'a crate::record::Layout) -> Self {
debug_assert!(len * layout.lane_stride() <= scratch.len() * 8);
Self { scratch, len, layout }
} }
}
impl<T> RecordLane<'_, T> { pub fn len(&self) -> usize {
pub fn layout(&self) -> &crate::record::Layout { self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn layout(&self) -> &'a crate::record::Layout {
self.layout self.layout
} }
/// Reads the lanes without giving up exclusivity.
pub fn share(&self) -> RecordBatch<'_> {
// SAFETY: the constructor's contract, narrowed to the reborrow's scope.
unsafe { RecordBatch::new(self.scratch.as_ptr().cast(), self.len, self.layout) }
}
/// Gives up exclusivity for the batch's whole lifetime.
pub fn into_shared(self) -> RecordBatch<'a> {
// SAFETY: the constructor's contract; the exclusive borrow is consumed.
unsafe { RecordBatch::new(self.scratch.as_ptr().cast(), self.len, self.layout) }
}
/// Lane `lane`'s frame for in-place writes through the layout's offsets.
pub fn lane_ptr(&mut self, lane: usize) -> *mut u8 {
assert!(lane < self.len, "lane {lane} out of bounds for a batch of {}", self.len);
// SAFETY: in-bounds by the assert against the constructor's contract.
unsafe { self.scratch.as_mut_ptr().cast::<u8>().add(lane * self.layout.lane_stride()) }
}
/// Reclaims the raw buffer, e.g. to rebind it under a same-stride output
/// layout for an in-place map.
pub fn into_scratch(self) -> &'a mut [MaybeUninit<u64>] {
self.scratch
}
} }
impl<'e> RecordLane<'_, crate::record::RecordValue<'e>> { /// One lane's record: its pointer paired with the batch's layout.
/// The record pointer, resolved through the layout. #[derive(Clone, Copy, Debug)]
pub struct RecordLane<'a> {
rec: crate::record::Rec,
layout: &'a crate::record::Layout,
}
impl<'a> RecordLane<'a> {
pub fn layout(&self) -> &'a crate::record::Layout {
self.layout
}
pub fn rec(&self) -> crate::record::Rec { pub fn rec(&self) -> crate::record::Rec {
self.layout.rec(self.value) self.rec
} }
/// The element at offset 0. /// The element at offset 0.
@@ -160,32 +155,32 @@ impl<'e> RecordLane<'_, crate::record::RecordValue<'e>> {
/// # Safety /// # Safety
/// `U` must be the record's element type, proven at the consumer's wiring. /// `U` must be the record's element type, proven at the consumer's wiring.
pub unsafe fn element<U: Copy>(&self) -> U { pub unsafe fn element<U: Copy>(&self) -> U {
unsafe { self.rec().element::<U>() } unsafe { self.rec.element::<U>() }
} }
/// Attribute `A` at the record's top level, or its census default when the /// Attribute `A` at the record's top level, or its census default when the
/// layout does not carry it. /// layout does not carry it.
pub fn attr<A: crate::attribute::Attribute>(&self) -> A::Value<'e> { pub fn attr<A: crate::attribute::Attribute>(&self) -> A::Value<'a> {
match self.layout.offset_of(A::NAME, 0) { match self.layout.offset_of(A::NAME, 0) {
Some(offset) => unsafe { self.rec().read::<A::Value<'e>>(offset) }, Some(offset) => unsafe { self.rec.read::<A::Value<'a>>(offset) },
None => A::default(), None => A::default(),
} }
} }
} }
/// A materialized nesting level handed to a folding kernel: a thin element-typed /// A materialized nesting level handed to a folding kernel: a thin element-typed
/// view over the [`RecordBatch`] the level was collected into. `'a` is the batch /// view over the [`RecordBatch`] the level was collected into. The eventual
/// view, `'e` the record payloads. The eventual `List` once `IList` is renamed. /// `List` once `IList` is renamed.
#[derive(Debug)] #[derive(Debug)]
pub struct List<'a, 'e, T> { pub struct List<'a, T> {
batch: RecordBatch<'a, crate::record::RecordValue<'e>>, batch: RecordBatch<'a>,
_element: PhantomData<T>, _element: PhantomData<T>,
} }
impl<'a, 'e, T: Copy> List<'a, 'e, T> { impl<'a, T: Copy> List<'a, T> {
/// # Safety /// # Safety
/// `T` must be the batch's record element type, proven at the consumer's wiring. /// `T` must be the batch's record element type, proven at the consumer's wiring.
pub unsafe fn new(batch: RecordBatch<'a, crate::record::RecordValue<'e>>) -> Self { pub unsafe fn new(batch: RecordBatch<'a>) -> Self {
Self { batch, _element: PhantomData } Self { batch, _element: PhantomData }
} }
@@ -199,7 +194,7 @@ impl<'a, 'e, T: Copy> List<'a, 'e, T> {
pub fn get(&self, index: usize) -> T { pub fn get(&self, index: usize) -> T {
// SAFETY: `List::new` established that `T` is the batch's element type. // SAFETY: `List::new` established that `T` is the batch's element type.
self.batch.get(index, |lane| unsafe { lane.element::<T>() }) unsafe { self.batch.get(index).element::<T>() }
} }
pub fn iter(&self) -> impl Iterator<Item = T> + '_ { pub fn iter(&self) -> impl Iterator<Item = T> + '_ {
@@ -207,21 +202,21 @@ impl<'a, 'e, T: Copy> List<'a, 'e, T> {
} }
} }
impl<'a, 'e, T: Copy> IntoIterator for List<'a, 'e, T> { impl<'a, T: Copy> IntoIterator for List<'a, T> {
type Item = T; type Item = T;
type IntoIter = ListIter<'a, 'e, T>; type IntoIter = ListIter<'a, T>;
fn into_iter(self) -> ListIter<'a, 'e, T> { fn into_iter(self) -> ListIter<'a, T> {
ListIter { list: self, position: 0 } ListIter { list: self, position: 0 }
} }
} }
pub struct ListIter<'a, 'e, T> { pub struct ListIter<'a, T> {
list: List<'a, 'e, T>, list: List<'a, T>,
position: usize, position: usize,
} }
impl<T: Copy> Iterator for ListIter<'_, '_, T> { impl<T: Copy> Iterator for ListIter<'_, T> {
type Item = T; type Item = T;
fn next(&mut self) -> Option<T> { fn next(&mut self) -> Option<T> {
@@ -281,47 +276,18 @@ pub trait Node<Input> {
/// Installs this node's resolved record layout; a no-op unless it produces records. /// Installs this node's resolved record layout; a no-op unless it produces records.
fn set_layout(&mut self, _layout: crate::record::RecordLayout) {} fn set_layout(&mut self, _layout: crate::record::RecordLayout) {}
fn eval_batch<'a>(&'a self, input: &'a Input, range: Range<u64>, scratch: Option<&'a mut [MaybeUninit<Self::Output>]>) -> BatchStatus<'a, Self::Output> /// Batched evaluation of `range` into caller-provided frame storage of
/// `range.len() * layout.lane_stride()` bytes; see [`BatchStatus`]. The
/// default advertises no support and drivers fall back to per-lane eval
/// with copy-out ([`crate::record::fill_frames`]); overrides exist to beat
/// that loop (resident lanes, direct fills, fewer erased calls), never for
/// correctness.
fn eval_batch<'a>(&'a self, input: &'a Input, range: Range<u64>, scratch: Option<&'a mut [MaybeUninit<u64>]>) -> BatchStatus<'a>
where where
Input: InjectIndex + Copy, Input: InjectIndex + Copy,
{ {
let Some(scratch) = scratch else { let _ = (input, range, scratch);
return BatchStatus::NeedBuffer; BatchStatus::Unbatched
};
let Some(len) = range.end.checked_sub(range.start).and_then(|len| usize::try_from(len).ok()) else {
return BatchStatus::InvalidRange;
};
if scratch.len() < len {
return BatchStatus::InvalidRange;
}
let mut local = *input;
let mut finality = Finality::AllFinal;
for offset in 0..len {
local.set_index(range.start + offset as u64);
let abort = match self.eval(&local) {
GPoll::Final(value) => {
scratch[offset].write(value);
None
}
GPoll::Partial(value) => {
scratch[offset].write(value);
finality = Finality::Partial;
None
}
GPoll::Pending => Some(BatchStatus::Pending),
GPoll::Fallback(boxed) => Some(BatchStatus::Error(boxed.1)),
GPoll::Error(e) => Some(BatchStatus::Error(*e)),
};
if let Some(status) = abort {
for written in scratch[..offset].iter_mut() {
// SAFETY: every lane before `offset` was written by this loop.
unsafe { written.assume_init_drop() };
}
return status;
}
}
// SAFETY: all `len` lanes were written by the loop above.
BatchStatus::Filled(RecordBatch::filled(unsafe { FilledBatch::new(scratch, len) }, self.layout()), finality)
} }
} }
@@ -347,7 +313,7 @@ where
(**self).layout() (**self).layout()
} }
fn eval_batch<'a>(&'a self, input: &'a Input, range: Range<u64>, scratch: Option<&'a mut [MaybeUninit<Self::Output>]>) -> BatchStatus<'a, Self::Output> fn eval_batch<'a>(&'a self, input: &'a Input, range: Range<u64>, scratch: Option<&'a mut [MaybeUninit<u64>]>) -> BatchStatus<'a>
where where
Input: InjectIndex + Copy, Input: InjectIndex + Copy,
{ {
@@ -377,7 +343,7 @@ where
(**self).layout() (**self).layout()
} }
fn eval_batch<'a>(&'a self, input: &'a Input, range: Range<u64>, scratch: Option<&'a mut [MaybeUninit<Self::Output>]>) -> BatchStatus<'a, Self::Output> fn eval_batch<'a>(&'a self, input: &'a Input, range: Range<u64>, scratch: Option<&'a mut [MaybeUninit<u64>]>) -> BatchStatus<'a>
where where
Input: InjectIndex + Copy, Input: InjectIndex + Copy,
{ {
@@ -407,7 +373,7 @@ where
(**self).layout() (**self).layout()
} }
fn eval_batch<'a>(&'a self, input: &'a Input, range: Range<u64>, scratch: Option<&'a mut [MaybeUninit<Self::Output>]>) -> BatchStatus<'a, Self::Output> fn eval_batch<'a>(&'a self, input: &'a Input, range: Range<u64>, scratch: Option<&'a mut [MaybeUninit<u64>]>) -> BatchStatus<'a>
where where
Input: InjectIndex + Copy, Input: InjectIndex + Copy,
{ {
@@ -511,7 +477,6 @@ impl<'a, N> LazyInput<'a, N> {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use std::sync::atomic::{AtomicU32, Ordering};
#[derive(Clone, Copy)] #[derive(Clone, Copy)]
struct TestInput { struct TestInput {
@@ -535,107 +500,11 @@ mod tests {
} }
#[test] #[test]
fn spec_loop_fills_scratch_per_lane() { fn the_default_advertises_no_batch_support() {
let input = TestInput { index: 0 }; let input = TestInput { index: 0 };
let mut scratch = [const { MaybeUninit::uninit() }; 4]; let mut scratch = [const { MaybeUninit::uninit() }; 4];
let status = Double.eval_batch(&input, 2..6, Some(&mut scratch)); assert!(matches!(Double.eval_batch(&input, 2..6, Some(&mut scratch)), BatchStatus::Unbatched));
let BatchStatus::Filled(batch, finality) = status else { assert!(matches!(Double.eval_batch(&input, 2..6, None), BatchStatus::Unbatched));
panic!("expected filled, got {status:?}");
};
let mut got = Vec::new();
batch.for_each(|_, lane| got.push(*lane));
assert_eq!(got, vec![4, 6, 8, 10]);
assert_eq!(finality, Finality::AllFinal);
}
#[test]
fn a_dropped_filled_batch_reclaims_every_lane() {
static DROPS: AtomicU32 = AtomicU32::new(0);
#[derive(Clone)]
struct Probe;
impl Drop for Probe {
fn drop(&mut self) {
DROPS.fetch_add(1, Ordering::Relaxed);
}
}
struct Probes;
impl Node<TestInput> for Probes {
type Output = Probe;
fn eval(&self, _input: &TestInput) -> GPoll<Probe> {
GPoll::Final(Probe)
}
}
let input = TestInput { index: 0 };
let mut scratch = [const { MaybeUninit::uninit() }; 3];
let status = Probes.eval_batch(&input, 0..3, Some(&mut scratch));
assert!(matches!(status, BatchStatus::Filled(..)));
drop(status);
assert_eq!(DROPS.load(Ordering::Relaxed), 3, "an unconsumed batch must not leak its lanes");
}
#[test]
fn probe_without_scratch_requests_a_buffer() {
let input = TestInput { index: 0 };
assert!(matches!(Double.eval_batch(&input, 0..4, None), BatchStatus::NeedBuffer));
}
#[test]
fn undersized_scratch_is_an_invalid_range() {
let input = TestInput { index: 0 };
let mut scratch = [const { MaybeUninit::uninit() }; 2];
assert!(matches!(Double.eval_batch(&input, 0..4, Some(&mut scratch)), BatchStatus::InvalidRange));
}
#[test]
fn partial_lane_downgrades_batch_finality() {
struct PartialAtThree;
impl Node<TestInput> for PartialAtThree {
type Output = u64;
fn eval(&self, input: &TestInput) -> GPoll<u64> {
match input.index {
3 => GPoll::Partial(input.index),
index => GPoll::Final(index),
}
}
}
let input = TestInput { index: 0 };
let mut scratch = [const { MaybeUninit::uninit() }; 4];
let status = PartialAtThree.eval_batch(&input, 0..4, Some(&mut scratch));
let BatchStatus::Filled(batch, finality) = status else {
panic!("expected filled, got {status:?}");
};
let mut got = Vec::new();
batch.for_each(|_, lane| got.push(*lane));
assert_eq!(got, vec![0, 1, 2, 3]);
assert_eq!(finality, Finality::Partial);
}
#[test]
fn abort_drops_already_written_lanes() {
static DROPS: AtomicU32 = AtomicU32::new(0);
struct Probe;
impl Drop for Probe {
fn drop(&mut self) {
DROPS.fetch_add(1, Ordering::Relaxed);
}
}
struct PendingAtTwo;
impl Node<TestInput> for PendingAtTwo {
type Output = Probe;
fn eval(&self, input: &TestInput) -> GPoll<Probe> {
match input.index {
2 => GPoll::Pending,
_ => GPoll::Final(Probe),
}
}
}
let input = TestInput { index: 0 };
let mut scratch = [const { MaybeUninit::uninit() }; 4];
let status = PendingAtTwo.eval_batch(&input, 0..4, Some(&mut scratch));
assert!(matches!(status, BatchStatus::Pending));
assert_eq!(DROPS.load(Ordering::Relaxed), 2);
} }
#[test] #[test]
@@ -643,8 +512,6 @@ mod tests {
let erased: Box<dyn Node<TestInput, Output = u64>> = Box::new(Double); let erased: Box<dyn Node<TestInput, Output = u64>> = Box::new(Double);
let input = TestInput { index: 21 }; let input = TestInput { index: 21 };
assert_eq!(erased.eval(&input), GPoll::Final(42)); assert_eq!(erased.eval(&input), GPoll::Final(42));
let mut scratch = [const { MaybeUninit::uninit() }; 2]; assert!(matches!(erased.eval_batch(&input, 0..2, None), BatchStatus::Unbatched));
let status = erased.eval_batch(&input, 0..2, Some(&mut scratch));
assert!(matches!(status, BatchStatus::Filled(_, Finality::AllFinal)));
} }
} }
@@ -120,6 +120,16 @@ impl Layout {
self.size.next_multiple_of(8) self.size.next_multiple_of(8)
} }
/// One batch lane's stride: a spilled record's frame, or the value itself
/// for records this layout keeps inline (`size == 0`), whose payload rides
/// in the `RecordValue`'s own storage exactly as [`Layout::rec`] resolves.
pub fn lane_stride(&self) -> usize {
match self.size == 0 {
true => size_of::<RecordValue<'static>>(),
false => self.frame_bytes(),
}
}
/// Resolves a value of this layout, which must be its wiring-proven one, /// Resolves a value of this layout, which must be its wiring-proven one,
/// to its record bytes. An empty record carries nothing and resolves to the /// to its record bytes. An empty record carries nothing and resolves to the
/// value's own storage; every other record spills and rides the pointer. /// value's own storage; every other record spills and rides the pointer.
@@ -454,6 +464,86 @@ where
} }
} }
/// Fills caller scratch with one frame per lane of `range`: the edge
/// evaluates at each index, the record's frame copies out, and the stack
/// rewinds, so the stack peak stays at one lane's need and every lane's bytes
/// are distinct. Frame bytes carry no drop glue, so the copy is a move.
pub fn fill_frames<'a, 'e, C, N>(node: &'a N, input: &C, range: std::ops::Range<u64>, scratch: Option<&'a mut [std::mem::MaybeUninit<u64>]>) -> crate::node::BatchStatus<'a>
where
C: crate::context::InjectIndex + Copy,
N: Node<C, Output = RecordValue<'e>>,
{
use crate::node::BatchStatus;
let Some(scratch) = scratch else {
return BatchStatus::NeedBuffer;
};
let Some(len) = range.end.checked_sub(range.start).and_then(|len| usize::try_from(len).ok()) else {
return BatchStatus::InvalidRange;
};
let layout = node.layout();
let stride = layout.lane_stride();
if scratch.len() * 8 < len * stride {
return BatchStatus::InvalidRange;
}
let base = scratch.as_mut_ptr().cast::<u8>();
let mut local = *input;
let mut finality = crate::gpoll::Finality::AllFinal;
for lane in 0..len {
local.set_index(range.start + lane as u64);
let mark = stack::sp();
let value = match node.eval(&local) {
GPoll::Final(value) => value,
GPoll::Partial(value) => {
finality = crate::gpoll::Finality::Partial;
value
}
GPoll::Pending => return BatchStatus::Pending,
GPoll::Fallback(boxed) => return BatchStatus::Error(boxed.1),
GPoll::Error(error) => return BatchStatus::Error(*error),
};
// SAFETY: the lane region is in-bounds by the scratch check, and the
// frame is fully copied out before the rewind releases it.
unsafe {
std::ptr::copy_nonoverlapping(layout.rec(&value).ptr(), base.add(lane * stride), stride);
stack::rewind(mark);
}
}
// SAFETY: all `len` lanes were filled above with records of `layout`.
BatchStatus::Filled(unsafe { crate::node::RecordBatchMut::new(scratch, len, layout) }, finality)
}
/// The driver a consumer runs on a record edge: a resident batch returns with
/// no allocation, a node's own batch impl gets `n * frame_bytes` of arena
/// scratch, and an unbatched edge falls back to the [`fill_frames`] loop.
pub fn materialize_batch<'a, 'e, C, N>(node: &'a N, input: &'a C, range: std::ops::Range<u64>, arena: &'a crate::arena::Arena) -> crate::node::BatchStatus<'a>
where
C: crate::context::InjectIndex + Copy,
N: Node<C, Output = RecordValue<'e>>,
{
use crate::node::BatchStatus;
let Some(len) = range.end.checked_sub(range.start).and_then(|len| usize::try_from(len).ok()) else {
return BatchStatus::InvalidRange;
};
let words = len * node.layout().lane_stride() / 8;
let exhausted = || {
BatchStatus::Error(crate::gpoll::GraphError {
kind: crate::gpoll::ErrorKind::ArenaExhausted,
trace: Vec::new(),
})
};
match node.eval_batch(input, range.clone(), None) {
BatchStatus::Unbatched => match arena.alloc_scratch::<u64>(words) {
Some(scratch) => fill_frames(node, input, range, Some(scratch)),
None => exhausted(),
},
BatchStatus::NeedBuffer => match arena.alloc_scratch::<u64>(words) {
Some(scratch) => node.eval_batch(input, range, Some(scratch)),
None => exhausted(),
},
status => status,
}
}
/// A record edge at a caller-chosen lifetime; the lifetime is a trait /// A record edge at a caller-chosen lifetime; the lifetime is a trait
/// parameter for the same constrained-position reason as /// parameter for the same constrained-position reason as
/// [`DerivedRecordEdge`]. /// [`DerivedRecordEdge`].
@@ -179,7 +179,7 @@ where
unsafe { self.ptr.as_ref() }.layout() unsafe { self.ptr.as_ref() }.layout()
} }
fn eval_batch<'a>(&'a self, input: &'a Input, range: std::ops::Range<u64>, scratch: Option<&'a mut [std::mem::MaybeUninit<Self::Output>]>) -> crate::node::BatchStatus<'a, Self::Output> fn eval_batch<'a>(&'a self, input: &'a Input, range: std::ops::Range<u64>, scratch: Option<&'a mut [std::mem::MaybeUninit<u64>]>) -> crate::node::BatchStatus<'a>
where where
Input: crate::context::InjectIndex + Copy, Input: crate::context::InjectIndex + Copy,
{ {
+27 -18
View File
@@ -934,7 +934,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
let pat = &field.pat_ident; let pat = &field.pat_ident;
match &field.ty { match &field.ty {
ParsedFieldType::Regular(RegularParsedField { ty, .. }) if ir::materialized_levels(&node, index) > 0 => { ParsedFieldType::Regular(RegularParsedField { ty, .. }) if ir::materialized_levels(&node, index) > 0 => {
quote!(#pat: #core_types::node::List<'_, '_, #ty>) quote!(#pat: #core_types::node::List<'_, #ty>)
} }
ParsedFieldType::Regular(RegularParsedField { ty, lend: Some(_), .. }) => quote!(#pat: &#ty), ParsedFieldType::Regular(RegularParsedField { ty, lend: Some(_), .. }) => quote!(#pat: &#ty),
ParsedFieldType::Regular(RegularParsedField { ty, .. }) if !field.attribute_reads.is_empty() => read_tuple_param(field, quote!(#pat), quote!(#ty)), ParsedFieldType::Regular(RegularParsedField { ty, .. }) if !field.attribute_reads.is_empty() => read_tuple_param(field, quote!(#pat), quote!(#ty)),
@@ -1082,12 +1082,9 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
#core_types::gpoll::GPoll::Pending => return #core_types::gpoll::GPoll::Pending, #core_types::gpoll::GPoll::Pending => return #core_types::gpoll::GPoll::Pending,
_ => return #core_types::gpoll::GPoll::Error(::std::boxed::Box::new(#core_types::gpoll::GraphError::new("reduce over a non-exact extent"))), _ => return #core_types::gpoll::GPoll::Error(::std::boxed::Box::new(#core_types::gpoll::GraphError::new("reduce over a non-exact extent"))),
}; };
let __scratch = match __arena.alloc_scratch::<#core_types::record::RecordValue<'__record>>(__count) { let __batch = match #core_types::record::materialize_batch(&self.#name, __input, 0..__count as u64, __arena) {
Some(__scratch) => __scratch, #core_types::node::BatchStatus::Lent(__batch, _) => __batch,
None => return #core_types::gpoll::GPoll::Error(::std::boxed::Box::new(#core_types::gpoll::GraphError::new("reduce scratch allocation failed"))), #core_types::node::BatchStatus::Filled(__batch, _) => __batch.into_shared(),
};
let __batch = match #core_types::node::Node::eval_batch(&self.#name, __input, 0..__count as u64, Some(__scratch)) {
#core_types::node::BatchStatus::Lent(__batch, _) | #core_types::node::BatchStatus::Filled(__batch, _) => __batch,
#core_types::node::BatchStatus::Pending => return #core_types::gpoll::GPoll::Pending, #core_types::node::BatchStatus::Pending => return #core_types::gpoll::GPoll::Pending,
#core_types::node::BatchStatus::Error(__error) => return #core_types::gpoll::GPoll::Error(::std::boxed::Box::new(__error)), #core_types::node::BatchStatus::Error(__error) => return #core_types::gpoll::GPoll::Error(::std::boxed::Box::new(__error)),
_ => return #core_types::gpoll::GPoll::Error(::std::boxed::Box::new(#core_types::gpoll::GraphError::new("reduce batch failed"))), _ => return #core_types::gpoll::GPoll::Error(::std::boxed::Box::new(#core_types::gpoll::GraphError::new("reduce batch failed"))),
@@ -1245,21 +1242,33 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
None => quote!(), None => quote!(),
}; };
let batch_impl = match &parsed.attributes.batch { let batch_signature = quote! {
Some(path) => quote! { fn eval_batch<'__batch>(
fn eval_batch<'__batch>( &'__batch self,
&'__batch self, __input: &'__batch #ctx_ident,
__input: &'__batch #ctx_ident, __range: ::std::ops::Range<u64>,
__range: ::std::ops::Range<u64>, __scratch: Option<&'__batch mut [::std::mem::MaybeUninit<u64>]>,
__scratch: Option<&'__batch mut [::std::mem::MaybeUninit<Self::Output>]>, ) -> #core_types::node::BatchStatus<'__batch>
) -> #core_types::node::BatchStatus<'__batch, Self::Output> where
where #ctx_ident: #core_types::context::InjectIndex + Copy,
#ctx_ident: #core_types::context::InjectIndex + Copy, };
let produces_records = record_io || routing_generic.is_some() || flip;
let batch_impl = match (&parsed.attributes.batch, produces_records) {
(Some(path), _) => quote! {
#batch_signature
{ {
#path(self, __input, __range, __scratch) #path(self, __input, __range, __scratch)
} }
}, },
None => quote!(), // The eager forward runs the shared copy-out loop with statically
// dispatched evals, so an erased batch costs one virtual call.
(None, true) => quote! {
#batch_signature
{
#core_types::record::fill_frames(self, __input, __range, __scratch)
}
},
(None, false) => quote!(),
}; };
let ctx_pat = &parsed.input.pat_ident; let ctx_pat = &parsed.input.pat_ident;
+36
View File
@@ -457,6 +457,42 @@ mod tests {
assert_eq!(unsafe { out.rec(&value).element::<f64>() }, 21.); assert_eq!(unsafe { out.rec(&value).element::<f64>() }, 21.);
} }
#[test]
fn reducer_folds_varying_copies_of_a_generic_repeat() {
let arena = Arena::new(1024).unwrap();
let generations = [];
let scope = scope_fixture(&generations, &arena);
let ctx = ContextImpl::root(&scope);
let base = f64_layout(&[]);
let (count_edge, count_layout) = lifted_value(4u32);
let out = f64_layout(&[]);
reserve_for(&[&base, &count_layout, &out]);
let meta = core_types::record::LayoutMeta {
sources: vec![0],
reads: vec![],
element: core_types::record::ElementSpec::Carried,
writes: vec![],
removes: vec![],
level_delta: 1,
};
let repeat = install(
RepeatNode::new(RecordSource::new(IndexSourceNode { layout: base.clone() }, &base, &base), count_edge, &base, &count_layout),
meta,
&[Some(&base)],
);
let leveled = Node::<ContextImpl>::layout(&repeat).clone();
let node = install_flip(SumNode::new(repeat, &leveled), &out);
let GPoll::Final(value) = node.eval(&ctx) else {
panic!("expected a final record");
};
// Every copy evaluates at its own index, so the lanes must be distinct
// storage: sum(0 + 1 + 2 + 3), not four aliases of the last copy.
assert_eq!(unsafe { out.rec(&value).element::<f64>() }, 6.);
}
#[test] #[test]
fn layout_meta_folds_to_construction() { fn layout_meta_folds_to_construction() {
let base = f64_layout(&[]); let base = f64_layout(&[]);
+2 -1
View File
@@ -1140,13 +1140,14 @@ mod graphene_test {
reserve_for(&[&li, &ls, &out]); reserve_for(&[&li, &ls, &out]);
let erased: Box<ErasedRecordNode> = Box::new(node); let erased: Box<ErasedRecordNode> = Box::new(node);
// One u64 word per lane: the uninstalled layout keeps the f64 inline.
let mut scratch = [const { MaybeUninit::uninit() }; 4]; let mut scratch = [const { MaybeUninit::uninit() }; 4];
let status = erased.eval_batch(&ctx, 2..6, Some(&mut scratch)); let status = erased.eval_batch(&ctx, 2..6, Some(&mut scratch));
let BatchStatus::Filled(batch, finality) = status else { let BatchStatus::Filled(batch, finality) = status else {
panic!("expected filled, got {status:?}"); panic!("expected filled, got {status:?}");
}; };
let mut got = Vec::new(); let mut got = Vec::new();
batch.for_each(|_, lane| got.push(unsafe { lane.element::<f64>() })); batch.share().for_each(|_, lane| got.push(unsafe { lane.element::<f64>() }));
assert_eq!(got, vec![12.0, 13.0, 14.0, 15.0]); assert_eq!(got, vec![12.0, 13.0, 14.0, 15.0]);
assert_eq!(finality, Finality::AllFinal); assert_eq!(finality, Finality::AllFinal);
} }