strided_basic/
dense_update.rs1use 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 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
53pub 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
81pub 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 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
151pub 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 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 for_each_chunk(dst.data_mut(), |_, chunk| {
215 chunk.fill(MaybeUninit::new(zero))
216 });
217 crate::map_into(&mut dst, &src, MaybeUninit::new)
218}