1use crate::ops_view::{axpy, copy_scale};
12use crate::{ElementOpApply, RawStridedMut, RawStridedRef, Result};
13use core::ops::{Add, Mul};
14
15use crate::maybe_sync::MaybeSendSync;
16
17pub const RAW_FUSED_RANK_LIMIT: usize = 8;
19
20#[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
161pub 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
186pub 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
211pub 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
236pub 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}