Skip to main content

strided_basic/
dense_update.rs

1//! Dense column-major update and structural kernels.
2//!
3//! AXPBY and triangular masking are adapted from tenferro-rs d8759f43,
4//! `crates/tenferro-cpu/src/{blas1,structural}.rs` (MIT OR Apache-2.0).
5//! Tensor allocation/placement stay with the caller. Diagonal embedding uses
6//! the existing strided map traversal rather than tenferro's coordinate loop.
7
8use core::mem::MaybeUninit;
9use core::ops::{Add, Mul};
10
11use crate::{MaybeSendSync, Result, StridedError, StridedView, StridedViewMut};
12
13fn same_len(actual: usize, expected: usize) -> Result<()> {
14    if actual == expected {
15        Ok(())
16    } else {
17        Err(StridedError::ShapeMismatch(vec![actual], vec![expected]))
18    }
19}
20
21fn product(shape: &[usize]) -> Result<usize> {
22    if shape.contains(&0) {
23        return Ok(0);
24    }
25    shape.iter().try_fold(1usize, |n, &d| {
26        n.checked_mul(d).ok_or(StridedError::OffsetOverflow)
27    })
28}
29
30fn for_each_chunk<T: MaybeSendSync, F: Fn(usize, &mut [T]) + crate::MaybeSync>(
31    output: &mut [T],
32    operation: F,
33) {
34    #[cfg(feature = "parallel")]
35    {
36        let threads = crate::threading::parallel_threads_for_len(output.len());
37        if threads > 1 {
38            let ptr = crate::threading::SendPtr(output.as_mut_ptr());
39            crate::threading::parallel_for_each(0..output.len(), threads, &|range| {
40                // SAFETY: the scoped scheduler partitions the exclusive output
41                // borrow into disjoint ranges; all tasks finish before return.
42                let chunk = unsafe {
43                    core::slice::from_raw_parts_mut(ptr.as_ptr().add(range.start), range.len())
44                };
45                operation(range.start, chunk);
46            });
47            return;
48        }
49    }
50    operation(0, output);
51}
52
53/// Update caller-owned contiguous storage: `y = alpha * x + beta * y`.
54///
55/// This is an accumulation operation: it reads the old `y`, even for beta=0
56/// (including ordinary IEEE NaN propagation). No temporary tensor is allocated.
57/// Execution obeys the active execution policy and shared parallel threshold.
58///
59/// # Examples
60/// ```
61/// let mut y = [3.0, 4.0];
62/// strided_basic::axpby_accum(&mut y, &[1.0, 2.0], 2.0, 3.0).unwrap();
63/// assert_eq!(y, [11.0, 16.0]);
64/// ```
65/// # Errors
66/// Returns `ShapeMismatch` for unequal slice lengths, before modifying `y`.
67pub fn axpby_accum<T>(y: &mut [T], x: &[T], alpha: T, beta: T) -> Result<()>
68where
69    T: Copy + Send + Sync + Add<Output = T> + Mul<Output = T>,
70{
71    same_len(x.len(), y.len())?;
72    for_each_chunk(y, |start, dst| {
73        let len = dst.len();
74        for (out, &src) in dst.iter_mut().zip(&x[start..start + len]) {
75            *out = alpha * src + beta * *out;
76        }
77    });
78    Ok(())
79}
80
81/// Copy dense column-major matrices, replacing the masked triangle with `fill`.
82///
83/// The first two dimensions are rows and columns; remaining dimensions are
84/// batches. `upper=true` keeps `row <= column-k`, otherwise the lower triangle.
85/// Every output element is initialized on success; old output is never read.
86///
87/// # Examples
88/// ```
89/// use core::mem::MaybeUninit;
90/// let mut out = [MaybeUninit::uninit(); 4];
91/// strided_basic::triangular_mask_into_uninit(
92///     &mut out, &[1, 2, 3, 4], &[2, 2], 0, false, 0).unwrap();
93/// // SAFETY: successful full-overwrite kernel initialized every element.
94/// assert_eq!(out.map(|v| unsafe { v.assume_init() }), [1, 2, 0, 4]);
95/// ```
96/// # Errors
97/// Returns `RankMismatch`, `ShapeMismatch`, or `OffsetOverflow` for invalid
98/// shape/buffer lengths, before writing output.
99pub fn triangular_mask_into_uninit<T: Copy + MaybeSendSync>(
100    output: &mut [MaybeUninit<T>],
101    input: &[T],
102    shape: &[usize],
103    k: i64,
104    upper: bool,
105    fill: T,
106) -> Result<()> {
107    if shape.len() < 2 {
108        return Err(StridedError::RankMismatch(shape.len(), 2));
109    }
110    let len = product(shape)?;
111    same_len(input.len(), len)?;
112    same_len(output.len(), len)?;
113    if len == 0 {
114        return Ok(());
115    }
116    let rows = shape[0];
117    let cols = shape[1];
118    for_each_chunk(output, |start, chunk| {
119        let mut pos = 0;
120        while pos < chunk.len() {
121            let flat = start + pos;
122            let row = flat % rows;
123            let col = (flat / rows) % cols;
124            let count = (rows - row).min(chunk.len() - pos);
125            let boundary = col as i128 - k as i128;
126            let split = if upper {
127                boundary.saturating_add(1)
128            } else {
129                boundary
130            }
131            .clamp(row as i128, (row + count) as i128) as usize
132                - row;
133            let (low, high) = chunk[pos..pos + count].split_at_mut(split);
134            let (kept, source, masked) = if upper {
135                (low, &input[flat..flat + split], high)
136            } else {
137                (high, &input[flat + split..flat + count], low)
138            };
139            // Full overwrite: never read masked input or copy it only to
140            // overwrite it a second time. Each output element is written once.
141            for (dst, &src) in kept.iter_mut().zip(source) {
142                dst.write(src);
143            }
144            masked.fill(MaybeUninit::new(fill));
145            pos += count;
146        }
147    });
148    Ok(())
149}
150
151/// Embed a dense column-major input on a diagonal in a new axis.
152///
153/// Insert an axis of size `shape[axis]` at `insert_axis`. Its coordinate must
154/// equal the original `axis` coordinate; other entries receive `zero`.
155/// Output is fully initialized without reading its previous contents. The
156/// active execution policy governs fill and strided diagonal-copy traversal.
157///
158/// # Examples
159/// ```
160/// use core::mem::MaybeUninit;
161/// let mut out = [MaybeUninit::uninit(); 4];
162/// strided_basic::embed_diagonal_into_uninit(&mut out, &[3, 5], &[2], 0, 1, 0).unwrap();
163/// // SAFETY: successful full-overwrite kernel initialized every element.
164/// assert_eq!(out.map(|v| unsafe { v.assume_init() }), [3, 0, 0, 5]);
165/// ```
166/// # Errors
167/// Returns `InvalidAxis`, `ShapeMismatch`, or `OffsetOverflow` for invalid
168/// axes, lengths or unrepresentable layouts. Destination layout validation
169/// may return `NonInjectiveOutputLayout` if injectivity cannot be established.
170pub fn embed_diagonal_into_uninit<T: Copy + MaybeSendSync + 'static>(
171    output: &mut [MaybeUninit<T>],
172    input: &[T],
173    shape: &[usize],
174    axis: usize,
175    insert_axis: usize,
176    zero: T,
177) -> Result<()> {
178    if axis >= shape.len() {
179        return Err(StridedError::InvalidAxis {
180            axis,
181            rank: shape.len(),
182        });
183    }
184    if insert_axis > shape.len() {
185        return Err(StridedError::InvalidAxis {
186            axis: insert_axis,
187            rank: shape.len() + 1,
188        });
189    }
190    let len = product(shape)?;
191    same_len(input.len(), len)?;
192    let out_len = len
193        .checked_mul(shape[axis])
194        .ok_or(StridedError::OffsetOverflow)?;
195    same_len(output.len(), out_len)?;
196    if out_len == 0 {
197        return Ok(());
198    }
199    isize::try_from(out_len).map_err(|_| StridedError::OffsetOverflow)?;
200    let mut out_shape = shape.to_vec();
201    out_shape.insert(insert_axis, shape[axis]);
202    // INVARIANT: nonempty products fit isize, so every prefix stride fits too.
203    let src_strides = crate::col_major_strides(shape);
204    let mut dst_strides = crate::col_major_strides(&out_shape);
205    let inserted_stride = dst_strides.remove(insert_axis);
206    dst_strides[axis] = dst_strides[axis]
207        .checked_add(inserted_stride)
208        .ok_or(StridedError::OffsetOverflow)?;
209    let src = StridedView::<T>::new(input, shape, &src_strides, 0)?;
210    let mut dst = StridedViewMut::new(output, shape, &dst_strides, 0)?;
211    crate::map_view::validate_destination_layout_without_alloc(shape, &dst_strides)?;
212    // INVARIANT: off-diagonal zeros are part of the mathematical result, not
213    // scratch initialization. Only the diagonal subset is overwritten below.
214    for_each_chunk(dst.data_mut(), |_, chunk| {
215        chunk.fill(MaybeUninit::new(zero))
216    });
217    crate::map_into(&mut dst, &src, MaybeUninit::new)
218}