1use super::{erased_raw_ref, map_output_dtype, raw_any, OneShotScalar};
11use crate::*;
12use core::mem::MaybeUninit;
13use core::num::Wrapping;
14use num_complex::{Complex32, Complex64};
15use strided_basic::execution::{
16 check_dtype, ensure_same_shape, map_raw_into_validated,
17 validate_destination_layout_without_alloc, validate_uninit_no_overlap,
18 zip_map2_raw_into_validated, ValidatedDestinationLayout,
19};
20
21pub fn erased_map_into_uninit(
58 input_dtype: KernelDType,
59 op: ErasedMapOp,
60 ctx: &ExecContext,
61 dest: &mut ErasedRawStridedUninitMut<'_>,
62 input: &ErasedRawStridedPtr<'_>,
63) -> Result<()> {
64 check_dtype(input_dtype, input.dtype())?;
65 check_dtype(map_output_dtype(input_dtype, op)?, dest.dtype())?;
66 validate_uninit_no_overlap(dest, input, 0)?;
67 let input = unsafe { input.try_as_ref_after_no_overlap() }?;
69
70 ctx.run(|| match (input_dtype, op) {
71 (KernelDType::C32, ErasedMapOp::Abs) => {
72 uninit_map_with::<f32, Complex32>(dest, &input, |value| value.norm())
73 }
74 (KernelDType::C64, ErasedMapOp::Abs) => {
75 uninit_map_with::<f64, Complex64>(dest, &input, |value| value.norm())
76 }
77 (KernelDType::F32, _) => uninit_map::<f32>(op, dest, &input),
78 (KernelDType::F64, _) => uninit_map::<f64>(op, dest, &input),
79 (KernelDType::I32, _) => uninit_map::<i32>(op, dest, &input),
80 (KernelDType::I64, _) => uninit_map::<i64>(op, dest, &input),
81 (KernelDType::Bool, _) => uninit_map::<bool>(op, dest, &input),
82 (KernelDType::C32, _) => uninit_map::<Complex32>(op, dest, &input),
83 (KernelDType::C64, _) => uninit_map::<Complex64>(op, dest, &input),
84 _ => Err(StridedError::UnsupportedDType {
85 dtype: input_dtype.label(),
86 }),
87 })
88}
89
90pub fn erased_zip_into_uninit(
131 dtype: KernelDType,
132 op: ErasedZipOp,
133 ctx: &ExecContext,
134 dest: &mut ErasedRawStridedUninitMut<'_>,
135 lhs: &ErasedRawStridedPtr<'_>,
136 rhs: &ErasedRawStridedPtr<'_>,
137) -> Result<()> {
138 check_dtype(dtype, dest.dtype())?;
139 check_dtype(dtype, lhs.dtype())?;
140 check_dtype(dtype, rhs.dtype())?;
141 validate_uninit_no_overlap(dest, lhs, 0)?;
142 validate_uninit_no_overlap(dest, rhs, 1)?;
143 let lhs = unsafe { lhs.try_as_ref_after_no_overlap() }?;
145 let rhs = unsafe { rhs.try_as_ref_after_no_overlap() }?;
147
148 ctx.run(|| match dtype {
149 KernelDType::F32 => uninit_zip::<f32>(op, dest, &lhs, &rhs),
150 KernelDType::F64 => uninit_zip::<f64>(op, dest, &lhs, &rhs),
151 KernelDType::I32 => uninit_zip::<i32>(op, dest, &lhs, &rhs),
152 KernelDType::I64 => uninit_zip::<i64>(op, dest, &lhs, &rhs),
153 KernelDType::Bool => uninit_zip::<bool>(op, dest, &lhs, &rhs),
154 KernelDType::C32 => uninit_zip::<Complex32>(op, dest, &lhs, &rhs),
155 KernelDType::C64 => uninit_zip::<Complex64>(op, dest, &lhs, &rhs),
156 _ => Err(StridedError::UnsupportedDType {
157 dtype: dtype.label(),
158 }),
159 })
160}
161
162pub fn erased_compare_into_uninit(
203 dtype: KernelDType,
204 op: CompareOp,
205 ctx: &ExecContext,
206 dest: &mut ErasedRawStridedUninitMut<'_>,
207 lhs: &ErasedRawStridedPtr<'_>,
208 rhs: &ErasedRawStridedPtr<'_>,
209) -> Result<()> {
210 check_dtype(KernelDType::Bool, dest.dtype())?;
211 check_dtype(dtype, lhs.dtype())?;
212 check_dtype(dtype, rhs.dtype())?;
213 validate_uninit_no_overlap(dest, lhs, 0)?;
214 validate_uninit_no_overlap(dest, rhs, 1)?;
215 let lhs = unsafe { lhs.try_as_ref_after_no_overlap() }?;
217 let rhs = unsafe { rhs.try_as_ref_after_no_overlap() }?;
219
220 ctx.run(|| match dtype {
221 KernelDType::F32 => uninit_compare_ordered::<f32>(op, dest, &lhs, &rhs),
222 KernelDType::F64 => uninit_compare_ordered::<f64>(op, dest, &lhs, &rhs),
223 KernelDType::I32 => uninit_compare_ordered::<i32>(op, dest, &lhs, &rhs),
224 KernelDType::I64 => uninit_compare_ordered::<i64>(op, dest, &lhs, &rhs),
225 KernelDType::Bool => uninit_compare_ordered::<bool>(op, dest, &lhs, &rhs),
226 KernelDType::C32 => uninit_compare_eq::<Complex32>(op, dest, &lhs, &rhs),
227 KernelDType::C64 => uninit_compare_eq::<Complex64>(op, dest, &lhs, &rhs),
228 _ => Err(StridedError::UnsupportedDType {
229 dtype: dtype.label(),
230 }),
231 })
232}
233
234pub fn erased_select_into_uninit(
275 dtype: KernelDType,
276 ctx: &ExecContext,
277 dest: &mut ErasedRawStridedUninitMut<'_>,
278 pred: &ErasedRawStridedPtr<'_>,
279 on_true: &ErasedRawStridedPtr<'_>,
280 on_false: &ErasedRawStridedPtr<'_>,
281) -> Result<()> {
282 check_dtype(dtype, dest.dtype())?;
283 check_dtype(KernelDType::Bool, pred.dtype())?;
284 check_dtype(dtype, on_true.dtype())?;
285 check_dtype(dtype, on_false.dtype())?;
286 validate_uninit_no_overlap(dest, pred, 0)?;
287 validate_uninit_no_overlap(dest, on_true, 1)?;
288 validate_uninit_no_overlap(dest, on_false, 2)?;
289 let pred = unsafe { pred.try_as_ref_after_no_overlap() }?;
291 let on_true = unsafe { on_true.try_as_ref_after_no_overlap() }?;
293 let on_false = unsafe { on_false.try_as_ref_after_no_overlap() }?;
295
296 ctx.run(|| match dtype {
297 KernelDType::F32 => uninit_select::<f32>(dest, &pred, &on_true, &on_false),
298 KernelDType::F64 => uninit_select::<f64>(dest, &pred, &on_true, &on_false),
299 KernelDType::I32 => uninit_select::<i32>(dest, &pred, &on_true, &on_false),
300 KernelDType::I64 => uninit_select::<i64>(dest, &pred, &on_true, &on_false),
301 KernelDType::Bool => uninit_select::<bool>(dest, &pred, &on_true, &on_false),
302 KernelDType::C32 => uninit_select::<Complex32>(dest, &pred, &on_true, &on_false),
303 KernelDType::C64 => uninit_select::<Complex64>(dest, &pred, &on_true, &on_false),
304 _ => Err(StridedError::UnsupportedDType {
305 dtype: dtype.label(),
306 }),
307 })
308}
309
310pub fn erased_clamp_into_uninit(
354 dtype: KernelDType,
355 ctx: &ExecContext,
356 dest: &mut ErasedRawStridedUninitMut<'_>,
357 x: &ErasedRawStridedPtr<'_>,
358 lo: &ErasedRawStridedPtr<'_>,
359 hi: &ErasedRawStridedPtr<'_>,
360) -> Result<()> {
361 check_dtype(dtype, dest.dtype())?;
362 check_dtype(dtype, x.dtype())?;
363 check_dtype(dtype, lo.dtype())?;
364 check_dtype(dtype, hi.dtype())?;
365 validate_uninit_no_overlap(dest, x, 0)?;
366 validate_uninit_no_overlap(dest, lo, 1)?;
367 validate_uninit_no_overlap(dest, hi, 2)?;
368 let x = unsafe { x.try_as_ref_after_no_overlap() }?;
370 let lo = unsafe { lo.try_as_ref_after_no_overlap() }?;
372 let hi = unsafe { hi.try_as_ref_after_no_overlap() }?;
374
375 ctx.run(|| match dtype {
376 KernelDType::F32 => uninit_clamp::<f32>(dest, &x, &lo, &hi),
377 KernelDType::F64 => uninit_clamp::<f64>(dest, &x, &lo, &hi),
378 KernelDType::I32 => uninit_clamp::<i32>(dest, &x, &lo, &hi),
379 KernelDType::I64 => uninit_clamp::<i64>(dest, &x, &lo, &hi),
380 _ => Err(StridedError::UnsupportedDType {
381 dtype: dtype.label(),
382 }),
383 })
384}
385
386pub fn erased_broadcast_mul_into_uninit(
431 dtype: KernelDType,
432 ctx: &ExecContext,
433 dest: &mut ErasedRawStridedUninitMut<'_>,
434 lhs: &ErasedRawStridedPtr<'_>,
435 lhs_axes: &[usize],
436 rhs: &ErasedRawStridedPtr<'_>,
437 rhs_axes: &[usize],
438) -> Result<()> {
439 check_dtype(dtype, dest.dtype())?;
440 check_dtype(dtype, lhs.dtype())?;
441 check_dtype(dtype, rhs.dtype())?;
442 validate_uninit_no_overlap(dest, lhs, 0)?;
443 validate_uninit_no_overlap(dest, rhs, 1)?;
444 let lhs = unsafe { lhs.try_as_ref_after_no_overlap() }?;
446 let rhs = unsafe { rhs.try_as_ref_after_no_overlap() }?;
448
449 ctx.run(|| match dtype {
450 KernelDType::F32 => uninit_broadcast_mul::<f32>(dest, &lhs, lhs_axes, &rhs, rhs_axes),
451 KernelDType::F64 => uninit_broadcast_mul::<f64>(dest, &lhs, lhs_axes, &rhs, rhs_axes),
452 KernelDType::C32 => uninit_broadcast_mul::<Complex32>(dest, &lhs, lhs_axes, &rhs, rhs_axes),
453 KernelDType::C64 => uninit_broadcast_mul::<Complex64>(dest, &lhs, lhs_axes, &rhs, rhs_axes),
454 KernelDType::I32 => {
455 uninit_broadcast_mul_wrapping::<i32>(dest, &lhs, lhs_axes, &rhs, rhs_axes)
456 }
457 KernelDType::I64 => {
458 uninit_broadcast_mul_wrapping::<i64>(dest, &lhs, lhs_axes, &rhs, rhs_axes)
459 }
460 _ => Err(StridedError::UnsupportedDType {
461 dtype: dtype.label(),
462 }),
463 })
464}
465
466fn prepare_destination(
470 dest: &ErasedRawStridedUninitMut<'_>,
471 inputs: &[&[usize]],
472) -> Result<Option<ValidatedDestinationLayout>> {
473 let validated = validate_destination_layout_without_alloc(dest.dims(), dest.strides())?;
474 for input_dims in inputs {
475 ensure_same_shape(dest.dims(), input_dims)?;
476 }
477 if dest.dims().contains(&0) {
478 Ok(None)
479 } else {
480 Ok(Some(validated))
481 }
482}
483
484fn uninit_raw_mut<'a, T: KernelStorageElement>(
485 dest: &'a mut ErasedRawStridedUninitMut<'_>,
486) -> Result<RawStridedMut<'a, MaybeUninit<T>>> {
487 let dims = dest.dims();
488 let strides = dest.strides();
489 let offset = dest.offset();
490 let data = dest.data_as_uninit_mut::<T>()?;
491 Ok(unsafe { RawStridedMut::new_unchecked(data, dims, strides, offset) })
494}
495
496fn uninit_view_mut<'a, T: KernelStorageElement>(
497 dest: &'a mut ErasedRawStridedUninitMut<'_>,
498) -> Result<StridedViewMut<'a, MaybeUninit<T>>> {
499 let dims = dest.dims();
500 let strides = dest.strides();
501 let offset = dest.offset();
502 let data = dest.data_as_uninit_mut::<T>()?;
503 Ok(unsafe { StridedViewMut::new_unchecked(data, dims, strides, offset) })
506}
507
508fn erased_view<'a, T: KernelStorageElement>(
509 src: &'a ErasedRawStridedRef<'a>,
510) -> Result<StridedView<'a, T>> {
511 let data = src.data_as::<T>()?;
512 Ok(unsafe { StridedView::new_unchecked(data, src.dims(), src.strides(), src.offset()) })
514}
515
516fn uninit_map<T: OneShotScalar>(
517 op: ErasedMapOp,
518 dest: &mut ErasedRawStridedUninitMut<'_>,
519 input: &ErasedRawStridedRef<'_>,
520) -> Result<()> {
521 if !T::supports_map(op) {
522 return Err(StridedError::UnsupportedOp {
523 op: op.label(),
524 dtype: T::one_shot_dtype_label(),
525 });
526 }
527 match op {
530 ErasedMapOp::Negate => {
531 uninit_map_with::<T, T>(dest, input, |value| T::map(ErasedMapOp::Negate, value))
532 }
533 ErasedMapOp::Conj => {
534 uninit_map_with::<T, T>(dest, input, |value| T::map(ErasedMapOp::Conj, value))
535 }
536 ErasedMapOp::Abs => {
537 uninit_map_with::<T, T>(dest, input, |value| T::map(ErasedMapOp::Abs, value))
538 }
539 ErasedMapOp::Sign => {
540 uninit_map_with::<T, T>(dest, input, |value| T::map(ErasedMapOp::Sign, value))
541 }
542 }
543}
544
545fn uninit_map_with<D, A>(
546 dest: &mut ErasedRawStridedUninitMut<'_>,
547 input: &ErasedRawStridedRef<'_>,
548 map: impl Fn(A) -> D + crate::MaybeSync,
549) -> Result<()>
550where
551 D: Copy + crate::MaybeSendSync + KernelStorageElement,
552 A: Copy + crate::MaybeSendSync + KernelStorageElement,
553{
554 let Some(validated) = prepare_destination(dest, &[input.dims()])? else {
555 return Ok(());
556 };
557 let input = erased_raw_ref::<A>(input)?;
558 let mut dest = uninit_raw_mut::<D>(dest)?;
559 unsafe {
562 map_raw_into_validated::<MaybeUninit<D>, A, Identity>(
563 &mut dest,
564 &input,
565 |value| MaybeUninit::new(map(value)),
566 validated,
567 )
568 }
569}
570
571fn uninit_zip<T: OneShotScalar>(
572 op: ErasedZipOp,
573 dest: &mut ErasedRawStridedUninitMut<'_>,
574 lhs: &ErasedRawStridedRef<'_>,
575 rhs: &ErasedRawStridedRef<'_>,
576) -> Result<()> {
577 if !T::supports_zip(op) {
578 return Err(StridedError::UnsupportedOp {
579 op: op.label(),
580 dtype: T::one_shot_dtype_label(),
581 });
582 }
583 let Some(validated) = prepare_destination(dest, &[lhs.dims(), rhs.dims()])? else {
584 return Ok(());
585 };
586 let lhs = erased_raw_ref::<T>(lhs)?;
587 let rhs = erased_raw_ref::<T>(rhs)?;
588 if matches!(op, ErasedZipOp::Divide | ErasedZipOp::Remainder)
589 && T::INTEGER
590 && raw_any(&rhs, T::is_zero)?
591 {
592 return Err(StridedError::IntegerDivisionByZero { op: op.label() });
593 }
594 let mut dest = uninit_raw_mut::<T>(dest)?;
595 macro_rules! replay {
598 ($op:ident) => {
599 unsafe {
602 zip_map2_raw_into_validated::<MaybeUninit<T>, T, T, Identity, Identity>(
603 &mut dest,
604 &lhs,
605 &rhs,
606 |lhs, rhs| MaybeUninit::new(T::zip(ErasedZipOp::$op, lhs, rhs)),
607 validated,
608 )
609 }
610 };
611 }
612 match op {
613 ErasedZipOp::Add => replay!(Add),
614 ErasedZipOp::Subtract => replay!(Subtract),
615 ErasedZipOp::Multiply => replay!(Multiply),
616 ErasedZipOp::Divide => replay!(Divide),
617 ErasedZipOp::Remainder => replay!(Remainder),
618 ErasedZipOp::Maximum => replay!(Maximum),
619 ErasedZipOp::Minimum => replay!(Minimum),
620 }
621}
622
623fn uninit_compare_with<T>(
624 dest: &mut ErasedRawStridedUninitMut<'_>,
625 lhs: &ErasedRawStridedRef<'_>,
626 rhs: &ErasedRawStridedRef<'_>,
627 compare: impl Fn(T, T) -> bool + crate::MaybeSync,
628) -> Result<()>
629where
630 T: Copy + crate::MaybeSendSync + KernelStorageElement,
631{
632 let Some(validated) = prepare_destination(dest, &[lhs.dims(), rhs.dims()])? else {
633 return Ok(());
634 };
635 let lhs = erased_raw_ref::<T>(lhs)?;
636 let rhs = erased_raw_ref::<T>(rhs)?;
637 let mut dest = uninit_raw_mut::<bool>(dest)?;
638 unsafe {
641 zip_map2_raw_into_validated::<MaybeUninit<bool>, T, T, Identity, Identity>(
642 &mut dest,
643 &lhs,
644 &rhs,
645 |lhs, rhs| MaybeUninit::new(compare(lhs, rhs)),
646 validated,
647 )
648 }
649}
650
651fn uninit_compare_ordered<T>(
652 op: CompareOp,
653 dest: &mut ErasedRawStridedUninitMut<'_>,
654 lhs: &ErasedRawStridedRef<'_>,
655 rhs: &ErasedRawStridedRef<'_>,
656) -> Result<()>
657where
658 T: Copy + crate::MaybeSendSync + KernelStorageElement + PartialOrd,
659{
660 match op {
662 CompareOp::Eq => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a == b),
663 CompareOp::Lt => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a < b),
664 CompareOp::Le => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a <= b),
665 CompareOp::Gt => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a > b),
666 CompareOp::Ge => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a >= b),
667 _ => Err(StridedError::UnsupportedOp {
668 op: compare_label(op),
669 dtype: T::DTYPE.label(),
670 }),
671 }
672}
673
674fn uninit_compare_eq<T>(
675 op: CompareOp,
676 dest: &mut ErasedRawStridedUninitMut<'_>,
677 lhs: &ErasedRawStridedRef<'_>,
678 rhs: &ErasedRawStridedRef<'_>,
679) -> Result<()>
680where
681 T: Copy + crate::MaybeSendSync + KernelStorageElement + PartialEq,
682{
683 match op {
684 CompareOp::Eq => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a == b),
685 _ => Err(StridedError::UnsupportedOp {
686 op: compare_label(op),
687 dtype: T::DTYPE.label(),
688 }),
689 }
690}
691
692fn compare_label(op: CompareOp) -> &'static str {
693 match op {
694 CompareOp::Eq => "eq",
695 CompareOp::Lt => "lt",
696 CompareOp::Le => "le",
697 CompareOp::Gt => "gt",
698 CompareOp::Ge => "ge",
699 _ => "compare",
700 }
701}
702
703fn uninit_select<T>(
704 dest: &mut ErasedRawStridedUninitMut<'_>,
705 pred: &ErasedRawStridedRef<'_>,
706 on_true: &ErasedRawStridedRef<'_>,
707 on_false: &ErasedRawStridedRef<'_>,
708) -> Result<()>
709where
710 T: Copy + crate::MaybeSendSync + KernelStorageElement,
711{
712 let dims: [&[usize]; 3] = [pred.dims(), on_true.dims(), on_false.dims()];
713 if prepare_destination(dest, &dims)?.is_none() {
714 return Ok(());
715 }
716 let pred = erased_view::<bool>(pred)?;
717 let on_true = erased_view::<T>(on_true)?;
718 let on_false = erased_view::<T>(on_false)?;
719 let mut dest = uninit_view_mut::<T>(dest)?;
720 zip_map3_into(&mut dest, &pred, &on_true, &on_false, |pred, a, b| {
721 MaybeUninit::new(if pred { a } else { b })
722 })
723}
724
725fn uninit_clamp<T: OneShotScalar>(
726 dest: &mut ErasedRawStridedUninitMut<'_>,
727 x: &ErasedRawStridedRef<'_>,
728 lo: &ErasedRawStridedRef<'_>,
729 hi: &ErasedRawStridedRef<'_>,
730) -> Result<()> {
731 let dims: [&[usize]; 3] = [x.dims(), lo.dims(), hi.dims()];
732 if prepare_destination(dest, &dims)?.is_none() {
733 return Ok(());
734 }
735 let x = erased_view::<T>(x)?;
736 let lo = erased_view::<T>(lo)?;
737 let hi = erased_view::<T>(hi)?;
738 let mut dest = uninit_view_mut::<T>(dest)?;
739 zip_map3_into(&mut dest, &x, &lo, &hi, |x, lo, hi| {
740 MaybeUninit::new(T::clamp(x, lo, hi))
741 })
742}
743
744fn uninit_broadcast_mul<T>(
745 dest: &mut ErasedRawStridedUninitMut<'_>,
746 lhs: &ErasedRawStridedRef<'_>,
747 lhs_axes: &[usize],
748 rhs: &ErasedRawStridedRef<'_>,
749 rhs_axes: &[usize],
750) -> Result<()>
751where
752 T: Copy + crate::MaybeSendSync + KernelStorageElement + core::ops::Mul<Output = T>,
753{
754 let lhs = erased_view::<T>(lhs)?;
755 let rhs = erased_view::<T>(rhs)?;
756 let mut dest = uninit_view_mut::<T>(dest)?;
757 broadcast_mul_into_uninit(&mut dest, &lhs, lhs_axes, &rhs, rhs_axes)
758}
759
760fn uninit_broadcast_mul_wrapping<T>(
761 dest: &mut ErasedRawStridedUninitMut<'_>,
762 lhs: &ErasedRawStridedRef<'_>,
763 lhs_axes: &[usize],
764 rhs: &ErasedRawStridedRef<'_>,
765 rhs_axes: &[usize],
766) -> Result<()>
767where
768 T: Copy + crate::MaybeSendSync + KernelStorageElement,
769 Wrapping<T>: core::ops::Mul<Output = Wrapping<T>> + crate::MaybeSendSync,
770{
771 let lhs_data = wrapping_slice(lhs.data_as::<T>()?);
772 let rhs_data = wrapping_slice(rhs.data_as::<T>()?);
773 let lhs = unsafe {
776 StridedView::<Wrapping<T>, Identity>::new_unchecked(
777 lhs_data,
778 lhs.dims(),
779 lhs.strides(),
780 lhs.offset(),
781 )
782 };
783 let rhs = unsafe {
785 StridedView::<Wrapping<T>, Identity>::new_unchecked(
786 rhs_data,
787 rhs.dims(),
788 rhs.strides(),
789 rhs.offset(),
790 )
791 };
792 let dims = dest.dims();
793 let strides = dest.strides();
794 let offset = dest.offset();
795 let data = dest.data_as_uninit_mut::<T>()?;
796 let data = unsafe {
800 core::slice::from_raw_parts_mut(
801 data.as_mut_ptr().cast::<MaybeUninit<Wrapping<T>>>(),
802 data.len(),
803 )
804 };
805 let mut dest = unsafe { StridedViewMut::new_unchecked(data, dims, strides, offset) };
807 broadcast_mul_into_uninit(&mut dest, &lhs, lhs_axes, &rhs, rhs_axes)
808}
809
810fn wrapping_slice<T>(data: &[T]) -> &[Wrapping<T>] {
811 unsafe { core::slice::from_raw_parts(data.as_ptr().cast::<Wrapping<T>>(), data.len()) }
814}