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
98pub(crate) fn apply_fused_pair<D, S, Apply, Op>(
99    dst: &mut RawStridedMut<'_, D>,
100    src: &RawStridedRef<'_, S>,
101    layout: &FusedPairLayout,
102    apply: Apply,
103    op: Op,
104) where
105    D: Copy,
106    S: Copy,
107    Apply: Fn(&mut D, S),
108    Op: Fn(S) -> S,
109{
110    if layout.dims[..layout.rank].iter().any(|&dim| dim == 0) {
111        return;
112    }
113    let inner_len = layout.dims[0];
114    let inner_dst = layout.dst_strides[0];
115    let inner_src = layout.src_strides[0];
116    let src_data = src.data();
117    let src_offset = src.offset();
118    let dst_offset = dst.offset();
119    let dst_data = dst.data_mut();
120    let mut index = [0usize; RAW_FUSED_RANK_LIMIT];
121    let mut dst_base = dst_offset;
122    let mut src_base = src_offset;
123    loop {
124        if inner_dst == 1 && inner_src == 1 {
125            let dst_start = dst_base as usize;
126            let src_start = src_base as usize;
127            let dst_run = &mut dst_data[dst_start..dst_start + inner_len];
128            let src_run = &src_data[src_start..src_start + inner_len];
129            for position in 0..inner_len {
130                apply(&mut dst_run[position], op(src_run[position]));
131            }
132        } else {
133            for position in 0..inner_len {
134                let dst_position = (dst_base + position as isize * inner_dst) as usize;
135                let src_position = (src_base + position as isize * inner_src) as usize;
136                apply(&mut dst_data[dst_position], op(src_data[src_position]));
137            }
138        }
139        // Advance an axis only when another position along it follows, and
140        // rewind it from its last position, so every intermediate base is a
141        // reachable offset of a validated layout (issue #243 follow-up): a
142        // layout ending within one stride of `isize::MAX` must not overflow.
143        let mut axis = 1;
144        loop {
145            if axis >= layout.rank {
146                return;
147            }
148            if index[axis] + 1 < layout.dims[axis] {
149                index[axis] += 1;
150                dst_base += layout.dst_strides[axis];
151                src_base += layout.src_strides[axis];
152                break;
153            }
154            let last = (layout.dims[axis] - 1) as isize;
155            dst_base -= last * layout.dst_strides[axis];
156            src_base -= last * layout.src_strides[axis];
157            index[axis] = 0;
158            axis += 1;
159        }
160    }
161}
162
163fn ensure_same_dims(dst: &[usize], src: &[usize]) -> Result<()> {
164    if dst != src {
165        return Err(crate::StridedError::ShapeMismatch(
166            dst.to_vec(),
167            src.to_vec(),
168        ));
169    }
170    // Reject an element count beyond usize (a huge stride-0 broadcast)
171    // instead of replaying it.
172    crate::kernel::total_len(dst)?;
173    Ok(())
174}
175
176/// `dest = scale * src` over borrowed raw strided layouts.
177pub fn copy_scale_raw<T>(
178    dest: &mut RawStridedMut<'_, T>,
179    src: &RawStridedRef<'_, T>,
180    scale: T,
181) -> Result<()>
182where
183    T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
184{
185    ensure_same_dims(dest.dims(), src.dims())?;
186    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
187        Some(layout) => {
188            apply_fused_pair(
189                dest,
190                src,
191                &layout,
192                |dst, value| *dst = value,
193                |value: T| scale * value,
194            );
195            Ok(())
196        }
197        None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
198    }
199}
200
201/// `dest = scale * conj(src)` over borrowed raw strided layouts.
202pub fn copy_scale_conj_raw<T>(
203    dest: &mut RawStridedMut<'_, T>,
204    src: &RawStridedRef<'_, T>,
205    scale: T,
206) -> Result<()>
207where
208    T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
209{
210    ensure_same_dims(dest.dims(), src.dims())?;
211    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
212        Some(layout) => {
213            apply_fused_pair(
214                dest,
215                src,
216                &layout,
217                |dst, value| *dst = value,
218                |value: T| scale * value.conj(),
219            );
220            Ok(())
221        }
222        None => copy_scale(&mut dest.as_view_mut(), &src.as_view().conj(), scale),
223    }
224}
225
226/// `dest = alpha * src + dest` over borrowed raw strided layouts.
227pub fn axpy_raw<T>(
228    dest: &mut RawStridedMut<'_, T>,
229    src: &RawStridedRef<'_, T>,
230    alpha: T,
231) -> Result<()>
232where
233    T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
234{
235    ensure_same_dims(dest.dims(), src.dims())?;
236    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
237        Some(layout) => {
238            apply_fused_pair(
239                dest,
240                src,
241                &layout,
242                |dst, value| *dst = *dst + value,
243                |value: T| alpha * value,
244            );
245            Ok(())
246        }
247        None => axpy(&mut dest.as_view_mut(), &src.as_view(), alpha),
248    }
249}
250
251/// `dest = alpha * conj(src) + dest` over borrowed raw strided layouts.
252pub fn axpy_conj_raw<T>(
253    dest: &mut RawStridedMut<'_, T>,
254    src: &RawStridedRef<'_, T>,
255    alpha: T,
256) -> Result<()>
257where
258    T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
259{
260    ensure_same_dims(dest.dims(), src.dims())?;
261    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
262        Some(layout) => {
263            apply_fused_pair(
264                dest,
265                src,
266                &layout,
267                |dst, value| *dst = *dst + value,
268                |value: T| alpha * value.conj(),
269            );
270            Ok(())
271        }
272        None => axpy(&mut dest.as_view_mut(), &src.as_view().conj(), alpha),
273    }
274}
275
276#[cfg(test)]
277#[path = "raw_ops/tests/tests.rs"]
278mod tests;