Skip to main content

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