Skip to main content

strided_basic/
ops_view.rs

1//! High-level operations on dynamic-rank strided views.
2
3use crate::kernel::{
4    build_plan_fused, ensure_same_shape, for_each_inner_block_preordered, same_contiguous_layout,
5    sequential_contiguous_layout, total_len,
6};
7use crate::map_view::{map_into, zip_map2_into};
8use crate::maybe_sync::{MaybeSendSync, MaybeSync};
9use crate::reduce_view::reduce;
10use crate::simd;
11use crate::view::{StridedView, StridedViewMut};
12use crate::{Result, StridedError};
13use num_traits::{One, Zero};
14use std::mem::MaybeUninit;
15use std::ops::{Add, Mul};
16use strided_view::{ElementOp, ElementOpApply};
17
18#[cfg(feature = "parallel")]
19use crate::fuse::compute_costs;
20#[cfg(feature = "parallel")]
21use crate::threading::{
22    for_each_inner_block_with_offsets, mapreduce_threaded, SendPtr, MINTHREADLENGTH,
23};
24// ============================================================================
25// Stride-specialized inner loop helpers for ops_view
26//
27// When all inner strides are 1 (contiguous in the innermost dimension),
28// we use slice-based iteration so LLVM can auto-vectorize effectively.
29// This mirrors the inner_loop_map* helpers in map_view.rs.
30// ============================================================================
31
32// INVARIANT: validated layouts keep all dereferences in bounds. Wrapping
33// advances permit an unused final cursor outside a reversed/gapped allocation.
34/// Inner loop for add: `dst[i] += Op::apply(src[i])`.
35#[inline(always)]
36unsafe fn inner_loop_add<D: Copy + Add<S, Output = D>, S: Copy, Op: ElementOp<S>>(
37    dp: *mut D,
38    ds: isize,
39    sp: *const S,
40    ss: isize,
41    len: usize,
42) {
43    if ds == 1 && ss == 1 {
44        let dst = std::slice::from_raw_parts_mut(dp, len);
45        let src = std::slice::from_raw_parts(sp, len);
46        simd::dispatch_if_large(len, || {
47            for i in 0..len {
48                dst[i] = dst[i] + Op::apply(src[i]);
49            }
50        });
51    } else {
52        let mut dp = dp;
53        let mut sp = sp;
54        for _ in 0..len {
55            *dp = *dp + Op::apply(*sp);
56            dp = dp.wrapping_offset(ds);
57            sp = sp.wrapping_offset(ss);
58        }
59    }
60}
61
62/// Inner loop for mul: `dst[i] *= Op::apply(src[i])`.
63#[inline(always)]
64unsafe fn inner_loop_mul<D: Copy + Mul<S, Output = D>, S: Copy, Op: ElementOp<S>>(
65    dp: *mut D,
66    ds: isize,
67    sp: *const S,
68    ss: isize,
69    len: usize,
70) {
71    if ds == 1 && ss == 1 {
72        let dst = std::slice::from_raw_parts_mut(dp, len);
73        let src = std::slice::from_raw_parts(sp, len);
74        simd::dispatch_if_large(len, || {
75            for i in 0..len {
76                dst[i] = dst[i] * Op::apply(src[i]);
77            }
78        });
79    } else {
80        let mut dp = dp;
81        let mut sp = sp;
82        for _ in 0..len {
83            *dp = *dp * Op::apply(*sp);
84            dp = dp.wrapping_offset(ds);
85            sp = sp.wrapping_offset(ss);
86        }
87    }
88}
89
90/// Inner loop for axpy: `dst[i] = alpha * Op::apply(src[i]) + dst[i]`.
91#[inline(always)]
92unsafe fn inner_loop_axpy<
93    D: Copy + Add<D, Output = D>,
94    S: Copy,
95    A: Copy + Mul<S, Output = D>,
96    Op: ElementOp<S>,
97>(
98    dp: *mut D,
99    ds: isize,
100    sp: *const S,
101    ss: isize,
102    len: usize,
103    alpha: A,
104) {
105    if ds == 1 && ss == 1 {
106        let dst = std::slice::from_raw_parts_mut(dp, len);
107        let src = std::slice::from_raw_parts(sp, len);
108        simd::dispatch_if_large(len, || {
109            for i in 0..len {
110                dst[i] = alpha * Op::apply(src[i]) + dst[i];
111            }
112        });
113    } else {
114        let mut dp = dp;
115        let mut sp = sp;
116        for _ in 0..len {
117            *dp = alpha * Op::apply(*sp) + *dp;
118            dp = dp.wrapping_offset(ds);
119            sp = sp.wrapping_offset(ss);
120        }
121    }
122}
123
124/// Inner loop for fma: `dst[i] += OpA::apply(a[i]) * OpB::apply(b[i])`.
125#[inline(always)]
126unsafe fn inner_loop_fma<
127    D: Copy + Add<D, Output = D>,
128    A: Copy + Mul<B, Output = D>,
129    B: Copy,
130    OpA: ElementOp<A>,
131    OpB: ElementOp<B>,
132>(
133    dp: *mut D,
134    ds: isize,
135    ap: *const A,
136    a_s: isize,
137    bp: *const B,
138    b_s: isize,
139    len: usize,
140) {
141    if ds == 1 && a_s == 1 && b_s == 1 {
142        let dst = std::slice::from_raw_parts_mut(dp, len);
143        let sa = std::slice::from_raw_parts(ap, len);
144        let sb = std::slice::from_raw_parts(bp, len);
145        simd::dispatch_if_large(len, || {
146            for i in 0..len {
147                dst[i] = dst[i] + OpA::apply(sa[i]) * OpB::apply(sb[i]);
148            }
149        });
150    } else {
151        let mut dp = dp;
152        let mut ap = ap;
153        let mut bp = bp;
154        for _ in 0..len {
155            *dp = *dp + OpA::apply(*ap) * OpB::apply(*bp);
156            dp = dp.wrapping_offset(ds);
157            ap = ap.wrapping_offset(a_s);
158            bp = bp.wrapping_offset(b_s);
159        }
160    }
161}
162
163/// Inner loop for dot: `acc += OpA::apply(a[i]) * OpB::apply(b[i])`.
164#[inline(always)]
165unsafe fn inner_loop_dot<
166    A: Copy + Mul<B, Output = R>,
167    B: Copy,
168    R: Copy + Add<R, Output = R>,
169    OpA: ElementOp<A>,
170    OpB: ElementOp<B>,
171>(
172    ap: *const A,
173    a_s: isize,
174    bp: *const B,
175    b_s: isize,
176    len: usize,
177    mut acc: R,
178) -> R {
179    if a_s == 1 && b_s == 1 {
180        let sa = std::slice::from_raw_parts(ap, len);
181        let sb = std::slice::from_raw_parts(bp, len);
182        simd::dispatch_if_large(len, || {
183            for i in 0..len {
184                acc = acc + OpA::apply(sa[i]) * OpB::apply(sb[i]);
185            }
186        });
187    } else {
188        let mut ap = ap;
189        let mut bp = bp;
190        for _ in 0..len {
191            acc = acc + OpA::apply(*ap) * OpB::apply(*bp);
192            ap = ap.wrapping_offset(a_s);
193            bp = bp.wrapping_offset(b_s);
194        }
195    }
196    acc
197}
198
199/// Copy elements from source to destination: `dest[i] = src[i]`.
200pub fn copy_into<T: Copy + MaybeSendSync, Op: ElementOp<T>>(
201    dest: &mut StridedViewMut<T>,
202    src: &StridedView<T, Op>,
203) -> Result<()> {
204    ensure_same_shape(dest.dims(), src.dims())?;
205
206    let dst_ptr = dest.as_mut_ptr();
207    let src_ptr = src.ptr();
208    let dst_dims = dest.dims();
209    let dst_strides = dest.strides();
210    let src_strides = src.strides();
211
212    if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides])?.is_some() {
213        let len = total_len(dst_dims)?;
214        if Op::IS_IDENTITY {
215            debug_assert!(
216                {
217                    let nbytes = len
218                        .checked_mul(std::mem::size_of::<T>())
219                        .expect("copy size must not overflow");
220                    let dst_start = dst_ptr as usize;
221                    let src_start = src_ptr as usize;
222                    let dst_end = dst_start.saturating_add(nbytes);
223                    let src_end = src_start.saturating_add(nbytes);
224                    dst_end <= src_start || src_end <= dst_start
225                },
226                "overlapping src/dest is not supported"
227            );
228            unsafe { std::ptr::copy_nonoverlapping(src_ptr, dst_ptr, len) };
229        } else {
230            let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
231            let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
232            simd::dispatch_if_large(len, || {
233                for i in 0..len {
234                    dst[i] = Op::apply(src[i]);
235                }
236            });
237        }
238        return Ok(());
239    }
240
241    map_into(dest, src, |x| x)
242}
243
244/// Copy initialized values into a potentially uninitialized destination.
245///
246/// On success every logical destination element is initialized; holes are
247/// untouched. Accepts identity views; use `map_into` for element operations
248/// such as conjugation. Sequential strided f32/f64 copies use the permutation
249/// engine; other types, contiguous and bounded parallel copies retain map.
250///
251/// # Errors
252/// Returns a shape/rank error for mismatched views, `NonInjectiveOutputLayout`
253/// for overlapping destination elements, or `OffsetOverflow` for size overflow.
254///
255/// # Examples
256/// ```
257/// use std::mem::MaybeUninit;
258/// use strided_basic::{copy_into_uninit, StridedView, StridedViewMut};
259/// let src = [1.0_f64, 2.0];
260/// let mut dst = [MaybeUninit::uninit(); 2];
261/// copy_into_uninit(
262///     &mut StridedViewMut::new(&mut dst, &[2], &[1], 0).unwrap(),
263///     &StridedView::new(&src, &[2], &[1], 0).unwrap(),
264/// ).unwrap();
265/// // SAFETY: the successful copy initialized both destination elements.
266/// assert_eq!(unsafe { dst[1].assume_init() }, 2.0);
267/// ```
268pub fn copy_into_uninit<T: Copy + MaybeSendSync + 'static>(
269    dest: &mut StridedViewMut<MaybeUninit<T>>,
270    src: &StridedView<T>,
271) -> Result<()> {
272    ensure_same_shape(dest.dims(), src.dims())?;
273    if !crate::layout_check::is_injective_layout(dest.dims(), dest.strides()) {
274        return Err(StridedError::NonInjectiveOutputLayout);
275    }
276    if src.is_empty() {
277        return Ok(());
278    }
279    let _len = src
280        .dims()
281        .iter()
282        .try_fold(1usize, |n, &dim| n.checked_mul(dim))
283        .ok_or(StridedError::OffsetOverflow)?;
284    // The permutation engine's 4/8-byte paths reinterpret storage as native
285    // floats. Restrict that path to those exact types: arbitrary Copy types can
286    // contain uninitialized padding or have weaker alignment (e.g. Complex32).
287    let native_float = std::any::TypeId::of::<T>() == std::any::TypeId::of::<f32>()
288        || std::any::TypeId::of::<T>() == std::any::TypeId::of::<f64>();
289    if !native_float
290        || sequential_contiguous_layout(dest.dims(), &[dest.strides(), src.strides()])?.is_some()
291    {
292        return map_into(dest, src, MaybeUninit::new);
293    }
294    #[cfg(feature = "parallel")]
295    if crate::threading::parallel_threads_for_len(_len) > 1 {
296        return map_into(dest, src, MaybeUninit::new);
297    }
298    let data = src.data();
299    // SAFETY: MaybeUninit<T> has T's size/alignment and accepts every initialized
300    // T representation. The source borrow and its extent are unchanged. No
301    // reference to initialized T is formed over destination storage.
302    let data =
303        unsafe { std::slice::from_raw_parts(data.as_ptr().cast::<MaybeUninit<T>>(), data.len()) };
304    let source = StridedView::new(data, src.dims(), src.strides(), src.offset())?;
305    copy_into_col_major(dest, &source)
306}
307
308/// Copy elements from `src` to `dst`, optimized for col-major destination.
309///
310/// Delegates to the centralized permutation-copy dispatch for the actual work.
311pub fn copy_into_col_major<T: Copy + MaybeSendSync>(
312    dst: &mut StridedViewMut<T>,
313    src: &StridedView<T>,
314) -> Result<()> {
315    crate::threading::copy_into_col_major(dst, src)
316}
317
318/// Element-wise addition: `dest[i] += src[i]`.
319///
320/// Source may have a different element type from destination.
321pub fn add<
322    D: Copy + Add<S, Output = D> + MaybeSendSync,
323    S: Copy + MaybeSendSync,
324    Op: ElementOp<S>,
325>(
326    dest: &mut StridedViewMut<D>,
327    src: &StridedView<S, Op>,
328) -> Result<()> {
329    ensure_same_shape(dest.dims(), src.dims())?;
330
331    let dst_ptr = dest.as_mut_ptr();
332    let src_ptr = src.ptr();
333    let dst_dims = dest.dims();
334    let dst_strides = dest.strides();
335    let src_strides = src.strides();
336
337    if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides])?.is_some() {
338        let len = total_len(dst_dims)?;
339        let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
340        let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
341        simd::dispatch_if_large(len, || {
342            for i in 0..len {
343                dst[i] = dst[i] + Op::apply(src[i]);
344            }
345        });
346        return Ok(());
347    }
348
349    let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
350    let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<S>());
351
352    let (fused_dims, ordered_strides, plan) =
353        build_plan_fused(dst_dims, &strides_list, Some(0), elem_size);
354
355    #[cfg(feature = "parallel")]
356    {
357        let total = total_len(&fused_dims)?;
358        let nthreads = crate::execution_policy::rayon_threads();
359        if total > MINTHREADLENGTH && nthreads > 1 {
360            let dst_send = SendPtr(dst_ptr);
361            let src_send = SendPtr(src_ptr as *mut S);
362
363            let costs = compute_costs(&ordered_strides);
364            let initial_offsets = vec![0isize; strides_list.len()];
365            return mapreduce_threaded(
366                &fused_dims,
367                &plan.block,
368                &ordered_strides,
369                &initial_offsets,
370                &costs,
371                nthreads,
372                0,
373                1,
374                &|dims, blocks, strides_list, offsets| {
375                    for_each_inner_block_with_offsets(
376                        dims,
377                        blocks,
378                        strides_list,
379                        offsets,
380                        |offsets, len, strides| {
381                            unsafe {
382                                inner_loop_add::<D, S, Op>(
383                                    dst_send.as_ptr().offset(offsets[0]),
384                                    strides[0],
385                                    src_send.as_const().offset(offsets[1]),
386                                    strides[1],
387                                    len,
388                                )
389                            };
390                            Ok(())
391                        },
392                    )
393                },
394            );
395        }
396    }
397
398    let initial_offsets = vec![0isize; ordered_strides.len()];
399    for_each_inner_block_preordered(
400        &fused_dims,
401        &plan.block,
402        &ordered_strides,
403        &initial_offsets,
404        |offsets, len, strides| {
405            unsafe {
406                inner_loop_add::<D, S, Op>(
407                    dst_ptr.offset(offsets[0]),
408                    strides[0],
409                    src_ptr.offset(offsets[1]),
410                    strides[1],
411                    len,
412                )
413            };
414            Ok(())
415        },
416    )
417}
418
419/// Element-wise multiplication: `dest[i] *= src[i]`.
420///
421/// Source may have a different element type from destination.
422pub fn mul<
423    D: Copy + Mul<S, Output = D> + MaybeSendSync,
424    S: Copy + MaybeSendSync,
425    Op: ElementOp<S>,
426>(
427    dest: &mut StridedViewMut<D>,
428    src: &StridedView<S, Op>,
429) -> Result<()> {
430    ensure_same_shape(dest.dims(), src.dims())?;
431
432    let dst_ptr = dest.as_mut_ptr();
433    let src_ptr = src.ptr();
434    let dst_dims = dest.dims();
435    let dst_strides = dest.strides();
436    let src_strides = src.strides();
437
438    if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides])?.is_some() {
439        let len = total_len(dst_dims)?;
440        let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
441        let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
442        simd::dispatch_if_large(len, || {
443            for i in 0..len {
444                dst[i] = dst[i] * Op::apply(src[i]);
445            }
446        });
447        return Ok(());
448    }
449
450    let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
451    let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<S>());
452
453    let (fused_dims, ordered_strides, plan) =
454        build_plan_fused(dst_dims, &strides_list, Some(0), elem_size);
455
456    #[cfg(feature = "parallel")]
457    {
458        let total = total_len(&fused_dims)?;
459        let nthreads = crate::execution_policy::rayon_threads();
460        if total > MINTHREADLENGTH && nthreads > 1 {
461            let dst_send = SendPtr(dst_ptr);
462            let src_send = SendPtr(src_ptr as *mut S);
463
464            let costs = compute_costs(&ordered_strides);
465            let initial_offsets = vec![0isize; strides_list.len()];
466            return mapreduce_threaded(
467                &fused_dims,
468                &plan.block,
469                &ordered_strides,
470                &initial_offsets,
471                &costs,
472                nthreads,
473                0,
474                1,
475                &|dims, blocks, strides_list, offsets| {
476                    for_each_inner_block_with_offsets(
477                        dims,
478                        blocks,
479                        strides_list,
480                        offsets,
481                        |offsets, len, strides| {
482                            unsafe {
483                                inner_loop_mul::<D, S, Op>(
484                                    dst_send.as_ptr().offset(offsets[0]),
485                                    strides[0],
486                                    src_send.as_const().offset(offsets[1]),
487                                    strides[1],
488                                    len,
489                                )
490                            };
491                            Ok(())
492                        },
493                    )
494                },
495            );
496        }
497    }
498
499    let initial_offsets = vec![0isize; ordered_strides.len()];
500    for_each_inner_block_preordered(
501        &fused_dims,
502        &plan.block,
503        &ordered_strides,
504        &initial_offsets,
505        |offsets, len, strides| {
506            unsafe {
507                inner_loop_mul::<D, S, Op>(
508                    dst_ptr.offset(offsets[0]),
509                    strides[0],
510                    src_ptr.offset(offsets[1]),
511                    strides[1],
512                    len,
513                )
514            };
515            Ok(())
516        },
517    )
518}
519
520/// AXPY: `dest[i] = alpha * src[i] + dest[i]`.
521///
522/// Alpha, source, and destination may have different element types.
523pub fn axpy<D, S, A, Op>(
524    dest: &mut StridedViewMut<D>,
525    src: &StridedView<S, Op>,
526    alpha: A,
527) -> Result<()>
528where
529    A: Copy + Mul<S, Output = D> + MaybeSync,
530    D: Copy + Add<D, Output = D> + MaybeSendSync,
531    S: Copy + MaybeSendSync,
532    Op: ElementOp<S>,
533{
534    ensure_same_shape(dest.dims(), src.dims())?;
535
536    let dst_ptr = dest.as_mut_ptr();
537    let src_ptr = src.ptr();
538    let dst_dims = dest.dims();
539    let dst_strides = dest.strides();
540    let src_strides = src.strides();
541
542    if sequential_contiguous_layout(dst_dims, &[dst_strides, src_strides])?.is_some() {
543        let len = total_len(dst_dims)?;
544        let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
545        let src = unsafe { std::slice::from_raw_parts(src_ptr, len) };
546        simd::dispatch_if_large(len, || {
547            for i in 0..len {
548                dst[i] = alpha * Op::apply(src[i]) + dst[i];
549            }
550        });
551        return Ok(());
552    }
553
554    let strides_list: [&[isize]; 2] = [dst_strides, src_strides];
555    let elem_size = std::mem::size_of::<D>().max(std::mem::size_of::<S>());
556
557    let (fused_dims, ordered_strides, plan) =
558        build_plan_fused(dst_dims, &strides_list, Some(0), elem_size);
559
560    #[cfg(feature = "parallel")]
561    {
562        let total = total_len(&fused_dims)?;
563        let nthreads = crate::execution_policy::rayon_threads();
564        if total > MINTHREADLENGTH && nthreads > 1 {
565            let dst_send = SendPtr(dst_ptr);
566            let src_send = SendPtr(src_ptr as *mut S);
567
568            let costs = compute_costs(&ordered_strides);
569            let initial_offsets = vec![0isize; strides_list.len()];
570            return mapreduce_threaded(
571                &fused_dims,
572                &plan.block,
573                &ordered_strides,
574                &initial_offsets,
575                &costs,
576                nthreads,
577                0,
578                1,
579                &|dims, blocks, strides_list, offsets| {
580                    for_each_inner_block_with_offsets(
581                        dims,
582                        blocks,
583                        strides_list,
584                        offsets,
585                        |offsets, len, strides| {
586                            unsafe {
587                                inner_loop_axpy::<D, S, A, Op>(
588                                    dst_send.as_ptr().offset(offsets[0]),
589                                    strides[0],
590                                    src_send.as_const().offset(offsets[1]),
591                                    strides[1],
592                                    len,
593                                    alpha,
594                                )
595                            };
596                            Ok(())
597                        },
598                    )
599                },
600            );
601        }
602    }
603
604    let initial_offsets = vec![0isize; ordered_strides.len()];
605    for_each_inner_block_preordered(
606        &fused_dims,
607        &plan.block,
608        &ordered_strides,
609        &initial_offsets,
610        |offsets, len, strides| {
611            unsafe {
612                inner_loop_axpy::<D, S, A, Op>(
613                    dst_ptr.offset(offsets[0]),
614                    strides[0],
615                    src_ptr.offset(offsets[1]),
616                    strides[1],
617                    len,
618                    alpha,
619                )
620            };
621            Ok(())
622        },
623    )
624}
625
626/// Fused multiply-add: `dest[i] += OpA::apply(a[i]) * OpB::apply(b[i])`.
627///
628/// Operands may have different element types. Element operations are applied lazily.
629pub fn fma<D, A, B, OpA, OpB>(
630    dest: &mut StridedViewMut<D>,
631    a: &StridedView<A, OpA>,
632    b: &StridedView<B, OpB>,
633) -> Result<()>
634where
635    A: Copy + Mul<B, Output = D> + MaybeSendSync,
636    B: Copy + MaybeSendSync,
637    D: Copy + Add<D, Output = D> + MaybeSendSync,
638    OpA: ElementOp<A>,
639    OpB: ElementOp<B>,
640{
641    ensure_same_shape(dest.dims(), a.dims())?;
642    ensure_same_shape(dest.dims(), b.dims())?;
643
644    let dst_ptr = dest.as_mut_ptr();
645    let a_ptr = a.ptr();
646    let b_ptr = b.ptr();
647    let dst_dims = dest.dims();
648    let dst_strides = dest.strides();
649    let a_strides = a.strides();
650    let b_strides = b.strides();
651
652    if sequential_contiguous_layout(dst_dims, &[dst_strides, a_strides, b_strides])?.is_some() {
653        let len = total_len(dst_dims)?;
654        let dst = unsafe { std::slice::from_raw_parts_mut(dst_ptr, len) };
655        let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
656        let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
657        simd::dispatch_if_large(len, || {
658            for i in 0..len {
659                dst[i] = dst[i] + OpA::apply(sa[i]) * OpB::apply(sb[i]);
660            }
661        });
662        return Ok(());
663    }
664
665    let strides_list: [&[isize]; 3] = [dst_strides, a_strides, b_strides];
666    let elem_size = std::mem::size_of::<D>()
667        .max(std::mem::size_of::<A>())
668        .max(std::mem::size_of::<B>());
669
670    let (fused_dims, ordered_strides, plan) =
671        build_plan_fused(dst_dims, &strides_list, Some(0), elem_size);
672
673    #[cfg(feature = "parallel")]
674    {
675        let total = total_len(&fused_dims)?;
676        let nthreads = crate::execution_policy::rayon_threads();
677        if total > MINTHREADLENGTH && nthreads > 1 {
678            let dst_send = SendPtr(dst_ptr);
679            let a_send = SendPtr(a_ptr as *mut A);
680            let b_send = SendPtr(b_ptr as *mut B);
681
682            let costs = compute_costs(&ordered_strides);
683            let initial_offsets = vec![0isize; strides_list.len()];
684            return mapreduce_threaded(
685                &fused_dims,
686                &plan.block,
687                &ordered_strides,
688                &initial_offsets,
689                &costs,
690                nthreads,
691                0,
692                1,
693                &|dims, blocks, strides_list, offsets| {
694                    for_each_inner_block_with_offsets(
695                        dims,
696                        blocks,
697                        strides_list,
698                        offsets,
699                        |offsets, len, strides| {
700                            unsafe {
701                                inner_loop_fma::<D, A, B, OpA, OpB>(
702                                    dst_send.as_ptr().offset(offsets[0]),
703                                    strides[0],
704                                    a_send.as_const().offset(offsets[1]),
705                                    strides[1],
706                                    b_send.as_const().offset(offsets[2]),
707                                    strides[2],
708                                    len,
709                                )
710                            };
711                            Ok(())
712                        },
713                    )
714                },
715            );
716        }
717    }
718
719    let initial_offsets = vec![0isize; ordered_strides.len()];
720    for_each_inner_block_preordered(
721        &fused_dims,
722        &plan.block,
723        &ordered_strides,
724        &initial_offsets,
725        |offsets, len, strides| {
726            unsafe {
727                inner_loop_fma::<D, A, B, OpA, OpB>(
728                    dst_ptr.offset(offsets[0]),
729                    strides[0],
730                    a_ptr.offset(offsets[1]),
731                    strides[1],
732                    b_ptr.offset(offsets[2]),
733                    strides[2],
734                    len,
735                )
736            };
737            Ok(())
738        },
739    )
740}
741
742#[cfg(feature = "parallel")]
743fn parallel_simd_sum<T: Copy + Zero + Add<Output = T> + simd::MaybeSimdOps + Send + Sync>(
744    src: &[T],
745) -> Option<T> {
746    // Check that T has SIMD support
747    if T::try_simd_sum(&[]).is_none() {
748        return None;
749    }
750    let nthreads = crate::execution_policy::rayon_threads();
751    let result = crate::threading::parallel_map_reduce(
752        0..src.len(),
753        nthreads,
754        &|range| T::try_simd_sum(&src[range]).unwrap(),
755        &|left, right| left + right,
756    );
757    Some(result)
758}
759
760/// Sum all elements: `sum(src)`.
761pub fn sum<
762    T: Copy + Zero + Add<Output = T> + MaybeSendSync + simd::MaybeSimdOps,
763    Op: ElementOp<T>,
764>(
765    src: &StridedView<T, Op>,
766) -> Result<T> {
767    // SIMD fast path: contiguous Identity view with SIMD support
768    if Op::IS_IDENTITY {
769        if same_contiguous_layout(src.dims(), &[src.strides()]).is_some() {
770            let len = total_len(src.dims())?;
771            let src_slice = unsafe { std::slice::from_raw_parts(src.ptr(), len) };
772
773            #[cfg(feature = "parallel")]
774            if len > MINTHREADLENGTH {
775                if let Some(result) = parallel_simd_sum(src_slice) {
776                    return Ok(result);
777                }
778            }
779
780            if let Some(result) = T::try_simd_sum(src_slice) {
781                return Ok(result);
782            }
783        }
784    }
785    reduce(src, |x| x, |a, b| a + b, T::zero())
786}
787
788/// Dot product: `sum(OpA::apply(a[i]) * OpB::apply(b[i]))`.
789///
790/// Operands may have different element types. Result type `R` must be `A * B`.
791/// SIMD fast path fires only when `A == B == R` (same type) and both Identity ops.
792pub fn dot<A, B, R, OpA, OpB>(a: &StridedView<A, OpA>, b: &StridedView<B, OpB>) -> Result<R>
793where
794    A: Copy + Mul<B, Output = R> + MaybeSendSync + 'static,
795    B: Copy + MaybeSendSync + 'static,
796    R: Copy + Zero + Add<Output = R> + MaybeSendSync + simd::MaybeSimdOps + 'static,
797    OpA: ElementOp<A>,
798    OpB: ElementOp<B>,
799{
800    ensure_same_shape(a.dims(), b.dims())?;
801
802    let a_ptr = a.ptr();
803    let b_ptr = b.ptr();
804    let a_strides = a.strides();
805    let b_strides = b.strides();
806    let a_dims = a.dims();
807    let len = total_len(a_dims)?;
808
809    if same_contiguous_layout(a_dims, &[a_strides, b_strides]).is_some() {
810        // SIMD fast path: both contiguous, both Identity ops, same type
811        if OpA::IS_IDENTITY
812            && OpB::IS_IDENTITY
813            && std::any::TypeId::of::<A>() == std::any::TypeId::of::<R>()
814            && std::any::TypeId::of::<B>() == std::any::TypeId::of::<R>()
815        {
816            let sa = unsafe { std::slice::from_raw_parts(a_ptr as *const R, len) };
817            let sb = unsafe { std::slice::from_raw_parts(b_ptr as *const R, len) };
818            if let Some(result) = R::try_simd_dot(sa, sb) {
819                return Ok(result);
820            }
821        }
822
823        // Generic contiguous fast path
824        let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
825        let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
826        let mut acc = R::zero();
827        simd::dispatch_if_large(len, || {
828            for i in 0..len {
829                acc = acc + OpA::apply(sa[i]) * OpB::apply(sb[i]);
830            }
831        });
832        return Ok(acc);
833    }
834
835    let strides_list: [&[isize]; 2] = [a_strides, b_strides];
836    let elem_size = std::mem::size_of::<A>()
837        .max(std::mem::size_of::<B>())
838        .max(std::mem::size_of::<R>());
839
840    let (fused_dims, ordered_strides, plan) =
841        build_plan_fused(a_dims, &strides_list, None, elem_size);
842
843    let mut acc = R::zero();
844    let initial_offsets = vec![0isize; ordered_strides.len()];
845    for_each_inner_block_preordered(
846        &fused_dims,
847        &plan.block,
848        &ordered_strides,
849        &initial_offsets,
850        |offsets, len, strides| {
851            acc = unsafe {
852                inner_loop_dot::<A, B, R, OpA, OpB>(
853                    a_ptr.offset(offsets[0]),
854                    strides[0],
855                    b_ptr.offset(offsets[1]),
856                    strides[1],
857                    len,
858                    acc,
859                )
860            };
861            Ok(())
862        },
863    )?;
864
865    Ok(acc)
866}
867
868/// Symmetrize a square matrix: `dest = (src + src^T) / 2`.
869pub fn symmetrize_into<T>(dest: &mut StridedViewMut<T>, src: &StridedView<T>) -> Result<()>
870where
871    T: Copy
872        + Add<Output = T>
873        + Mul<Output = T>
874        + num_traits::FromPrimitive
875        + std::ops::Div<Output = T>
876        + MaybeSendSync,
877{
878    if src.ndim() != 2 {
879        return Err(StridedError::RankMismatch(src.ndim(), 2));
880    }
881    let rows = src.dims()[0];
882    let cols = src.dims()[1];
883    if rows != cols {
884        return Err(StridedError::NonSquare { rows, cols });
885    }
886
887    let src_t = src.permute(&[1, 0])?;
888    let half = T::from_f64(0.5).ok_or(StridedError::ScalarConversion)?;
889
890    zip_map2_into(dest, src, &src_t, |a, b| (a + b) * half)
891}
892
893/// Conjugate-symmetrize a square matrix: `dest = (src + conj(src^T)) / 2`.
894pub fn symmetrize_conj_into<T>(dest: &mut StridedViewMut<T>, src: &StridedView<T>) -> Result<()>
895where
896    T: Copy
897        + ElementOpApply
898        + Add<Output = T>
899        + Mul<Output = T>
900        + num_traits::FromPrimitive
901        + std::ops::Div<Output = T>
902        + MaybeSendSync,
903{
904    if src.ndim() != 2 {
905        return Err(StridedError::RankMismatch(src.ndim(), 2));
906    }
907    let rows = src.dims()[0];
908    let cols = src.dims()[1];
909    if rows != cols {
910        return Err(StridedError::NonSquare { rows, cols });
911    }
912
913    // adjoint = conj + transpose
914    let src_adj = src.adjoint_2d()?;
915    let half = T::from_f64(0.5).ok_or(StridedError::ScalarConversion)?;
916
917    zip_map2_into(dest, src, &src_adj, |a, b| (a + b) * half)
918}
919
920/// Copy with scaling: `dest[i] = scale * src[i]`.
921///
922/// Scale, source, and destination may have different element types.
923pub fn copy_scale<D, S, A, Op>(
924    dest: &mut StridedViewMut<D>,
925    src: &StridedView<S, Op>,
926    scale: A,
927) -> Result<()>
928where
929    A: Copy + Mul<S, Output = D> + MaybeSync,
930    D: Copy + MaybeSendSync,
931    S: Copy + MaybeSendSync,
932    Op: ElementOp<S>,
933{
934    map_into(dest, src, |x| scale * x)
935}
936
937/// Copy with complex conjugation: `dest[i] = conj(src[i])`.
938pub fn copy_conj<T: Copy + ElementOpApply + MaybeSendSync>(
939    dest: &mut StridedViewMut<T>,
940    src: &StridedView<T>,
941) -> Result<()> {
942    let src_conj = src.conj();
943    copy_into(dest, &src_conj)
944}
945
946#[inline]
947fn element_transpose_is_identity<T: 'static>() -> bool {
948    use std::any::TypeId;
949
950    macro_rules! matches_type {
951        ($($ty:ty),* $(,)?) => {{
952            let id = TypeId::of::<T>();
953            false $(|| id == TypeId::of::<$ty>())*
954        }};
955    }
956
957    matches_type!(
958        f32,
959        f64,
960        i8,
961        i16,
962        i32,
963        i64,
964        i128,
965        isize,
966        u8,
967        u16,
968        u32,
969        u64,
970        u128,
971        usize,
972        num_complex::Complex32,
973        num_complex::Complex64,
974    )
975}
976
977#[inline]
978fn element_zero_is_all_bits_zero<T: 'static>() -> bool {
979    use std::any::TypeId;
980
981    macro_rules! matches_type {
982        ($($ty:ty),* $(,)?) => {{
983            let id = TypeId::of::<T>();
984            false $(|| id == TypeId::of::<$ty>())*
985        }};
986    }
987
988    matches_type!(
989        f32,
990        f64,
991        i8,
992        i16,
993        i32,
994        i64,
995        i128,
996        isize,
997        u8,
998        u16,
999        u32,
1000        u64,
1001        u128,
1002        usize,
1003        num_complex::Complex32,
1004        num_complex::Complex64,
1005    )
1006}
1007
1008#[inline]
1009unsafe fn fill_2d<T: Copy + MaybeSendSync>(
1010    dst: *mut T,
1011    dim0: usize,
1012    dim1: usize,
1013    dst_stride0: isize,
1014    dst_stride1: isize,
1015    value: T,
1016) {
1017    #[cfg(feature = "parallel")]
1018    {
1019        let total = dim0.saturating_mul(dim1);
1020        let nthreads = crate::execution_policy::rayon_threads();
1021        if total > MINTHREADLENGTH && nthreads > 1 {
1022            let dst_send = SendPtr(dst);
1023            if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1024                crate::threading::parallel_for_each(0..dim1, nthreads, &|columns| {
1025                    for j in columns {
1026                        let dst = dst_send.as_ptr();
1027                        unsafe {
1028                            let base = j as isize * dst_stride1;
1029                            for i in 0..dim0 {
1030                                *dst.offset(base + i as isize * dst_stride0) = value;
1031                            }
1032                        }
1033                    }
1034                });
1035            } else {
1036                crate::threading::parallel_for_each(0..dim0, nthreads, &|rows| {
1037                    for i in rows {
1038                        let dst = dst_send.as_ptr();
1039                        unsafe {
1040                            let base = i as isize * dst_stride0;
1041                            for j in 0..dim1 {
1042                                *dst.offset(base + j as isize * dst_stride1) = value;
1043                            }
1044                        }
1045                    }
1046                });
1047            }
1048            return;
1049        }
1050    }
1051
1052    if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1053        for j in 0..dim1 {
1054            let base = j as isize * dst_stride1;
1055            for i in 0..dim0 {
1056                *dst.offset(base + i as isize * dst_stride0) = value;
1057            }
1058        }
1059    } else {
1060        for i in 0..dim0 {
1061            let base = i as isize * dst_stride0;
1062            for j in 0..dim1 {
1063                *dst.offset(base + j as isize * dst_stride1) = value;
1064            }
1065        }
1066    }
1067}
1068
1069#[inline]
1070unsafe fn fill_contiguous<T>(dst: *mut T, len: usize, value: T)
1071where
1072    T: Copy + Zero + PartialEq + MaybeSendSync + 'static,
1073{
1074    if element_zero_is_all_bits_zero::<T>() && value == T::zero() {
1075        std::ptr::write_bytes(dst, 0, len);
1076        return;
1077    }
1078
1079    let dst = std::slice::from_raw_parts_mut(dst, len);
1080    dst.fill(value);
1081}
1082
1083#[inline(always)]
1084unsafe fn transpose_scale_4x4_f64(
1085    dst: *mut f64,
1086    dst_stride0: isize,
1087    dst_stride1: isize,
1088    src: *const f64,
1089    src_stride0: isize,
1090    src_stride1: isize,
1091    i: usize,
1092    j: usize,
1093    scale: f64,
1094) {
1095    let src_base = src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1096
1097    let s00 = *src_base;
1098    let s10 = *src_base.offset(src_stride0);
1099    let s20 = *src_base.offset(2 * src_stride0);
1100    let s30 = *src_base.offset(3 * src_stride0);
1101
1102    let src_col1 = src_base.offset(src_stride1);
1103    let s01 = *src_col1;
1104    let s11 = *src_col1.offset(src_stride0);
1105    let s21 = *src_col1.offset(2 * src_stride0);
1106    let s31 = *src_col1.offset(3 * src_stride0);
1107
1108    let src_col2 = src_base.offset(2 * src_stride1);
1109    let s02 = *src_col2;
1110    let s12 = *src_col2.offset(src_stride0);
1111    let s22 = *src_col2.offset(2 * src_stride0);
1112    let s32 = *src_col2.offset(3 * src_stride0);
1113
1114    let src_col3 = src_base.offset(3 * src_stride1);
1115    let s03 = *src_col3;
1116    let s13 = *src_col3.offset(src_stride0);
1117    let s23 = *src_col3.offset(2 * src_stride0);
1118    let s33 = *src_col3.offset(3 * src_stride0);
1119
1120    let dst_row0 = dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1);
1121    *dst_row0 = scale * s00;
1122    *dst_row0.offset(dst_stride0) = scale * s01;
1123    *dst_row0.offset(2 * dst_stride0) = scale * s02;
1124    *dst_row0.offset(3 * dst_stride0) = scale * s03;
1125
1126    let dst_row1 = dst_row0.offset(dst_stride1);
1127    *dst_row1 = scale * s10;
1128    *dst_row1.offset(dst_stride0) = scale * s11;
1129    *dst_row1.offset(2 * dst_stride0) = scale * s12;
1130    *dst_row1.offset(3 * dst_stride0) = scale * s13;
1131
1132    let dst_row2 = dst_row0.offset(2 * dst_stride1);
1133    *dst_row2 = scale * s20;
1134    *dst_row2.offset(dst_stride0) = scale * s21;
1135    *dst_row2.offset(2 * dst_stride0) = scale * s22;
1136    *dst_row2.offset(3 * dst_stride0) = scale * s23;
1137
1138    let dst_row3 = dst_row0.offset(3 * dst_stride1);
1139    *dst_row3 = scale * s30;
1140    *dst_row3.offset(dst_stride0) = scale * s31;
1141    *dst_row3.offset(2 * dst_stride0) = scale * s32;
1142    *dst_row3.offset(3 * dst_stride0) = scale * s33;
1143}
1144
1145#[inline]
1146unsafe fn copy_transpose_scale_2d_f64_tiled_raw(
1147    dst: *mut f64,
1148    dst_stride0: isize,
1149    dst_stride1: isize,
1150    src: *const f64,
1151    src_stride0: isize,
1152    src_stride1: isize,
1153    src_rows: usize,
1154    src_cols: usize,
1155    scale: f64,
1156) {
1157    const TILE: usize = 4;
1158    let row_full = src_rows / TILE * TILE;
1159    let col_full = src_cols / TILE * TILE;
1160
1161    #[cfg(feature = "parallel")]
1162    {
1163        let total = src_rows.saturating_mul(src_cols);
1164        let nthreads = crate::execution_policy::rayon_threads();
1165        if total > MINTHREADLENGTH && nthreads > 1 {
1166            let dst_send = SendPtr(dst);
1167            let src_send = SendPtr(src as *mut f64);
1168            let row_tiles = row_full / TILE;
1169            crate::threading::parallel_for_each(0..row_tiles, nthreads, &|tiles| {
1170                for tile_i in tiles {
1171                    let i = tile_i * TILE;
1172                    let dst = dst_send.as_ptr();
1173                    let src = src_send.as_const();
1174                    unsafe {
1175                        let mut j = 0;
1176                        while j < col_full {
1177                            transpose_scale_4x4_f64(
1178                                dst,
1179                                dst_stride0,
1180                                dst_stride1,
1181                                src,
1182                                src_stride0,
1183                                src_stride1,
1184                                i,
1185                                j,
1186                                scale,
1187                            );
1188                            j += TILE;
1189                        }
1190                        for j in col_full..src_cols {
1191                            for ii in i..i + TILE {
1192                                *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1193                                    scale
1194                                        * *src.offset(
1195                                            ii as isize * src_stride0 + j as isize * src_stride1,
1196                                        );
1197                            }
1198                        }
1199                    }
1200                }
1201            });
1202            for i in row_full..src_rows {
1203                for j in 0..src_cols {
1204                    *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1205                        scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1206                }
1207            }
1208            return;
1209        }
1210    }
1211
1212    let mut i = 0;
1213    while i < row_full {
1214        let mut j = 0;
1215        while j < col_full {
1216            transpose_scale_4x4_f64(
1217                dst,
1218                dst_stride0,
1219                dst_stride1,
1220                src,
1221                src_stride0,
1222                src_stride1,
1223                i,
1224                j,
1225                scale,
1226            );
1227            j += TILE;
1228        }
1229        for j in col_full..src_cols {
1230            for ii in i..i + TILE {
1231                *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1232                    scale * *src.offset(ii as isize * src_stride0 + j as isize * src_stride1);
1233            }
1234        }
1235        i += TILE;
1236    }
1237    for i in row_full..src_rows {
1238        for j in 0..src_cols {
1239            *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1240                scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1241        }
1242    }
1243}
1244
1245#[inline]
1246#[cfg(test)]
1247unsafe fn try_copy_transpose_scale_2d_f64_tiled(
1248    dest: &mut StridedViewMut<f64>,
1249    src: &StridedView<f64>,
1250    scale: f64,
1251) -> bool {
1252    if src.ndim() != 2 || dest.ndim() != 2 {
1253        return false;
1254    }
1255    let src_dims = src.dims();
1256    if dest.dims() != [src_dims[1], src_dims[0]] {
1257        return false;
1258    }
1259    if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1260        return false;
1261    }
1262
1263    copy_transpose_scale_2d_f64_tiled_raw(
1264        dest.as_mut_ptr(),
1265        dest.strides()[0],
1266        dest.strides()[1],
1267        src.ptr(),
1268        src.strides()[0],
1269        src.strides()[1],
1270        src_dims[0],
1271        src_dims[1],
1272        scale,
1273    );
1274    true
1275}
1276
1277#[inline]
1278unsafe fn try_copy_transpose_scale_2d_f64_tiled_typed<T>(
1279    dest: &mut StridedViewMut<T>,
1280    src: &StridedView<T>,
1281    scale: T,
1282) -> bool
1283where
1284    T: Copy + 'static,
1285{
1286    if std::any::TypeId::of::<T>() != std::any::TypeId::of::<f64>() {
1287        return false;
1288    }
1289
1290    let scale = *(&scale as *const T).cast::<f64>();
1291    if src.ndim() != 2 || dest.ndim() != 2 {
1292        return false;
1293    }
1294    let src_dims = src.dims();
1295    if dest.dims() != [src_dims[1], src_dims[0]] {
1296        return false;
1297    }
1298    if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1299        return false;
1300    }
1301
1302    copy_transpose_scale_2d_f64_tiled_raw(
1303        dest.as_mut_ptr().cast::<f64>(),
1304        dest.strides()[0],
1305        dest.strides()[1],
1306        src.ptr().cast::<f64>(),
1307        src.strides()[0],
1308        src.strides()[1],
1309        src_dims[0],
1310        src_dims[1],
1311        scale,
1312    );
1313    true
1314}
1315
1316#[inline(always)]
1317unsafe fn transpose_scale_4x4_identity<T>(
1318    dst: *mut T,
1319    dst_stride0: isize,
1320    dst_stride1: isize,
1321    src: *const T,
1322    src_stride0: isize,
1323    src_stride1: isize,
1324    i: usize,
1325    j: usize,
1326    scale: T,
1327) where
1328    T: Copy + Mul<Output = T>,
1329{
1330    let src_base = src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1331
1332    let s00 = *src_base;
1333    let s10 = *src_base.offset(src_stride0);
1334    let s20 = *src_base.offset(2 * src_stride0);
1335    let s30 = *src_base.offset(3 * src_stride0);
1336
1337    let src_col1 = src_base.offset(src_stride1);
1338    let s01 = *src_col1;
1339    let s11 = *src_col1.offset(src_stride0);
1340    let s21 = *src_col1.offset(2 * src_stride0);
1341    let s31 = *src_col1.offset(3 * src_stride0);
1342
1343    let src_col2 = src_base.offset(2 * src_stride1);
1344    let s02 = *src_col2;
1345    let s12 = *src_col2.offset(src_stride0);
1346    let s22 = *src_col2.offset(2 * src_stride0);
1347    let s32 = *src_col2.offset(3 * src_stride0);
1348
1349    let src_col3 = src_base.offset(3 * src_stride1);
1350    let s03 = *src_col3;
1351    let s13 = *src_col3.offset(src_stride0);
1352    let s23 = *src_col3.offset(2 * src_stride0);
1353    let s33 = *src_col3.offset(3 * src_stride0);
1354
1355    let dst_row0 = dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1);
1356    *dst_row0 = scale * s00;
1357    *dst_row0.offset(dst_stride0) = scale * s01;
1358    *dst_row0.offset(2 * dst_stride0) = scale * s02;
1359    *dst_row0.offset(3 * dst_stride0) = scale * s03;
1360
1361    let dst_row1 = dst_row0.offset(dst_stride1);
1362    *dst_row1 = scale * s10;
1363    *dst_row1.offset(dst_stride0) = scale * s11;
1364    *dst_row1.offset(2 * dst_stride0) = scale * s12;
1365    *dst_row1.offset(3 * dst_stride0) = scale * s13;
1366
1367    let dst_row2 = dst_row0.offset(2 * dst_stride1);
1368    *dst_row2 = scale * s20;
1369    *dst_row2.offset(dst_stride0) = scale * s21;
1370    *dst_row2.offset(2 * dst_stride0) = scale * s22;
1371    *dst_row2.offset(3 * dst_stride0) = scale * s23;
1372
1373    let dst_row3 = dst_row0.offset(3 * dst_stride1);
1374    *dst_row3 = scale * s30;
1375    *dst_row3.offset(dst_stride0) = scale * s31;
1376    *dst_row3.offset(2 * dst_stride0) = scale * s32;
1377    *dst_row3.offset(3 * dst_stride0) = scale * s33;
1378}
1379
1380#[inline]
1381unsafe fn copy_transpose_scale_2d_identity_tiled_raw<T>(
1382    dst: *mut T,
1383    dst_stride0: isize,
1384    dst_stride1: isize,
1385    src: *const T,
1386    src_stride0: isize,
1387    src_stride1: isize,
1388    src_rows: usize,
1389    src_cols: usize,
1390    scale: T,
1391) where
1392    T: Copy + Mul<Output = T> + MaybeSendSync,
1393{
1394    const TILE: usize = 4;
1395    let row_full = src_rows / TILE * TILE;
1396    let col_full = src_cols / TILE * TILE;
1397
1398    #[cfg(feature = "parallel")]
1399    {
1400        let total = src_rows.saturating_mul(src_cols);
1401        let nthreads = crate::execution_policy::rayon_threads();
1402        if total > MINTHREADLENGTH && nthreads > 1 {
1403            let dst_send = SendPtr(dst);
1404            let src_send = SendPtr(src as *mut T);
1405            let row_tiles = row_full / TILE;
1406            crate::threading::parallel_for_each(0..row_tiles, nthreads, &|tiles| {
1407                for tile_i in tiles {
1408                    let i = tile_i * TILE;
1409                    let dst = dst_send.as_ptr();
1410                    let src = src_send.as_const();
1411                    unsafe {
1412                        let mut j = 0;
1413                        while j < col_full {
1414                            transpose_scale_4x4_identity(
1415                                dst,
1416                                dst_stride0,
1417                                dst_stride1,
1418                                src,
1419                                src_stride0,
1420                                src_stride1,
1421                                i,
1422                                j,
1423                                scale,
1424                            );
1425                            j += TILE;
1426                        }
1427                        for j in col_full..src_cols {
1428                            for ii in i..i + TILE {
1429                                *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1430                                    scale
1431                                        * *src.offset(
1432                                            ii as isize * src_stride0 + j as isize * src_stride1,
1433                                        );
1434                            }
1435                        }
1436                    }
1437                }
1438            });
1439            for i in row_full..src_rows {
1440                for j in 0..src_cols {
1441                    *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1442                        scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1443                }
1444            }
1445            return;
1446        }
1447    }
1448
1449    let mut i = 0;
1450    while i < row_full {
1451        let mut j = 0;
1452        while j < col_full {
1453            transpose_scale_4x4_identity(
1454                dst,
1455                dst_stride0,
1456                dst_stride1,
1457                src,
1458                src_stride0,
1459                src_stride1,
1460                i,
1461                j,
1462                scale,
1463            );
1464            j += TILE;
1465        }
1466        for j in col_full..src_cols {
1467            for ii in i..i + TILE {
1468                *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1469                    scale * *src.offset(ii as isize * src_stride0 + j as isize * src_stride1);
1470            }
1471        }
1472        i += TILE;
1473    }
1474    for i in row_full..src_rows {
1475        for j in 0..src_cols {
1476            *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1477                scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1478        }
1479    }
1480}
1481
1482#[inline]
1483unsafe fn try_copy_transpose_scale_2d_identity_tiled<T>(
1484    dest: &mut StridedViewMut<T>,
1485    src: &StridedView<T>,
1486    scale: T,
1487) -> bool
1488where
1489    T: Copy + Mul<Output = T> + MaybeSendSync,
1490{
1491    if src.ndim() != 2 || dest.ndim() != 2 {
1492        return false;
1493    }
1494    let src_dims = src.dims();
1495    if dest.dims() != [src_dims[1], src_dims[0]] {
1496        return false;
1497    }
1498    if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1499        return false;
1500    }
1501
1502    copy_transpose_scale_2d_identity_tiled_raw(
1503        dest.as_mut_ptr(),
1504        dest.strides()[0],
1505        dest.strides()[1],
1506        src.ptr(),
1507        src.strides()[0],
1508        src.strides()[1],
1509        src_dims[0],
1510        src_dims[1],
1511        scale,
1512    );
1513    true
1514}
1515
1516#[inline]
1517unsafe fn copy_transpose_scale_2d_loop<T>(
1518    dst: *mut T,
1519    dst_stride0: isize,
1520    dst_stride1: isize,
1521    src: *const T,
1522    src_stride0: isize,
1523    src_stride1: isize,
1524    src_rows: usize,
1525    src_cols: usize,
1526    scale: T,
1527) where
1528    T: Copy + ElementOpApply + Mul<Output = T> + MaybeSendSync,
1529{
1530    #[cfg(feature = "parallel")]
1531    {
1532        let total = src_rows.saturating_mul(src_cols);
1533        let nthreads = crate::execution_policy::rayon_threads();
1534        if total > MINTHREADLENGTH && nthreads > 1 {
1535            let dst_send = SendPtr(dst);
1536            let src_send = SendPtr(src as *mut T);
1537            if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1538                crate::threading::parallel_for_each(0..src_rows, nthreads, &|rows| {
1539                    for i in rows {
1540                        let dst = dst_send.as_ptr();
1541                        let src = src_send.as_const();
1542                        unsafe {
1543                            for j in 0..src_cols {
1544                                let value = (*src
1545                                    .offset(i as isize * src_stride0 + j as isize * src_stride1))
1546                                .transpose();
1547                                *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1548                                    scale * value;
1549                            }
1550                        }
1551                    }
1552                });
1553            } else {
1554                crate::threading::parallel_for_each(0..src_cols, nthreads, &|columns| {
1555                    for j in columns {
1556                        let dst = dst_send.as_ptr();
1557                        let src = src_send.as_const();
1558                        unsafe {
1559                            for i in 0..src_rows {
1560                                let value = (*src
1561                                    .offset(i as isize * src_stride0 + j as isize * src_stride1))
1562                                .transpose();
1563                                *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1564                                    scale * value;
1565                            }
1566                        }
1567                    }
1568                });
1569            }
1570            return;
1571        }
1572    }
1573
1574    if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1575        for i in 0..src_rows {
1576            for j in 0..src_cols {
1577                let value =
1578                    (*src.offset(i as isize * src_stride0 + j as isize * src_stride1)).transpose();
1579                *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) = scale * value;
1580            }
1581        }
1582    } else {
1583        for j in 0..src_cols {
1584            for i in 0..src_rows {
1585                let value =
1586                    (*src.offset(i as isize * src_stride0 + j as isize * src_stride1)).transpose();
1587                *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) = scale * value;
1588            }
1589        }
1590    }
1591}
1592
1593/// Copy with transpose and scaling: `dest[j,i] = scale * src[i,j]`.
1594pub fn copy_transpose_scale_into<T>(
1595    dest: &mut StridedViewMut<T>,
1596    src: &StridedView<T>,
1597    scale: T,
1598) -> Result<()>
1599where
1600    T: Copy + ElementOpApply + Mul<Output = T> + Zero + One + PartialEq + MaybeSendSync + 'static,
1601{
1602    if src.ndim() != 2 || dest.ndim() != 2 {
1603        return Err(StridedError::RankMismatch(src.ndim(), 2));
1604    }
1605
1606    let src_dims = src.dims();
1607    let expected_dims = [src_dims[1], src_dims[0]];
1608    ensure_same_shape(dest.dims(), &expected_dims)?;
1609    // Reject an element count beyond usize (a huge stride-0 broadcast) up
1610    // front instead of looping over it.
1611    total_len(src_dims)?;
1612
1613    if scale == T::zero() {
1614        unsafe {
1615            if same_contiguous_layout(dest.dims(), &[dest.strides()]).is_some() {
1616                fill_contiguous(dest.as_mut_ptr(), total_len(dest.dims())?, T::zero());
1617            } else {
1618                fill_2d(
1619                    dest.as_mut_ptr(),
1620                    dest.dims()[0],
1621                    dest.dims()[1],
1622                    dest.strides()[0],
1623                    dest.strides()[1],
1624                    T::zero(),
1625                );
1626            }
1627        }
1628        return Ok(());
1629    }
1630
1631    let transpose_is_identity = element_transpose_is_identity::<T>();
1632
1633    unsafe {
1634        if transpose_is_identity && try_copy_transpose_scale_2d_f64_tiled_typed(dest, src, scale) {
1635            return Ok(());
1636        }
1637        if transpose_is_identity && try_copy_transpose_scale_2d_identity_tiled(dest, src, scale) {
1638            return Ok(());
1639        }
1640    }
1641
1642    if scale == T::one() && transpose_is_identity {
1643        let src_t = src.permute(&[1, 0])?;
1644        #[cfg(feature = "parallel")]
1645        {
1646            return crate::threading::copy_permuted_with_active_policy(dest, &src_t);
1647        }
1648        #[cfg(not(feature = "parallel"))]
1649        return crate::threading::copy_permuted_serial(dest, &src_t);
1650    }
1651
1652    unsafe {
1653        copy_transpose_scale_2d_loop(
1654            dest.as_mut_ptr(),
1655            dest.strides()[0],
1656            dest.strides()[1],
1657            src.ptr(),
1658            src.strides()[0],
1659            src.strides()[1],
1660            src_dims[0],
1661            src_dims[1],
1662            scale,
1663        );
1664    }
1665    Ok(())
1666}
1667
1668#[cfg(test)]
1669#[path = "ops_view/tests/tiled_tests.rs"]
1670mod tiled_tests;