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: usize = fused_dims.iter().product();
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: usize = fused_dims.iter().product();
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: usize = fused_dims.iter().product();
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: usize = fused_dims.iter().product();
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
806    if same_contiguous_layout(a_dims, &[a_strides, b_strides]).is_some() {
807        let len = total_len(a_dims);
808
809        // SIMD fast path: both contiguous, both Identity ops, same type
810        if OpA::IS_IDENTITY
811            && OpB::IS_IDENTITY
812            && std::any::TypeId::of::<A>() == std::any::TypeId::of::<R>()
813            && std::any::TypeId::of::<B>() == std::any::TypeId::of::<R>()
814        {
815            let sa = unsafe { std::slice::from_raw_parts(a_ptr as *const R, len) };
816            let sb = unsafe { std::slice::from_raw_parts(b_ptr as *const R, len) };
817            if let Some(result) = R::try_simd_dot(sa, sb) {
818                return Ok(result);
819            }
820        }
821
822        // Generic contiguous fast path
823        let sa = unsafe { std::slice::from_raw_parts(a_ptr, len) };
824        let sb = unsafe { std::slice::from_raw_parts(b_ptr, len) };
825        let mut acc = R::zero();
826        simd::dispatch_if_large(len, || {
827            for i in 0..len {
828                acc = acc + OpA::apply(sa[i]) * OpB::apply(sb[i]);
829            }
830        });
831        return Ok(acc);
832    }
833
834    let strides_list: [&[isize]; 2] = [a_strides, b_strides];
835    let elem_size = std::mem::size_of::<A>()
836        .max(std::mem::size_of::<B>())
837        .max(std::mem::size_of::<R>());
838
839    let (fused_dims, ordered_strides, plan) =
840        build_plan_fused(a_dims, &strides_list, None, elem_size);
841
842    let mut acc = R::zero();
843    let initial_offsets = vec![0isize; ordered_strides.len()];
844    for_each_inner_block_preordered(
845        &fused_dims,
846        &plan.block,
847        &ordered_strides,
848        &initial_offsets,
849        |offsets, len, strides| {
850            acc = unsafe {
851                inner_loop_dot::<A, B, R, OpA, OpB>(
852                    a_ptr.offset(offsets[0]),
853                    strides[0],
854                    b_ptr.offset(offsets[1]),
855                    strides[1],
856                    len,
857                    acc,
858                )
859            };
860            Ok(())
861        },
862    )?;
863
864    Ok(acc)
865}
866
867/// Symmetrize a square matrix: `dest = (src + src^T) / 2`.
868pub fn symmetrize_into<T>(dest: &mut StridedViewMut<T>, src: &StridedView<T>) -> Result<()>
869where
870    T: Copy
871        + Add<Output = T>
872        + Mul<Output = T>
873        + num_traits::FromPrimitive
874        + std::ops::Div<Output = T>
875        + MaybeSendSync,
876{
877    if src.ndim() != 2 {
878        return Err(StridedError::RankMismatch(src.ndim(), 2));
879    }
880    let rows = src.dims()[0];
881    let cols = src.dims()[1];
882    if rows != cols {
883        return Err(StridedError::NonSquare { rows, cols });
884    }
885
886    let src_t = src.permute(&[1, 0])?;
887    let half = T::from_f64(0.5).ok_or(StridedError::ScalarConversion)?;
888
889    zip_map2_into(dest, src, &src_t, |a, b| (a + b) * half)
890}
891
892/// Conjugate-symmetrize a square matrix: `dest = (src + conj(src^T)) / 2`.
893pub fn symmetrize_conj_into<T>(dest: &mut StridedViewMut<T>, src: &StridedView<T>) -> Result<()>
894where
895    T: Copy
896        + ElementOpApply
897        + Add<Output = T>
898        + Mul<Output = T>
899        + num_traits::FromPrimitive
900        + std::ops::Div<Output = T>
901        + MaybeSendSync,
902{
903    if src.ndim() != 2 {
904        return Err(StridedError::RankMismatch(src.ndim(), 2));
905    }
906    let rows = src.dims()[0];
907    let cols = src.dims()[1];
908    if rows != cols {
909        return Err(StridedError::NonSquare { rows, cols });
910    }
911
912    // adjoint = conj + transpose
913    let src_adj = src.adjoint_2d()?;
914    let half = T::from_f64(0.5).ok_or(StridedError::ScalarConversion)?;
915
916    zip_map2_into(dest, src, &src_adj, |a, b| (a + b) * half)
917}
918
919/// Copy with scaling: `dest[i] = scale * src[i]`.
920///
921/// Scale, source, and destination may have different element types.
922pub fn copy_scale<D, S, A, Op>(
923    dest: &mut StridedViewMut<D>,
924    src: &StridedView<S, Op>,
925    scale: A,
926) -> Result<()>
927where
928    A: Copy + Mul<S, Output = D> + MaybeSync,
929    D: Copy + MaybeSendSync,
930    S: Copy + MaybeSendSync,
931    Op: ElementOp<S>,
932{
933    map_into(dest, src, |x| scale * x)
934}
935
936/// Copy with complex conjugation: `dest[i] = conj(src[i])`.
937pub fn copy_conj<T: Copy + ElementOpApply + MaybeSendSync>(
938    dest: &mut StridedViewMut<T>,
939    src: &StridedView<T>,
940) -> Result<()> {
941    let src_conj = src.conj();
942    copy_into(dest, &src_conj)
943}
944
945#[inline]
946fn element_transpose_is_identity<T: 'static>() -> bool {
947    use std::any::TypeId;
948
949    macro_rules! matches_type {
950        ($($ty:ty),* $(,)?) => {{
951            let id = TypeId::of::<T>();
952            false $(|| id == TypeId::of::<$ty>())*
953        }};
954    }
955
956    matches_type!(
957        f32,
958        f64,
959        i8,
960        i16,
961        i32,
962        i64,
963        i128,
964        isize,
965        u8,
966        u16,
967        u32,
968        u64,
969        u128,
970        usize,
971        num_complex::Complex32,
972        num_complex::Complex64,
973    )
974}
975
976#[inline]
977fn element_zero_is_all_bits_zero<T: 'static>() -> bool {
978    use std::any::TypeId;
979
980    macro_rules! matches_type {
981        ($($ty:ty),* $(,)?) => {{
982            let id = TypeId::of::<T>();
983            false $(|| id == TypeId::of::<$ty>())*
984        }};
985    }
986
987    matches_type!(
988        f32,
989        f64,
990        i8,
991        i16,
992        i32,
993        i64,
994        i128,
995        isize,
996        u8,
997        u16,
998        u32,
999        u64,
1000        u128,
1001        usize,
1002        num_complex::Complex32,
1003        num_complex::Complex64,
1004    )
1005}
1006
1007#[inline]
1008unsafe fn fill_2d<T: Copy + MaybeSendSync>(
1009    dst: *mut T,
1010    dim0: usize,
1011    dim1: usize,
1012    dst_stride0: isize,
1013    dst_stride1: isize,
1014    value: T,
1015) {
1016    #[cfg(feature = "parallel")]
1017    {
1018        let total = dim0.saturating_mul(dim1);
1019        let nthreads = crate::execution_policy::rayon_threads();
1020        if total > MINTHREADLENGTH && nthreads > 1 {
1021            let dst_send = SendPtr(dst);
1022            if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1023                crate::threading::parallel_for_each(0..dim1, nthreads, &|columns| {
1024                    for j in columns {
1025                        let dst = dst_send.as_ptr();
1026                        unsafe {
1027                            let base = j as isize * dst_stride1;
1028                            for i in 0..dim0 {
1029                                *dst.offset(base + i as isize * dst_stride0) = value;
1030                            }
1031                        }
1032                    }
1033                });
1034            } else {
1035                crate::threading::parallel_for_each(0..dim0, nthreads, &|rows| {
1036                    for i in rows {
1037                        let dst = dst_send.as_ptr();
1038                        unsafe {
1039                            let base = i as isize * dst_stride0;
1040                            for j in 0..dim1 {
1041                                *dst.offset(base + j as isize * dst_stride1) = value;
1042                            }
1043                        }
1044                    }
1045                });
1046            }
1047            return;
1048        }
1049    }
1050
1051    if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1052        for j in 0..dim1 {
1053            let base = j as isize * dst_stride1;
1054            for i in 0..dim0 {
1055                *dst.offset(base + i as isize * dst_stride0) = value;
1056            }
1057        }
1058    } else {
1059        for i in 0..dim0 {
1060            let base = i as isize * dst_stride0;
1061            for j in 0..dim1 {
1062                *dst.offset(base + j as isize * dst_stride1) = value;
1063            }
1064        }
1065    }
1066}
1067
1068#[inline]
1069unsafe fn fill_contiguous<T>(dst: *mut T, len: usize, value: T)
1070where
1071    T: Copy + Zero + PartialEq + MaybeSendSync + 'static,
1072{
1073    if element_zero_is_all_bits_zero::<T>() && value == T::zero() {
1074        std::ptr::write_bytes(dst, 0, len);
1075        return;
1076    }
1077
1078    let dst = std::slice::from_raw_parts_mut(dst, len);
1079    dst.fill(value);
1080}
1081
1082#[inline(always)]
1083unsafe fn transpose_scale_4x4_f64(
1084    dst: *mut f64,
1085    dst_stride0: isize,
1086    dst_stride1: isize,
1087    src: *const f64,
1088    src_stride0: isize,
1089    src_stride1: isize,
1090    i: usize,
1091    j: usize,
1092    scale: f64,
1093) {
1094    let src_base = src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1095
1096    let s00 = *src_base;
1097    let s10 = *src_base.offset(src_stride0);
1098    let s20 = *src_base.offset(2 * src_stride0);
1099    let s30 = *src_base.offset(3 * src_stride0);
1100
1101    let src_col1 = src_base.offset(src_stride1);
1102    let s01 = *src_col1;
1103    let s11 = *src_col1.offset(src_stride0);
1104    let s21 = *src_col1.offset(2 * src_stride0);
1105    let s31 = *src_col1.offset(3 * src_stride0);
1106
1107    let src_col2 = src_base.offset(2 * src_stride1);
1108    let s02 = *src_col2;
1109    let s12 = *src_col2.offset(src_stride0);
1110    let s22 = *src_col2.offset(2 * src_stride0);
1111    let s32 = *src_col2.offset(3 * src_stride0);
1112
1113    let src_col3 = src_base.offset(3 * src_stride1);
1114    let s03 = *src_col3;
1115    let s13 = *src_col3.offset(src_stride0);
1116    let s23 = *src_col3.offset(2 * src_stride0);
1117    let s33 = *src_col3.offset(3 * src_stride0);
1118
1119    let dst_row0 = dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1);
1120    *dst_row0 = scale * s00;
1121    *dst_row0.offset(dst_stride0) = scale * s01;
1122    *dst_row0.offset(2 * dst_stride0) = scale * s02;
1123    *dst_row0.offset(3 * dst_stride0) = scale * s03;
1124
1125    let dst_row1 = dst_row0.offset(dst_stride1);
1126    *dst_row1 = scale * s10;
1127    *dst_row1.offset(dst_stride0) = scale * s11;
1128    *dst_row1.offset(2 * dst_stride0) = scale * s12;
1129    *dst_row1.offset(3 * dst_stride0) = scale * s13;
1130
1131    let dst_row2 = dst_row0.offset(2 * dst_stride1);
1132    *dst_row2 = scale * s20;
1133    *dst_row2.offset(dst_stride0) = scale * s21;
1134    *dst_row2.offset(2 * dst_stride0) = scale * s22;
1135    *dst_row2.offset(3 * dst_stride0) = scale * s23;
1136
1137    let dst_row3 = dst_row0.offset(3 * dst_stride1);
1138    *dst_row3 = scale * s30;
1139    *dst_row3.offset(dst_stride0) = scale * s31;
1140    *dst_row3.offset(2 * dst_stride0) = scale * s32;
1141    *dst_row3.offset(3 * dst_stride0) = scale * s33;
1142}
1143
1144#[inline]
1145unsafe fn copy_transpose_scale_2d_f64_tiled_raw(
1146    dst: *mut f64,
1147    dst_stride0: isize,
1148    dst_stride1: isize,
1149    src: *const f64,
1150    src_stride0: isize,
1151    src_stride1: isize,
1152    src_rows: usize,
1153    src_cols: usize,
1154    scale: f64,
1155) {
1156    const TILE: usize = 4;
1157    let row_full = src_rows / TILE * TILE;
1158    let col_full = src_cols / TILE * TILE;
1159
1160    #[cfg(feature = "parallel")]
1161    {
1162        let total = src_rows.saturating_mul(src_cols);
1163        let nthreads = crate::execution_policy::rayon_threads();
1164        if total > MINTHREADLENGTH && nthreads > 1 {
1165            let dst_send = SendPtr(dst);
1166            let src_send = SendPtr(src as *mut f64);
1167            let row_tiles = row_full / TILE;
1168            crate::threading::parallel_for_each(0..row_tiles, nthreads, &|tiles| {
1169                for tile_i in tiles {
1170                    let i = tile_i * TILE;
1171                    let dst = dst_send.as_ptr();
1172                    let src = src_send.as_const();
1173                    unsafe {
1174                        let mut j = 0;
1175                        while j < col_full {
1176                            transpose_scale_4x4_f64(
1177                                dst,
1178                                dst_stride0,
1179                                dst_stride1,
1180                                src,
1181                                src_stride0,
1182                                src_stride1,
1183                                i,
1184                                j,
1185                                scale,
1186                            );
1187                            j += TILE;
1188                        }
1189                        for j in col_full..src_cols {
1190                            for ii in i..i + TILE {
1191                                *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1192                                    scale
1193                                        * *src.offset(
1194                                            ii as isize * src_stride0 + j as isize * src_stride1,
1195                                        );
1196                            }
1197                        }
1198                    }
1199                }
1200            });
1201            for i in row_full..src_rows {
1202                for j in 0..src_cols {
1203                    *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1204                        scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1205                }
1206            }
1207            return;
1208        }
1209    }
1210
1211    let mut i = 0;
1212    while i < row_full {
1213        let mut j = 0;
1214        while j < col_full {
1215            transpose_scale_4x4_f64(
1216                dst,
1217                dst_stride0,
1218                dst_stride1,
1219                src,
1220                src_stride0,
1221                src_stride1,
1222                i,
1223                j,
1224                scale,
1225            );
1226            j += TILE;
1227        }
1228        for j in col_full..src_cols {
1229            for ii in i..i + TILE {
1230                *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1231                    scale * *src.offset(ii as isize * src_stride0 + j as isize * src_stride1);
1232            }
1233        }
1234        i += TILE;
1235    }
1236    for i in row_full..src_rows {
1237        for j in 0..src_cols {
1238            *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1239                scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1240        }
1241    }
1242}
1243
1244#[inline]
1245#[cfg(test)]
1246unsafe fn try_copy_transpose_scale_2d_f64_tiled(
1247    dest: &mut StridedViewMut<f64>,
1248    src: &StridedView<f64>,
1249    scale: f64,
1250) -> bool {
1251    if src.ndim() != 2 || dest.ndim() != 2 {
1252        return false;
1253    }
1254    let src_dims = src.dims();
1255    if dest.dims() != [src_dims[1], src_dims[0]] {
1256        return false;
1257    }
1258    if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1259        return false;
1260    }
1261
1262    copy_transpose_scale_2d_f64_tiled_raw(
1263        dest.as_mut_ptr(),
1264        dest.strides()[0],
1265        dest.strides()[1],
1266        src.ptr(),
1267        src.strides()[0],
1268        src.strides()[1],
1269        src_dims[0],
1270        src_dims[1],
1271        scale,
1272    );
1273    true
1274}
1275
1276#[inline]
1277unsafe fn try_copy_transpose_scale_2d_f64_tiled_typed<T>(
1278    dest: &mut StridedViewMut<T>,
1279    src: &StridedView<T>,
1280    scale: T,
1281) -> bool
1282where
1283    T: Copy + 'static,
1284{
1285    if std::any::TypeId::of::<T>() != std::any::TypeId::of::<f64>() {
1286        return false;
1287    }
1288
1289    let scale = *(&scale as *const T).cast::<f64>();
1290    if src.ndim() != 2 || dest.ndim() != 2 {
1291        return false;
1292    }
1293    let src_dims = src.dims();
1294    if dest.dims() != [src_dims[1], src_dims[0]] {
1295        return false;
1296    }
1297    if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1298        return false;
1299    }
1300
1301    copy_transpose_scale_2d_f64_tiled_raw(
1302        dest.as_mut_ptr().cast::<f64>(),
1303        dest.strides()[0],
1304        dest.strides()[1],
1305        src.ptr().cast::<f64>(),
1306        src.strides()[0],
1307        src.strides()[1],
1308        src_dims[0],
1309        src_dims[1],
1310        scale,
1311    );
1312    true
1313}
1314
1315#[inline(always)]
1316unsafe fn transpose_scale_4x4_identity<T>(
1317    dst: *mut T,
1318    dst_stride0: isize,
1319    dst_stride1: isize,
1320    src: *const T,
1321    src_stride0: isize,
1322    src_stride1: isize,
1323    i: usize,
1324    j: usize,
1325    scale: T,
1326) where
1327    T: Copy + Mul<Output = T>,
1328{
1329    let src_base = src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1330
1331    let s00 = *src_base;
1332    let s10 = *src_base.offset(src_stride0);
1333    let s20 = *src_base.offset(2 * src_stride0);
1334    let s30 = *src_base.offset(3 * src_stride0);
1335
1336    let src_col1 = src_base.offset(src_stride1);
1337    let s01 = *src_col1;
1338    let s11 = *src_col1.offset(src_stride0);
1339    let s21 = *src_col1.offset(2 * src_stride0);
1340    let s31 = *src_col1.offset(3 * src_stride0);
1341
1342    let src_col2 = src_base.offset(2 * src_stride1);
1343    let s02 = *src_col2;
1344    let s12 = *src_col2.offset(src_stride0);
1345    let s22 = *src_col2.offset(2 * src_stride0);
1346    let s32 = *src_col2.offset(3 * src_stride0);
1347
1348    let src_col3 = src_base.offset(3 * src_stride1);
1349    let s03 = *src_col3;
1350    let s13 = *src_col3.offset(src_stride0);
1351    let s23 = *src_col3.offset(2 * src_stride0);
1352    let s33 = *src_col3.offset(3 * src_stride0);
1353
1354    let dst_row0 = dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1);
1355    *dst_row0 = scale * s00;
1356    *dst_row0.offset(dst_stride0) = scale * s01;
1357    *dst_row0.offset(2 * dst_stride0) = scale * s02;
1358    *dst_row0.offset(3 * dst_stride0) = scale * s03;
1359
1360    let dst_row1 = dst_row0.offset(dst_stride1);
1361    *dst_row1 = scale * s10;
1362    *dst_row1.offset(dst_stride0) = scale * s11;
1363    *dst_row1.offset(2 * dst_stride0) = scale * s12;
1364    *dst_row1.offset(3 * dst_stride0) = scale * s13;
1365
1366    let dst_row2 = dst_row0.offset(2 * dst_stride1);
1367    *dst_row2 = scale * s20;
1368    *dst_row2.offset(dst_stride0) = scale * s21;
1369    *dst_row2.offset(2 * dst_stride0) = scale * s22;
1370    *dst_row2.offset(3 * dst_stride0) = scale * s23;
1371
1372    let dst_row3 = dst_row0.offset(3 * dst_stride1);
1373    *dst_row3 = scale * s30;
1374    *dst_row3.offset(dst_stride0) = scale * s31;
1375    *dst_row3.offset(2 * dst_stride0) = scale * s32;
1376    *dst_row3.offset(3 * dst_stride0) = scale * s33;
1377}
1378
1379#[inline]
1380unsafe fn copy_transpose_scale_2d_identity_tiled_raw<T>(
1381    dst: *mut T,
1382    dst_stride0: isize,
1383    dst_stride1: isize,
1384    src: *const T,
1385    src_stride0: isize,
1386    src_stride1: isize,
1387    src_rows: usize,
1388    src_cols: usize,
1389    scale: T,
1390) where
1391    T: Copy + Mul<Output = T> + MaybeSendSync,
1392{
1393    const TILE: usize = 4;
1394    let row_full = src_rows / TILE * TILE;
1395    let col_full = src_cols / TILE * TILE;
1396
1397    #[cfg(feature = "parallel")]
1398    {
1399        let total = src_rows.saturating_mul(src_cols);
1400        let nthreads = crate::execution_policy::rayon_threads();
1401        if total > MINTHREADLENGTH && nthreads > 1 {
1402            let dst_send = SendPtr(dst);
1403            let src_send = SendPtr(src as *mut T);
1404            let row_tiles = row_full / TILE;
1405            crate::threading::parallel_for_each(0..row_tiles, nthreads, &|tiles| {
1406                for tile_i in tiles {
1407                    let i = tile_i * TILE;
1408                    let dst = dst_send.as_ptr();
1409                    let src = src_send.as_const();
1410                    unsafe {
1411                        let mut j = 0;
1412                        while j < col_full {
1413                            transpose_scale_4x4_identity(
1414                                dst,
1415                                dst_stride0,
1416                                dst_stride1,
1417                                src,
1418                                src_stride0,
1419                                src_stride1,
1420                                i,
1421                                j,
1422                                scale,
1423                            );
1424                            j += TILE;
1425                        }
1426                        for j in col_full..src_cols {
1427                            for ii in i..i + TILE {
1428                                *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1429                                    scale
1430                                        * *src.offset(
1431                                            ii as isize * src_stride0 + j as isize * src_stride1,
1432                                        );
1433                            }
1434                        }
1435                    }
1436                }
1437            });
1438            for i in row_full..src_rows {
1439                for j in 0..src_cols {
1440                    *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1441                        scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1442                }
1443            }
1444            return;
1445        }
1446    }
1447
1448    let mut i = 0;
1449    while i < row_full {
1450        let mut j = 0;
1451        while j < col_full {
1452            transpose_scale_4x4_identity(
1453                dst,
1454                dst_stride0,
1455                dst_stride1,
1456                src,
1457                src_stride0,
1458                src_stride1,
1459                i,
1460                j,
1461                scale,
1462            );
1463            j += TILE;
1464        }
1465        for j in col_full..src_cols {
1466            for ii in i..i + TILE {
1467                *dst.offset(j as isize * dst_stride0 + ii as isize * dst_stride1) =
1468                    scale * *src.offset(ii as isize * src_stride0 + j as isize * src_stride1);
1469            }
1470        }
1471        i += TILE;
1472    }
1473    for i in row_full..src_rows {
1474        for j in 0..src_cols {
1475            *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1476                scale * *src.offset(i as isize * src_stride0 + j as isize * src_stride1);
1477        }
1478    }
1479}
1480
1481#[inline]
1482unsafe fn try_copy_transpose_scale_2d_identity_tiled<T>(
1483    dest: &mut StridedViewMut<T>,
1484    src: &StridedView<T>,
1485    scale: T,
1486) -> bool
1487where
1488    T: Copy + Mul<Output = T> + MaybeSendSync,
1489{
1490    if src.ndim() != 2 || dest.ndim() != 2 {
1491        return false;
1492    }
1493    let src_dims = src.dims();
1494    if dest.dims() != [src_dims[1], src_dims[0]] {
1495        return false;
1496    }
1497    if src.strides()[0] != 1 || dest.strides()[0] != 1 {
1498        return false;
1499    }
1500
1501    copy_transpose_scale_2d_identity_tiled_raw(
1502        dest.as_mut_ptr(),
1503        dest.strides()[0],
1504        dest.strides()[1],
1505        src.ptr(),
1506        src.strides()[0],
1507        src.strides()[1],
1508        src_dims[0],
1509        src_dims[1],
1510        scale,
1511    );
1512    true
1513}
1514
1515#[inline]
1516unsafe fn copy_transpose_scale_2d_loop<T>(
1517    dst: *mut T,
1518    dst_stride0: isize,
1519    dst_stride1: isize,
1520    src: *const T,
1521    src_stride0: isize,
1522    src_stride1: isize,
1523    src_rows: usize,
1524    src_cols: usize,
1525    scale: T,
1526) where
1527    T: Copy + ElementOpApply + Mul<Output = T> + MaybeSendSync,
1528{
1529    #[cfg(feature = "parallel")]
1530    {
1531        let total = src_rows.saturating_mul(src_cols);
1532        let nthreads = crate::execution_policy::rayon_threads();
1533        if total > MINTHREADLENGTH && nthreads > 1 {
1534            let dst_send = SendPtr(dst);
1535            let src_send = SendPtr(src as *mut T);
1536            if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1537                crate::threading::parallel_for_each(0..src_rows, nthreads, &|rows| {
1538                    for i in rows {
1539                        let dst = dst_send.as_ptr();
1540                        let src = src_send.as_const();
1541                        unsafe {
1542                            for j in 0..src_cols {
1543                                let value = (*src
1544                                    .offset(i as isize * src_stride0 + j as isize * src_stride1))
1545                                .transpose();
1546                                *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1547                                    scale * value;
1548                            }
1549                        }
1550                    }
1551                });
1552            } else {
1553                crate::threading::parallel_for_each(0..src_cols, nthreads, &|columns| {
1554                    for j in columns {
1555                        let dst = dst_send.as_ptr();
1556                        let src = src_send.as_const();
1557                        unsafe {
1558                            for i in 0..src_rows {
1559                                let value = (*src
1560                                    .offset(i as isize * src_stride0 + j as isize * src_stride1))
1561                                .transpose();
1562                                *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) =
1563                                    scale * value;
1564                            }
1565                        }
1566                    }
1567                });
1568            }
1569            return;
1570        }
1571    }
1572
1573    if dst_stride0.unsigned_abs() <= dst_stride1.unsigned_abs() {
1574        for i in 0..src_rows {
1575            for j in 0..src_cols {
1576                let value =
1577                    (*src.offset(i as isize * src_stride0 + j as isize * src_stride1)).transpose();
1578                *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) = scale * value;
1579            }
1580        }
1581    } else {
1582        for j in 0..src_cols {
1583            for i in 0..src_rows {
1584                let value =
1585                    (*src.offset(i as isize * src_stride0 + j as isize * src_stride1)).transpose();
1586                *dst.offset(j as isize * dst_stride0 + i as isize * dst_stride1) = scale * value;
1587            }
1588        }
1589    }
1590}
1591
1592/// Copy with transpose and scaling: `dest[j,i] = scale * src[i,j]`.
1593pub fn copy_transpose_scale_into<T>(
1594    dest: &mut StridedViewMut<T>,
1595    src: &StridedView<T>,
1596    scale: T,
1597) -> Result<()>
1598where
1599    T: Copy + ElementOpApply + Mul<Output = T> + Zero + One + PartialEq + MaybeSendSync + 'static,
1600{
1601    if src.ndim() != 2 || dest.ndim() != 2 {
1602        return Err(StridedError::RankMismatch(src.ndim(), 2));
1603    }
1604
1605    let src_dims = src.dims();
1606    let expected_dims = [src_dims[1], src_dims[0]];
1607    ensure_same_shape(dest.dims(), &expected_dims)?;
1608
1609    if scale == T::zero() {
1610        unsafe {
1611            if same_contiguous_layout(dest.dims(), &[dest.strides()]).is_some() {
1612                fill_contiguous(dest.as_mut_ptr(), total_len(dest.dims()), T::zero());
1613            } else {
1614                fill_2d(
1615                    dest.as_mut_ptr(),
1616                    dest.dims()[0],
1617                    dest.dims()[1],
1618                    dest.strides()[0],
1619                    dest.strides()[1],
1620                    T::zero(),
1621                );
1622            }
1623        }
1624        return Ok(());
1625    }
1626
1627    let transpose_is_identity = element_transpose_is_identity::<T>();
1628
1629    unsafe {
1630        if transpose_is_identity && try_copy_transpose_scale_2d_f64_tiled_typed(dest, src, scale) {
1631            return Ok(());
1632        }
1633        if transpose_is_identity && try_copy_transpose_scale_2d_identity_tiled(dest, src, scale) {
1634            return Ok(());
1635        }
1636    }
1637
1638    if scale == T::one() && transpose_is_identity {
1639        let src_t = src.permute(&[1, 0])?;
1640        #[cfg(feature = "parallel")]
1641        {
1642            return crate::threading::copy_permuted_with_active_policy(dest, &src_t);
1643        }
1644        #[cfg(not(feature = "parallel"))]
1645        return crate::threading::copy_permuted_serial(dest, &src_t);
1646    }
1647
1648    unsafe {
1649        copy_transpose_scale_2d_loop(
1650            dest.as_mut_ptr(),
1651            dest.strides()[0],
1652            dest.strides()[1],
1653            src.ptr(),
1654            src.strides()[0],
1655            src.strides()[1],
1656            src_dims[0],
1657            src_dims[1],
1658            scale,
1659        );
1660    }
1661    Ok(())
1662}
1663
1664#[cfg(test)]
1665#[path = "ops_view/tests/tiled_tests.rs"]
1666mod tiled_tests;