Skip to main content

strided_basic/
raw_ops.rs

1//! Allocation-free copy/axpy over borrowed raw strided layouts.
2//!
3//! [`StridedView`]/[`StridedViewMut`] own their metadata (`Arc<[usize]>` /
4//! `Arc<[isize]>`) and the map/zip kernels build a traversal plan per call;
5//! for small replay copies that fixed cost dominates. These entry points take
6//! [`RawStridedRef`]/[`RawStridedMut`] (borrowed metadata), fuse the stride
7//! pair into a stack-allocated loop nest, and run plain loops - no heap
8//! allocation on any call path with rank at most [`RAW_FUSED_RANK_LIMIT`].
9//! Higher ranks fall back to the view-based kernels.
10
11use crate::ops_view::{axpy, copy_scale};
12use crate::{ElementOpApply, RawStridedMut, RawStridedRef, Result};
13use core::ops::{Add, Mul};
14
15use crate::maybe_sync::MaybeSendSync;
16
17/// Maximum rank fused on the stack before falling back to the view kernels.
18pub const RAW_FUSED_RANK_LIMIT: usize = 8;
19
20/// Stack-allocated fused stride pair (dims ordered by destination stride,
21/// adjacent contiguous axes merged). Built once and replayed by both the
22/// per-call raw kernels and the prepared [`crate::CopyPlan`].
23#[derive(Clone, Copy, Debug)]
24pub(crate) struct FusedPairLayout {
25    pub(crate) rank: usize,
26    pub(crate) dims: [usize; RAW_FUSED_RANK_LIMIT],
27    pub(crate) dst_strides: [isize; RAW_FUSED_RANK_LIMIT],
28    pub(crate) src_strides: [isize; RAW_FUSED_RANK_LIMIT],
29}
30
31pub(crate) fn fuse_pair_layout(
32    dims: &[usize],
33    dst_strides: &[isize],
34    src_strides: &[isize],
35) -> Option<FusedPairLayout> {
36    if dims.len() > RAW_FUSED_RANK_LIMIT {
37        return None;
38    }
39    let mut layout = FusedPairLayout {
40        rank: 0,
41        dims: [1; RAW_FUSED_RANK_LIMIT],
42        dst_strides: [0; RAW_FUSED_RANK_LIMIT],
43        src_strides: [0; RAW_FUSED_RANK_LIMIT],
44    };
45    for axis in 0..dims.len() {
46        if dims[axis] == 1 {
47            continue;
48        }
49        if dims[axis] == 0 {
50            return Some(FusedPairLayout {
51                rank: 1,
52                dims: [0; RAW_FUSED_RANK_LIMIT],
53                dst_strides: [0; RAW_FUSED_RANK_LIMIT],
54                src_strides: [0; RAW_FUSED_RANK_LIMIT],
55            });
56        }
57        let mut position = layout.rank;
58        while position > 0 && layout.dst_strides[position - 1] > dst_strides[axis] {
59            layout.dims[position] = layout.dims[position - 1];
60            layout.dst_strides[position] = layout.dst_strides[position - 1];
61            layout.src_strides[position] = layout.src_strides[position - 1];
62            position -= 1;
63        }
64        layout.dims[position] = dims[axis];
65        layout.dst_strides[position] = dst_strides[axis];
66        layout.src_strides[position] = src_strides[axis];
67        layout.rank += 1;
68    }
69    if layout.rank == 0 {
70        layout.rank = 1;
71        layout.dims[0] = 1;
72    }
73    let mut fused = 0usize;
74    for axis in 1..layout.rank {
75        // A merged extent that overflows (possible for a stride-0 broadcast
76        // whose element count exceeds usize) leaves the axes unfused.
77        let merged = isize::try_from(layout.dims[fused])
78            .ok()
79            .filter(|&extent| {
80                layout.dst_strides[fused].checked_mul(extent) == Some(layout.dst_strides[axis])
81                    && layout.src_strides[fused].checked_mul(extent)
82                        == Some(layout.src_strides[axis])
83            })
84            .and_then(|_| layout.dims[fused].checked_mul(layout.dims[axis]));
85        if let Some(merged) = merged {
86            layout.dims[fused] = merged;
87        } else {
88            fused += 1;
89            layout.dims[fused] = layout.dims[axis];
90            layout.dst_strides[fused] = layout.dst_strides[axis];
91            layout.src_strides[fused] = layout.src_strides[axis];
92        }
93    }
94    layout.rank = fused + 1;
95    Some(layout)
96}
97
98/// Number of logical elements covered by a fused layout.
99///
100/// Every constructor of a [`FusedPairLayout`] starts from dims whose product
101/// was checked by [`crate::kernel::total_len`], and fusion only merges axes
102/// with a checked product, so the product cannot overflow.
103#[inline]
104pub(crate) fn fused_total(layout: &FusedPairLayout) -> usize {
105    layout.dims[..layout.rank].iter().product()
106}
107
108/// Serial replay of a fused pair layout over borrowed raw views.
109///
110/// The destination layout need not be injective: runs are replayed in order,
111/// so accumulating callers such as [`axpy_raw`] keep their sequential
112/// semantics. Parallel replay is only offered by [`crate::CopyPlan`], whose
113/// compile step proves destination injectivity.
114pub(crate) fn apply_fused_pair<D, S, Apply, Op>(
115    dst: &mut RawStridedMut<'_, D>,
116    src: &RawStridedRef<'_, S>,
117    layout: &FusedPairLayout,
118    apply: Apply,
119    op: Op,
120) where
121    D: Copy,
122    S: Copy,
123    Apply: Fn(&mut D, S),
124    Op: Fn(S) -> S,
125{
126    let total = fused_total(layout);
127    if total == 0 {
128        return;
129    }
130    let src_ptr = src.data().as_ptr();
131    let src_base = src.offset();
132    let dst_base = dst.offset();
133    let dst_ptr = dst.data_mut().as_mut_ptr();
134    // SAFETY: `RawStridedRef`/`RawStridedMut` guarantee (by `new`, or by the
135    // caller of `new_unchecked`) that every offset reachable from their offset
136    // through their dims/strides lies inside their data. Callers pass a layout
137    // fused from exactly those dims/strides, so every logical index in
138    // `0..total` maps to in-bounds source and destination slots. The shared
139    // and exclusive borrows cannot overlap, and this is the only writer.
140    unsafe {
141        apply_fused_range(
142            dst_ptr, dst_base, src_ptr, src_base, layout, 0, total, &apply, &op,
143        );
144    }
145}
146
147/// Replay logical indices `start..start + len` of a fused pair layout, in
148/// column-major order of the fused axes (axis 0 fastest).
149///
150/// The start coordinate is decoded once; runs along axis 0 then advance the
151/// outer coordinates incrementally, so no per-element range checks or
152/// coordinate rebuilds remain in the hot loop.
153///
154/// # Safety
155///
156/// - `start + len` must not exceed [`fused_total`] of `layout`.
157/// - For every logical index `i` in the range, `dst_base` plus the fused
158///   destination offset of `i` must be an in-bounds, writable slot of the
159///   allocation behind `dst_ptr`, and likewise `src_base` plus the source
160///   offset must be an in-bounds, readable slot behind `src_ptr`.
161/// - No other thread may access the destination slots of this range while it
162///   runs, and the source slots must not be written concurrently. Destination
163///   and source slots must not overlap.
164#[allow(clippy::too_many_arguments)]
165pub(crate) unsafe fn apply_fused_range<D, S, Apply, Op>(
166    dst_ptr: *mut D,
167    dst_base: isize,
168    src_ptr: *const S,
169    src_base: isize,
170    layout: &FusedPairLayout,
171    start: usize,
172    len: usize,
173    apply: &Apply,
174    op: &Op,
175) where
176    D: Copy,
177    S: Copy,
178    Apply: Fn(&mut D, S),
179    Op: Fn(S) -> S,
180{
181    if len == 0 {
182        return;
183    }
184    let rank = layout.rank;
185    let inner_len = layout.dims[0];
186    let inner_dst = layout.dst_strides[0];
187    let inner_src = layout.src_strides[0];
188
189    // Decode the start coordinate once per range. `dst_offset`/`src_offset`
190    // track the offset of the current run start (outer coordinates plus the
191    // inner coordinate `inner_start`); every value is a reachable offset.
192    let mut index = [0usize; RAW_FUSED_RANK_LIMIT];
193    let mut rest = start;
194    let mut dst_outer = dst_base;
195    let mut src_outer = src_base;
196    for axis in 0..rank {
197        let dim = layout.dims[axis];
198        index[axis] = rest % dim;
199        rest /= dim;
200        if axis > 0 {
201            dst_outer += index[axis] as isize * layout.dst_strides[axis];
202            src_outer += index[axis] as isize * layout.src_strides[axis];
203        }
204    }
205    let mut inner_start = index[0];
206    let mut remaining = len;
207    loop {
208        let run = (inner_len - inner_start).min(remaining);
209        let dst_run = dst_outer + inner_start as isize * inner_dst;
210        let src_run = src_outer + inner_start as isize * inner_src;
211        // SAFETY: the run covers logical indices inside the caller's range, so
212        // by this function's contract every slot touched below is in-bounds,
213        // exclusively owned by this call (destination) and unaliased.
214        unsafe {
215            apply_run(
216                dst_ptr, dst_run, inner_dst, src_ptr, src_run, inner_src, run, apply, op,
217            )
218        };
219        remaining -= run;
220        if remaining == 0 {
221            return;
222        }
223        inner_start = 0;
224        // Advance an axis only when another position along it follows, and
225        // rewind it from its last position, so every intermediate base is a
226        // reachable offset of a validated layout (issue #243 follow-up): a
227        // layout ending within one stride of `isize::MAX` must not overflow.
228        // `remaining > 0` guarantees a following run exists, so the loop
229        // always finds an axis to advance before running out of rank.
230        let mut axis = 1;
231        while axis < rank {
232            if index[axis] + 1 < layout.dims[axis] {
233                index[axis] += 1;
234                dst_outer += layout.dst_strides[axis];
235                src_outer += layout.src_strides[axis];
236                break;
237            }
238            let last = (layout.dims[axis] - 1) as isize;
239            dst_outer -= last * layout.dst_strides[axis];
240            src_outer -= last * layout.src_strides[axis];
241            index[axis] = 0;
242            axis += 1;
243        }
244    }
245}
246
247/// One run along the fused inner axis.
248///
249/// # Safety
250///
251/// Every `dst_offset + k * dst_stride` and `src_offset + k * src_stride` for
252/// `k < len` must be an in-bounds slot as described in [`apply_fused_range`],
253/// with exclusive access to the destination slots.
254#[allow(clippy::too_many_arguments)]
255#[inline(always)]
256unsafe fn apply_run<D, S, Apply, Op>(
257    dst_ptr: *mut D,
258    dst_offset: isize,
259    dst_stride: isize,
260    src_ptr: *const S,
261    src_offset: isize,
262    src_stride: isize,
263    len: usize,
264    apply: &Apply,
265    op: &Op,
266) where
267    D: Copy,
268    S: Copy,
269    Apply: Fn(&mut D, S),
270    Op: Fn(S) -> S,
271{
272    if dst_stride == 1 {
273        // SAFETY: a unit destination stride makes the run `len` consecutive
274        // in-bounds slots starting at `dst_offset`, exclusively owned here.
275        let dst_run = unsafe { core::slice::from_raw_parts_mut(dst_ptr.offset(dst_offset), len) };
276        match src_stride {
277            1 => {
278                // SAFETY: `len` consecutive in-bounds source slots.
279                let src_run =
280                    unsafe { core::slice::from_raw_parts(src_ptr.offset(src_offset), len) };
281                for (dst, &value) in dst_run.iter_mut().zip(src_run) {
282                    apply(dst, op(value));
283                }
284            }
285            -1 => {
286                // SAFETY: the run reads `src_offset - (len - 1) ..= src_offset`,
287                // all in-bounds; the lowest one is the run's last source slot.
288                let src_run = unsafe {
289                    core::slice::from_raw_parts(
290                        src_ptr.offset(src_offset - (len as isize - 1)),
291                        len,
292                    )
293                };
294                for (dst, &value) in dst_run.iter_mut().zip(src_run.iter().rev()) {
295                    apply(dst, op(value));
296                }
297            }
298            _ => {
299                let src_start = unsafe { src_ptr.offset(src_offset) };
300                for (position, dst) in dst_run.iter_mut().enumerate() {
301                    // SAFETY: `position < len`, so this is a run source slot.
302                    let value = unsafe { *src_start.offset(position as isize * src_stride) };
303                    apply(dst, op(value));
304                }
305            }
306        }
307        return;
308    }
309    // SAFETY: the run start is an in-bounds slot of each allocation.
310    let dst_start = unsafe { dst_ptr.offset(dst_offset) };
311    let src_start = unsafe { src_ptr.offset(src_offset) };
312    for position in 0..len as isize {
313        // SAFETY: `position < len`, so both are slots of this run.
314        unsafe {
315            let value = *src_start.offset(position * src_stride);
316            apply(&mut *dst_start.offset(position * dst_stride), op(value));
317        }
318    }
319}
320
321fn ensure_same_dims(dst: &[usize], src: &[usize]) -> Result<()> {
322    if dst != src {
323        return Err(crate::StridedError::ShapeMismatch(
324            dst.to_vec(),
325            src.to_vec(),
326        ));
327    }
328    // Reject an element count beyond usize (a huge stride-0 broadcast)
329    // instead of replaying it.
330    crate::kernel::total_len(dst)?;
331    Ok(())
332}
333
334/// `dest = scale * src` over borrowed raw strided layouts.
335pub fn copy_scale_raw<T>(
336    dest: &mut RawStridedMut<'_, T>,
337    src: &RawStridedRef<'_, T>,
338    scale: T,
339) -> Result<()>
340where
341    T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
342{
343    ensure_same_dims(dest.dims(), src.dims())?;
344    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
345        Some(layout) => {
346            apply_fused_pair(
347                dest,
348                src,
349                &layout,
350                |dst, value| *dst = value,
351                |value: T| scale * value,
352            );
353            Ok(())
354        }
355        None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
356    }
357}
358
359/// `dest = scale * conj(src)` over borrowed raw strided layouts.
360pub fn copy_scale_conj_raw<T>(
361    dest: &mut RawStridedMut<'_, T>,
362    src: &RawStridedRef<'_, T>,
363    scale: T,
364) -> Result<()>
365where
366    T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
367{
368    ensure_same_dims(dest.dims(), src.dims())?;
369    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
370        Some(layout) => {
371            apply_fused_pair(
372                dest,
373                src,
374                &layout,
375                |dst, value| *dst = value,
376                |value: T| scale * value.conj(),
377            );
378            Ok(())
379        }
380        None => copy_scale(&mut dest.as_view_mut(), &src.as_view().conj(), scale),
381    }
382}
383
384/// `dest = alpha * src + dest` over borrowed raw strided layouts.
385pub fn axpy_raw<T>(
386    dest: &mut RawStridedMut<'_, T>,
387    src: &RawStridedRef<'_, T>,
388    alpha: T,
389) -> Result<()>
390where
391    T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
392{
393    ensure_same_dims(dest.dims(), src.dims())?;
394    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
395        Some(layout) => {
396            apply_fused_pair(
397                dest,
398                src,
399                &layout,
400                |dst, value| *dst = *dst + value,
401                |value: T| alpha * value,
402            );
403            Ok(())
404        }
405        None => axpy(&mut dest.as_view_mut(), &src.as_view(), alpha),
406    }
407}
408
409/// `dest = alpha * conj(src) + dest` over borrowed raw strided layouts.
410pub fn axpy_conj_raw<T>(
411    dest: &mut RawStridedMut<'_, T>,
412    src: &RawStridedRef<'_, T>,
413    alpha: T,
414) -> Result<()>
415where
416    T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
417{
418    ensure_same_dims(dest.dims(), src.dims())?;
419    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
420        Some(layout) => {
421            apply_fused_pair(
422                dest,
423                src,
424                &layout,
425                |dst, value| *dst = *dst + value,
426                |value: T| alpha * value.conj(),
427            );
428            Ok(())
429        }
430        None => axpy(&mut dest.as_view_mut(), &src.as_view().conj(), alpha),
431    }
432}
433
434#[cfg(test)]
435#[path = "raw_ops/tests/tests.rs"]
436mod tests;