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