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 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 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 crate::kernel::total_len(dst)?;
173 Ok(())
174}
175
176pub 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
201pub 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
226pub 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
251pub 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;