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
98#[inline]
104pub(crate) fn fused_total(layout: &FusedPairLayout) -> usize {
105 layout.dims[..layout.rank].iter().product()
106}
107
108pub(crate) fn apply_fused_pair<D, S, Apply, Op>(
115 dst: &mut RawStridedMut<'_, D>,
116 src: &RawStridedRef<'_, S>,
117 layout: &FusedPairLayout,
118 apply: Apply,
119 op: Op,
120) where
121 D: Copy,
122 S: Copy,
123 Apply: Fn(&mut D, S),
124 Op: Fn(S) -> S,
125{
126 let total = fused_total(layout);
127 if total == 0 {
128 return;
129 }
130 let src_ptr = src.data().as_ptr();
131 let src_base = src.offset();
132 let dst_base = dst.offset();
133 let dst_ptr = dst.data_mut().as_mut_ptr();
134 unsafe {
141 apply_fused_range(
142 dst_ptr, dst_base, src_ptr, src_base, layout, 0, total, &apply, &op,
143 );
144 }
145}
146
147#[allow(clippy::too_many_arguments)]
165pub(crate) unsafe fn apply_fused_range<D, S, Apply, Op>(
166 dst_ptr: *mut D,
167 dst_base: isize,
168 src_ptr: *const S,
169 src_base: isize,
170 layout: &FusedPairLayout,
171 start: usize,
172 len: usize,
173 apply: &Apply,
174 op: &Op,
175) where
176 D: Copy,
177 S: Copy,
178 Apply: Fn(&mut D, S),
179 Op: Fn(S) -> S,
180{
181 if len == 0 {
182 return;
183 }
184 let rank = layout.rank;
185 let inner_len = layout.dims[0];
186 let inner_dst = layout.dst_strides[0];
187 let inner_src = layout.src_strides[0];
188
189 let mut index = [0usize; RAW_FUSED_RANK_LIMIT];
193 let mut rest = start;
194 let mut dst_outer = dst_base;
195 let mut src_outer = src_base;
196 for axis in 0..rank {
197 let dim = layout.dims[axis];
198 index[axis] = rest % dim;
199 rest /= dim;
200 if axis > 0 {
201 dst_outer += index[axis] as isize * layout.dst_strides[axis];
202 src_outer += index[axis] as isize * layout.src_strides[axis];
203 }
204 }
205 let mut inner_start = index[0];
206 let mut remaining = len;
207 loop {
208 let run = (inner_len - inner_start).min(remaining);
209 let dst_run = dst_outer + inner_start as isize * inner_dst;
210 let src_run = src_outer + inner_start as isize * inner_src;
211 unsafe {
215 apply_run(
216 dst_ptr, dst_run, inner_dst, src_ptr, src_run, inner_src, run, apply, op,
217 )
218 };
219 remaining -= run;
220 if remaining == 0 {
221 return;
222 }
223 inner_start = 0;
224 let mut axis = 1;
231 while axis < rank {
232 if index[axis] + 1 < layout.dims[axis] {
233 index[axis] += 1;
234 dst_outer += layout.dst_strides[axis];
235 src_outer += layout.src_strides[axis];
236 break;
237 }
238 let last = (layout.dims[axis] - 1) as isize;
239 dst_outer -= last * layout.dst_strides[axis];
240 src_outer -= last * layout.src_strides[axis];
241 index[axis] = 0;
242 axis += 1;
243 }
244 }
245}
246
247#[allow(clippy::too_many_arguments)]
255#[inline(always)]
256unsafe fn apply_run<D, S, Apply, Op>(
257 dst_ptr: *mut D,
258 dst_offset: isize,
259 dst_stride: isize,
260 src_ptr: *const S,
261 src_offset: isize,
262 src_stride: isize,
263 len: usize,
264 apply: &Apply,
265 op: &Op,
266) where
267 D: Copy,
268 S: Copy,
269 Apply: Fn(&mut D, S),
270 Op: Fn(S) -> S,
271{
272 if dst_stride == 1 {
273 let dst_run = unsafe { core::slice::from_raw_parts_mut(dst_ptr.offset(dst_offset), len) };
276 match src_stride {
277 1 => {
278 let src_run =
280 unsafe { core::slice::from_raw_parts(src_ptr.offset(src_offset), len) };
281 for (dst, &value) in dst_run.iter_mut().zip(src_run) {
282 apply(dst, op(value));
283 }
284 }
285 -1 => {
286 let src_run = unsafe {
289 core::slice::from_raw_parts(
290 src_ptr.offset(src_offset - (len as isize - 1)),
291 len,
292 )
293 };
294 for (dst, &value) in dst_run.iter_mut().zip(src_run.iter().rev()) {
295 apply(dst, op(value));
296 }
297 }
298 _ => {
299 let src_start = unsafe { src_ptr.offset(src_offset) };
300 for (position, dst) in dst_run.iter_mut().enumerate() {
301 let value = unsafe { *src_start.offset(position as isize * src_stride) };
303 apply(dst, op(value));
304 }
305 }
306 }
307 return;
308 }
309 let dst_start = unsafe { dst_ptr.offset(dst_offset) };
311 let src_start = unsafe { src_ptr.offset(src_offset) };
312 for position in 0..len as isize {
313 unsafe {
315 let value = *src_start.offset(position * src_stride);
316 apply(&mut *dst_start.offset(position * dst_stride), op(value));
317 }
318 }
319}
320
321fn ensure_same_dims(dst: &[usize], src: &[usize]) -> Result<()> {
322 if dst != src {
323 return Err(crate::StridedError::ShapeMismatch(
324 dst.to_vec(),
325 src.to_vec(),
326 ));
327 }
328 crate::kernel::total_len(dst)?;
331 Ok(())
332}
333
334pub fn copy_scale_raw<T>(
336 dest: &mut RawStridedMut<'_, T>,
337 src: &RawStridedRef<'_, T>,
338 scale: T,
339) -> Result<()>
340where
341 T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
342{
343 ensure_same_dims(dest.dims(), src.dims())?;
344 match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
345 Some(layout) => {
346 apply_fused_pair(
347 dest,
348 src,
349 &layout,
350 |dst, value| *dst = value,
351 |value: T| scale * value,
352 );
353 Ok(())
354 }
355 None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
356 }
357}
358
359pub fn copy_scale_conj_raw<T>(
361 dest: &mut RawStridedMut<'_, T>,
362 src: &RawStridedRef<'_, T>,
363 scale: T,
364) -> Result<()>
365where
366 T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
367{
368 ensure_same_dims(dest.dims(), src.dims())?;
369 match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
370 Some(layout) => {
371 apply_fused_pair(
372 dest,
373 src,
374 &layout,
375 |dst, value| *dst = value,
376 |value: T| scale * value.conj(),
377 );
378 Ok(())
379 }
380 None => copy_scale(&mut dest.as_view_mut(), &src.as_view().conj(), scale),
381 }
382}
383
384pub fn axpy_raw<T>(
386 dest: &mut RawStridedMut<'_, T>,
387 src: &RawStridedRef<'_, T>,
388 alpha: T,
389) -> Result<()>
390where
391 T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
392{
393 ensure_same_dims(dest.dims(), src.dims())?;
394 match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
395 Some(layout) => {
396 apply_fused_pair(
397 dest,
398 src,
399 &layout,
400 |dst, value| *dst = *dst + value,
401 |value: T| alpha * value,
402 );
403 Ok(())
404 }
405 None => axpy(&mut dest.as_view_mut(), &src.as_view(), alpha),
406 }
407}
408
409pub fn axpy_conj_raw<T>(
411 dest: &mut RawStridedMut<'_, T>,
412 src: &RawStridedRef<'_, T>,
413 alpha: T,
414) -> Result<()>
415where
416 T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
417{
418 ensure_same_dims(dest.dims(), src.dims())?;
419 match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
420 Some(layout) => {
421 apply_fused_pair(
422 dest,
423 src,
424 &layout,
425 |dst, value| *dst = *dst + value,
426 |value: T| alpha * value.conj(),
427 );
428 Ok(())
429 }
430 None => axpy(&mut dest.as_view_mut(), &src.as_view().conj(), alpha),
431 }
432}
433
434#[cfg(test)]
435#[path = "raw_ops/tests/tests.rs"]
436mod tests;