strided_basic/copy_plan.rs
1//! Prepared (compile-once, execute-many) copy plans over raw strided layouts.
2//!
3//! [`copy_scale_raw`](crate::copy_scale_raw) and friends rebuild the fused
4//! loop nest on every call; for prepared-replay consumers that issue many
5//! small copies with a fixed layout, that per-call planning dominates
6//! (see issue #139). [`CopyPlan`] splits the work: [`CopyPlan::compile`]
7//! validates the layout pair and builds the fused traversal once,
8//! [`CopyPlan::execute`]/[`CopyPlan::execute_scale`]/[`CopyPlan::execute_conj`]
9//! replay it with no planning and no heap allocation for ranks at most
10//! [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
11
12use core::{
13 marker::PhantomData,
14 mem::MaybeUninit,
15 ops::{Add, Mul},
16};
17
18use crate::map_view::map_raw_into;
19use crate::ops_view::{copy_conj, copy_into, copy_scale};
20use crate::raw_ops::{apply_fused_range, fuse_pair_layout, fused_total, FusedPairLayout};
21use crate::{
22 ElementOpApply, Identity, MaybeSendSync, RawStridedMut, RawStridedRef, Result, StridedError,
23};
24
25// Same pattern as map_view.rs / outer_product.rs: stack storage when the
26// parallel feature pulls in smallvec, plain Vec otherwise. Only `compile`
27// touches these; `execute*` never allocates either way.
28#[cfg(feature = "parallel")]
29type AxisVec<T> = smallvec::SmallVec<[T; crate::RAW_FUSED_RANK_LIMIT]>;
30#[cfg(not(feature = "parallel"))]
31type AxisVec<T> = Vec<T>;
32
33pub(crate) trait OverwriteWriter<T> {
34 fn dims(&self) -> &[usize];
35 fn strides(&self) -> &[isize];
36 fn offset(&self) -> isize;
37 /// # Safety
38 /// The caller must use the pointer only for the validated allocation and
39 /// logical layout represented by this writer.
40 unsafe fn data_ptr(&mut self) -> *mut T;
41 /// # Safety
42 /// The offset must be an in-bounds logical destination proven by layout.
43 unsafe fn write_at(&mut self, offset: isize, value: T);
44}
45
46pub(crate) trait ReadModifyWrite<T>: OverwriteWriter<T> {
47 /// # Safety
48 /// The offset must be an in-bounds initialized slot covered by the
49 /// traversal's copy and disjointness proof.
50 unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
51 where
52 F: FnOnce(T, T) -> T;
53}
54
55impl<'a, T> OverwriteWriter<T> for RawStridedMut<'a, T> {
56 fn dims(&self) -> &[usize] {
57 self.dims()
58 }
59 fn strides(&self) -> &[isize] {
60 self.strides()
61 }
62 fn offset(&self) -> isize {
63 self.offset()
64 }
65 unsafe fn data_ptr(&mut self) -> *mut T {
66 self.data_mut().as_mut_ptr()
67 }
68 unsafe fn write_at(&mut self, offset: isize, value: T) {
69 // SAFETY: the prepared layout validates every logical destination.
70 unsafe { self.data_mut().as_mut_ptr().offset(offset).write(value) }
71 }
72}
73
74impl<'a, T> ReadModifyWrite<T> for RawStridedMut<'a, T>
75where
76 T: Add<Output = T>,
77{
78 #[inline(always)]
79 unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
80 where
81 F: FnOnce(T, T) -> T,
82 {
83 // SAFETY: the copy or initialized caller proves this logical slot.
84 unsafe {
85 let ptr = self.data_mut().as_mut_ptr().offset(offset);
86 ptr.write(combine(ptr.read(), value));
87 }
88 }
89}
90
91impl<'a, T> OverwriteWriter<T> for RawStridedMut<'a, MaybeUninit<T>> {
92 fn dims(&self) -> &[usize] {
93 self.dims()
94 }
95 fn strides(&self) -> &[isize] {
96 self.strides()
97 }
98 fn offset(&self) -> isize {
99 self.offset()
100 }
101 unsafe fn data_ptr(&mut self) -> *mut T {
102 self.data_mut().as_mut_ptr().cast()
103 }
104 unsafe fn write_at(&mut self, offset: isize, value: T) {
105 // SAFETY: the prepared layout validates every logical destination.
106 unsafe {
107 self.data_mut()
108 .as_mut_ptr()
109 .offset(offset)
110 .write(MaybeUninit::new(value))
111 }
112 }
113}
114
115pub(crate) struct InitializedRawDest<'a, T> {
116 ptr: *mut T,
117 extent: usize,
118 dims: &'a [usize],
119 strides: &'a [isize],
120 offset: isize,
121 _marker: PhantomData<&'a mut [MaybeUninit<T>]>,
122}
123
124impl<'a, T> OverwriteWriter<T> for InitializedRawDest<'a, T> {
125 fn dims(&self) -> &[usize] {
126 self.dims
127 }
128 fn strides(&self) -> &[isize] {
129 self.strides
130 }
131 fn offset(&self) -> isize {
132 self.offset
133 }
134 unsafe fn data_ptr(&mut self) -> *mut T {
135 self.ptr
136 }
137 unsafe fn write_at(&mut self, offset: isize, value: T) {
138 debug_assert!(offset >= 0 && (offset as usize) < self.extent);
139 // SAFETY: the copy proof and extent check cover this logical slot.
140 unsafe { self.ptr.offset(offset).write(value) }
141 }
142}
143
144impl<'a, T> ReadModifyWrite<T> for InitializedRawDest<'a, T>
145where
146 T: Add<Output = T>,
147{
148 #[inline(always)]
149 unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
150 where
151 F: FnOnce(T, T) -> T,
152 {
153 debug_assert!(offset >= 0 && (offset as usize) < self.extent);
154 // SAFETY: the copy proof and extent check cover this logical slot.
155 unsafe {
156 let ptr = self.ptr.offset(offset);
157 ptr.write(combine(ptr.read(), value));
158 }
159 }
160}
161
162/// A compiled copy traversal for one `(dims, dst_strides, src_strides)`
163/// layout pair.
164///
165/// `compile` proves the layout facts once (rank agreement, extent overflow,
166/// destination injectivity) and fuses/orders the loop nest; each `execute*`
167/// call then only re-checks the per-call facts (that the supplied views carry
168/// exactly the compiled layout) before replaying the prepared loops.
169///
170/// Overlapping `src`/`dest` memory is not supported, matching the rest of the
171/// crate.
172///
173/// Ranks above [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT) are supported through the view-based
174/// kernels; only that fallback path may allocate.
175///
176/// # Example
177///
178/// ```rust
179/// use strided_basic::{CopyPlan, RawStridedMut, RawStridedRef};
180///
181/// let dims = [2usize, 3];
182/// let src_strides = [3isize, 1];
183/// let dst_strides = [1isize, 2]; // transposed destination
184/// let plan = CopyPlan::compile(&dims, &dst_strides, &src_strides).unwrap();
185///
186/// let src = [0.0f64, 1.0, 2.0, 10.0, 11.0, 12.0];
187/// let mut dst = [0.0f64; 6];
188/// let src_ref = RawStridedRef::new(&src, &dims, &src_strides, 0).unwrap();
189/// let mut dst_mut = RawStridedMut::new(&mut dst, &dims, &dst_strides, 0).unwrap();
190/// plan.execute(&mut dst_mut, &src_ref).unwrap();
191/// assert_eq!(dst, [0.0, 10.0, 1.0, 11.0, 2.0, 12.0]);
192/// ```
193#[derive(Clone, Debug)]
194pub struct CopyPlan {
195 dims: AxisVec<usize>,
196 dst_strides: AxisVec<isize>,
197 src_strides: AxisVec<isize>,
198 /// `None` when rank exceeds [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT); `execute*` then
199 /// falls back to the view-based kernels.
200 fused: Option<FusedPairLayout>,
201}
202
203impl CopyPlan {
204 /// The fused traversal, or `None` above
205 /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
206 #[cfg(feature = "parallel")]
207 #[inline]
208 pub(crate) fn fused_layout(&self) -> Option<&FusedPairLayout> {
209 self.fused.as_ref()
210 }
211
212 pub(crate) fn execute_uninit_then<'a, T, R>(
213 &self,
214 dest: &'a mut RawStridedMut<'a, MaybeUninit<T>>,
215 src: &RawStridedRef<'_, T>,
216 f: impl for<'b> FnOnce(InitializedRawDest<'b, T>) -> R,
217 ) -> Result<R>
218 where
219 T: Copy + MaybeSendSync,
220 {
221 self.execute_uninit(dest, src)?;
222 let data = dest.data_mut();
223 let receipt = InitializedRawDest {
224 ptr: data.as_mut_ptr().cast(),
225 extent: data.len(),
226 dims: dest.dims(),
227 strides: dest.strides(),
228 offset: dest.offset(),
229 _marker: PhantomData,
230 };
231 Ok(f(receipt))
232 }
233
234 /// Compile a copy plan for the given layout pair.
235 ///
236 /// Performs the layout validation and traversal construction
237 /// (fuse + order) once:
238 ///
239 /// - `dims`, `dst_strides`, and `src_strides` must have equal length
240 /// ([`StridedError::StrideLengthMismatch`]);
241 /// - the total element count must not overflow `usize`
242 /// ([`StridedError::OffsetOverflow`]);
243 /// - the destination layout must be injective, i.e. map distinct logical
244 /// indices to distinct offsets
245 /// ([`StridedError::NonInjectiveOutputLayout`]).
246 pub fn compile(dims: &[usize], dst_strides: &[isize], src_strides: &[isize]) -> Result<Self> {
247 if dims.len() != dst_strides.len() || dims.len() != src_strides.len() {
248 return Err(StridedError::StrideLengthMismatch);
249 }
250 // A zero-sized layout has zero elements even when its other extents
251 // overflow; only a nonempty overflowing count is rejected.
252 crate::kernel::total_len(dims)?;
253 if !crate::layout_check::is_injective_layout(dims, dst_strides) {
254 return Err(StridedError::NonInjectiveOutputLayout);
255 }
256 Ok(Self {
257 dims: dims.into(),
258 dst_strides: dst_strides.into(),
259 src_strides: src_strides.into(),
260 fused: fuse_pair_layout(dims, dst_strides, src_strides),
261 })
262 }
263
264 /// Check the per-call facts: the supplied views must carry exactly the
265 /// compiled layout. Buffer bounds against that layout were already proven
266 /// by [`RawStridedRef::new`]/[`RawStridedMut::new`] (or asserted by the
267 /// caller of the `new_unchecked` constructors), so layout equality is the
268 /// complete precondition for the pointer-based fused replay below.
269 fn check_call<D, S>(
270 &self,
271 dest: &RawStridedMut<'_, D>,
272 src: &RawStridedRef<'_, S>,
273 ) -> Result<()> {
274 if dest.dims() != &self.dims[..]
275 || src.dims() != &self.dims[..]
276 || dest.strides() != &self.dst_strides[..]
277 || src.strides() != &self.src_strides[..]
278 {
279 return Err(StridedError::PlanLayoutMismatch);
280 }
281 Ok(())
282 }
283
284 /// `dest = src` into a potentially uninitialized destination.
285 ///
286 /// On success every reachable logical destination element is initialized.
287 /// Non-reachable holes in the backing allocation are not written.
288 pub fn execute_uninit<T>(
289 &self,
290 dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
291 src: &RawStridedRef<'_, T>,
292 ) -> Result<()>
293 where
294 T: Copy + MaybeSendSync,
295 {
296 self.check_call(dest, src)?;
297 match &self.fused {
298 Some(layout) => {
299 replay_fused(
300 dest,
301 src,
302 layout,
303 |dst, value| {
304 dst.write(value);
305 },
306 |value| value,
307 );
308 Ok(())
309 }
310 None => map_raw_into::<MaybeUninit<T>, T, Identity>(dest, src, MaybeUninit::new),
311 }
312 }
313
314 /// `dest = src`. Allocation-free for ranks at most [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
315 pub fn execute<T>(
316 &self,
317 dest: &mut RawStridedMut<'_, T>,
318 src: &RawStridedRef<'_, T>,
319 ) -> Result<()>
320 where
321 T: Copy + MaybeSendSync,
322 {
323 self.check_call(dest, src)?;
324 match &self.fused {
325 Some(layout) => {
326 replay_fused(
327 dest,
328 src,
329 layout,
330 |dst, value| *dst = value,
331 |value: T| value,
332 );
333 Ok(())
334 }
335 None => copy_into(&mut dest.as_view_mut(), &src.as_view()),
336 }
337 }
338
339 /// `dest = scale * src`. Allocation-free for ranks at most
340 /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
341 pub fn execute_scale<T>(
342 &self,
343 dest: &mut RawStridedMut<'_, T>,
344 src: &RawStridedRef<'_, T>,
345 scale: T,
346 ) -> Result<()>
347 where
348 T: Copy + Mul<T, Output = T> + MaybeSendSync,
349 {
350 self.check_call(dest, src)?;
351 match &self.fused {
352 Some(layout) => {
353 replay_fused(
354 dest,
355 src,
356 layout,
357 |dst, value| *dst = value,
358 |value: T| scale * value,
359 );
360 Ok(())
361 }
362 None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
363 }
364 }
365
366 /// `dest = conj(src)`. Allocation-free for ranks at most
367 /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
368 pub fn execute_conj<T>(
369 &self,
370 dest: &mut RawStridedMut<'_, T>,
371 src: &RawStridedRef<'_, T>,
372 ) -> Result<()>
373 where
374 T: Copy + ElementOpApply + MaybeSendSync,
375 {
376 self.check_call(dest, src)?;
377 match &self.fused {
378 Some(layout) => {
379 replay_fused(
380 dest,
381 src,
382 layout,
383 |dst, value| *dst = value,
384 |value: T| value.conj(),
385 );
386 Ok(())
387 }
388 None => copy_conj(&mut dest.as_view_mut(), &src.as_view()),
389 }
390 }
391}
392
393/// Replay a compiled fused layout, splitting it across worker threads when the
394/// repository threshold and the active execution policy allow it.
395///
396/// Parallel replay is sound only because [`CopyPlan::compile`] proved the
397/// destination layout injective: disjoint logical ranges then write disjoint
398/// destination slots. The logical index space `0..total` (column-major over
399/// the fused axes) is split into contiguous worker ranges, which divides the
400/// outer fused axes and, when there are fewer outer positions than workers,
401/// also chunks long inner runs (for example a rank-1 contiguous copy).
402fn replay_fused<D, S, Apply, Op>(
403 dest: &mut RawStridedMut<'_, D>,
404 src: &RawStridedRef<'_, S>,
405 layout: &FusedPairLayout,
406 apply: Apply,
407 op: Op,
408) where
409 D: Copy + MaybeSendSync,
410 S: Copy + MaybeSendSync,
411 Apply: Fn(&mut D, S) + MaybeSendSync,
412 Op: Fn(S) -> S + MaybeSendSync,
413{
414 let src_ptr = src.data().as_ptr();
415 let src_base = src.offset();
416 let dst_base = dest.offset();
417 let dst_ptr = dest.data_mut().as_mut_ptr();
418 // SAFETY: `check_call` proved both views carry exactly the compiled
419 // layout, and `RawStridedRef`/`RawStridedMut` guarantee every offset
420 // reachable through that layout lies inside their data. `layout` is the
421 // fusion of that layout, so every logical index in `0..total` is an
422 // in-bounds slot of both. The compiled destination layout is injective
423 // and the exclusive destination borrow cannot overlap the shared source
424 // borrow.
425 unsafe { replay_fused_raw(dst_ptr, dst_base, src_ptr, src_base, layout, apply, op) }
426}
427
428/// Pointer-level fused replay shared by [`CopyPlan`] and the dynamic slice
429/// plans, which replay a fused window layout at a runtime base offset.
430///
431/// Splits `0..total` across workers when the repository threshold and the
432/// active execution policy allow more than one thread, otherwise runs the
433/// serial kernel directly.
434///
435/// # Safety
436///
437/// Every logical index of `layout`, mapped through its destination strides
438/// from `dst_base` and its source strides from `src_base`, must be an
439/// in-bounds slot of the destination and source allocations. The
440/// destination strides must be injective over `layout`, the destination
441/// must not overlap the source, and nothing else may access the destination
442/// slots during the call.
443pub(crate) unsafe fn replay_fused_raw<D, S, Apply, Op>(
444 dst_ptr: *mut D,
445 dst_base: isize,
446 src_ptr: *const S,
447 src_base: isize,
448 layout: &FusedPairLayout,
449 apply: Apply,
450 op: Op,
451) where
452 D: Copy + MaybeSendSync,
453 S: Copy + MaybeSendSync,
454 Apply: Fn(&mut D, S) + MaybeSendSync,
455 Op: Fn(S) -> S + MaybeSendSync,
456{
457 let total = fused_total(layout);
458 if total == 0 {
459 return;
460 }
461 #[cfg(feature = "parallel")]
462 {
463 let nthreads = crate::threading::parallel_threads_for_len(total);
464 if nthreads > 1 {
465 let dst_ptr = crate::threading::SendPtr(dst_ptr);
466 let src_ptr = crate::threading::SendPtr(src_ptr as *mut S);
467 crate::threading::parallel_for_each(0..total, nthreads, &|range| {
468 // SAFETY: the caller contract makes every index in bounds;
469 // the worker ranges are disjoint and the destination strides
470 // injective, so no two workers write the same slot, and the
471 // source is only read.
472 unsafe {
473 apply_fused_range(
474 dst_ptr.as_ptr(),
475 dst_base,
476 src_ptr.as_const(),
477 src_base,
478 layout,
479 range.start,
480 range.len(),
481 &apply,
482 &op,
483 );
484 }
485 });
486 return;
487 }
488 }
489 // SAFETY: forwarded caller contract over the full range `0..total`.
490 unsafe {
491 apply_fused_range(
492 dst_ptr, dst_base, src_ptr, src_base, layout, 0, total, &apply, &op,
493 );
494 }
495}
496
497#[cfg(test)]
498#[path = "copy_plan/tests/tests.rs"]
499mod tests;
500
501#[cfg(all(test, feature = "parallel"))]
502#[path = "copy_plan/tests/parallel_tests.rs"]
503mod parallel_tests;