Skip to main content

strided_kernel/
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        let extent = layout.dims[fused] as isize;
76        if layout.dst_strides[fused] * extent == layout.dst_strides[axis]
77            && layout.src_strides[fused] * extent == layout.src_strides[axis]
78        {
79            layout.dims[fused] *= layout.dims[axis];
80        } else {
81            fused += 1;
82            layout.dims[fused] = layout.dims[axis];
83            layout.dst_strides[fused] = layout.dst_strides[axis];
84            layout.src_strides[fused] = layout.src_strides[axis];
85        }
86    }
87    layout.rank = fused + 1;
88    Some(layout)
89}
90
91pub(crate) fn apply_fused_pair<D, S, Apply, Op>(
92    dst: &mut RawStridedMut<'_, D>,
93    src: &RawStridedRef<'_, S>,
94    layout: &FusedPairLayout,
95    apply: Apply,
96    op: Op,
97) where
98    D: Copy,
99    S: Copy,
100    Apply: Fn(&mut D, S),
101    Op: Fn(S) -> S,
102{
103    if layout.dims[..layout.rank].iter().any(|&dim| dim == 0) {
104        return;
105    }
106    let inner_len = layout.dims[0];
107    let inner_dst = layout.dst_strides[0];
108    let inner_src = layout.src_strides[0];
109    let src_data = src.data();
110    let src_offset = src.offset();
111    let dst_offset = dst.offset();
112    let dst_data = dst.data_mut();
113    let mut index = [0usize; RAW_FUSED_RANK_LIMIT];
114    let mut dst_base = dst_offset;
115    let mut src_base = src_offset;
116    loop {
117        if inner_dst == 1 && inner_src == 1 {
118            let dst_start = dst_base as usize;
119            let src_start = src_base as usize;
120            let dst_run = &mut dst_data[dst_start..dst_start + inner_len];
121            let src_run = &src_data[src_start..src_start + inner_len];
122            for position in 0..inner_len {
123                apply(&mut dst_run[position], op(src_run[position]));
124            }
125        } else {
126            for position in 0..inner_len {
127                let dst_position = (dst_base + position as isize * inner_dst) as usize;
128                let src_position = (src_base + position as isize * inner_src) as usize;
129                apply(&mut dst_data[dst_position], op(src_data[src_position]));
130            }
131        }
132        let mut axis = 1;
133        loop {
134            if axis >= layout.rank {
135                return;
136            }
137            index[axis] += 1;
138            dst_base += layout.dst_strides[axis];
139            src_base += layout.src_strides[axis];
140            if index[axis] < layout.dims[axis] {
141                break;
142            }
143            dst_base -= layout.dims[axis] as isize * layout.dst_strides[axis];
144            src_base -= layout.dims[axis] as isize * layout.src_strides[axis];
145            index[axis] = 0;
146            axis += 1;
147        }
148    }
149}
150
151fn ensure_same_dims(dst: &[usize], src: &[usize]) -> Result<()> {
152    if dst != src {
153        return Err(crate::StridedError::ShapeMismatch(
154            dst.to_vec(),
155            src.to_vec(),
156        ));
157    }
158    Ok(())
159}
160
161/// `dest = scale * src` over borrowed raw strided layouts.
162pub fn copy_scale_raw<T>(
163    dest: &mut RawStridedMut<'_, T>,
164    src: &RawStridedRef<'_, T>,
165    scale: T,
166) -> Result<()>
167where
168    T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
169{
170    ensure_same_dims(dest.dims(), src.dims())?;
171    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
172        Some(layout) => {
173            apply_fused_pair(
174                dest,
175                src,
176                &layout,
177                |dst, value| *dst = value,
178                |value: T| scale * value,
179            );
180            Ok(())
181        }
182        None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
183    }
184}
185
186/// `dest = scale * conj(src)` over borrowed raw strided layouts.
187pub fn copy_scale_conj_raw<T>(
188    dest: &mut RawStridedMut<'_, T>,
189    src: &RawStridedRef<'_, T>,
190    scale: T,
191) -> Result<()>
192where
193    T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
194{
195    ensure_same_dims(dest.dims(), src.dims())?;
196    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
197        Some(layout) => {
198            apply_fused_pair(
199                dest,
200                src,
201                &layout,
202                |dst, value| *dst = value,
203                |value: T| scale * value.conj(),
204            );
205            Ok(())
206        }
207        None => copy_scale(&mut dest.as_view_mut(), &src.as_view().conj(), scale),
208    }
209}
210
211/// `dest = alpha * src + dest` over borrowed raw strided layouts.
212pub fn axpy_raw<T>(
213    dest: &mut RawStridedMut<'_, T>,
214    src: &RawStridedRef<'_, T>,
215    alpha: T,
216) -> Result<()>
217where
218    T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
219{
220    ensure_same_dims(dest.dims(), src.dims())?;
221    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
222        Some(layout) => {
223            apply_fused_pair(
224                dest,
225                src,
226                &layout,
227                |dst, value| *dst = *dst + value,
228                |value: T| alpha * value,
229            );
230            Ok(())
231        }
232        None => axpy(&mut dest.as_view_mut(), &src.as_view(), alpha),
233    }
234}
235
236/// `dest = alpha * conj(src) + dest` over borrowed raw strided layouts.
237pub fn axpy_conj_raw<T>(
238    dest: &mut RawStridedMut<'_, T>,
239    src: &RawStridedRef<'_, T>,
240    alpha: T,
241) -> Result<()>
242where
243    T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
244{
245    ensure_same_dims(dest.dims(), src.dims())?;
246    match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
247        Some(layout) => {
248            apply_fused_pair(
249                dest,
250                src,
251                &layout,
252                |dst, value| *dst = *dst + value,
253                |value: T| alpha * value.conj(),
254            );
255            Ok(())
256        }
257        None => axpy(&mut dest.as_view_mut(), &src.as_view().conj(), alpha),
258    }
259}
260
261#[cfg(test)]
262mod tests {
263    use super::*;
264    use crate::{StridedView, StridedViewMut};
265
266    fn reference_copy_scale(
267        dst: &mut [f64],
268        src: &[f64],
269        dims: &[usize],
270        dst_strides: &[isize],
271        src_strides: &[isize],
272        scale: f64,
273    ) {
274        let mut dest_view = StridedViewMut::new(dst, dims, dst_strides, 0).unwrap();
275        let src_view: StridedView<'_, f64> = StridedView::new(src, dims, src_strides, 0).unwrap();
276        copy_scale(&mut dest_view, &src_view, scale).unwrap();
277    }
278
279    #[test]
280    fn raw_copy_scale_matches_view_kernel() {
281        let dims = [2usize, 3, 2];
282        let src_strides = [1isize, 2, 6];
283        let dst_strides = [6isize, 2, 1];
284        let src: Vec<f64> = (0..12).map(|value| value as f64 - 3.0).collect();
285        let mut expected = vec![0.0; 12];
286        reference_copy_scale(&mut expected, &src, &dims, &dst_strides, &src_strides, 1.5);
287
288        let mut actual = vec![0.0; 12];
289        let mut dest = RawStridedMut::new(&mut actual, &dims, &dst_strides, 0).unwrap();
290        let source = RawStridedRef::new(&src, &dims, &src_strides, 0).unwrap();
291        copy_scale_raw(&mut dest, &source, 1.5).unwrap();
292
293        assert_eq!(actual, expected);
294    }
295
296    #[test]
297    fn raw_axpy_accumulates() {
298        let dims = [4usize];
299        let strides = [1isize];
300        let src = [1.0f64, 2.0, 3.0, 4.0];
301        let mut dst = [10.0f64, 20.0, 30.0, 40.0];
302        let mut dest = RawStridedMut::new(&mut dst, &dims, &strides, 0).unwrap();
303        let source = RawStridedRef::new(&src, &dims, &strides, 0).unwrap();
304        axpy_raw(&mut dest, &source, 2.0).unwrap();
305
306        assert_eq!(dst, [12.0, 24.0, 36.0, 48.0]);
307    }
308
309    #[test]
310    fn raw_copy_scale_conjugates_complex_sources() {
311        use num_complex::Complex64;
312        let dims = [2usize];
313        let strides = [1isize];
314        let src = [Complex64::new(1.0, 2.0), Complex64::new(-3.0, 4.0)];
315        let mut dst = [Complex64::new(0.0, 0.0); 2];
316        let mut dest = RawStridedMut::new(&mut dst, &dims, &strides, 0).unwrap();
317        let source = RawStridedRef::new(&src, &dims, &strides, 0).unwrap();
318        copy_scale_conj_raw(&mut dest, &source, Complex64::new(2.0, 0.0)).unwrap();
319
320        assert_eq!(dst[0], Complex64::new(2.0, -4.0));
321        assert_eq!(dst[1], Complex64::new(-6.0, -8.0));
322    }
323}