1use crate::*;
2use core::ops::Add;
3use num_complex::{Complex32, Complex64};
4use strided_basic::execution::check_static_indexing_dtype;
5use strided_basic::execution::{check_dtype, validate_uninit_no_overlap};
6
7mod uninit;
8pub use uninit::{
9 erased_broadcast_mul_into_uninit, erased_clamp_into_uninit, erased_compare_into_uninit,
10 erased_map_into_uninit, erased_select_into_uninit, erased_zip_into_uninit,
11};
12#[non_exhaustive]
14#[derive(Clone, Copy, Debug, Eq, PartialEq)]
15pub enum ErasedMapOp {
16 Negate,
17 Conj,
18 Abs,
19 Sign,
20}
21
22impl ErasedMapOp {
23 const fn label(self) -> &'static str {
24 match self {
25 Self::Negate => "negate",
26 Self::Conj => "conj",
27 Self::Abs => "abs",
28 Self::Sign => "sign",
29 }
30 }
31}
32
33#[non_exhaustive]
35#[derive(Clone, Copy, Debug, Eq, PartialEq)]
36pub enum ErasedZipOp {
37 Add,
38 Subtract,
39 Multiply,
40 Divide,
41 Remainder,
42 Maximum,
43 Minimum,
44}
45
46impl ErasedZipOp {
47 const fn label(self) -> &'static str {
48 match self {
49 Self::Add => "add",
50 Self::Subtract => "subtract",
51 Self::Multiply => "multiply",
52 Self::Divide => "divide",
53 Self::Remainder => "remainder",
54 Self::Maximum => "maximum",
55 Self::Minimum => "minimum",
56 }
57 }
58}
59
60pub fn erased_map_into(
73 input_dtype: KernelDType,
74 op: ErasedMapOp,
75 ctx: &ExecContext,
76 dest: &mut ErasedRawStridedMut<'_>,
77 input: &ErasedRawStridedPtr<'_>,
78) -> Result<()> {
79 check_dtype(input_dtype, input.dtype())?;
80 check_dtype(map_output_dtype(input_dtype, op)?, dest.dtype())?;
81 validate_no_overlap(dest, input, 0)?;
82 let input = unsafe { input.try_as_ref_after_no_overlap() }?;
84
85 let result = ctx.run(|| match (input_dtype, op) {
86 (KernelDType::C32, ErasedMapOp::Abs) => {
87 execute_one_shot_map_with::<f32, Complex32>(dest, &input, |value| value.norm())
88 }
89 (KernelDType::C64, ErasedMapOp::Abs) => {
90 execute_one_shot_map_with::<f64, Complex64>(dest, &input, |value| value.norm())
91 }
92 (KernelDType::F32, _) => execute_one_shot_map::<f32>(op, dest, &input),
93 (KernelDType::F64, _) => execute_one_shot_map::<f64>(op, dest, &input),
94 (KernelDType::I32, _) => execute_one_shot_map::<i32>(op, dest, &input),
95 (KernelDType::I64, _) => execute_one_shot_map::<i64>(op, dest, &input),
96 (KernelDType::Bool, _) => execute_one_shot_map::<bool>(op, dest, &input),
97 (KernelDType::C32, _) => execute_one_shot_map::<Complex32>(op, dest, &input),
98 (KernelDType::C64, _) => execute_one_shot_map::<Complex64>(op, dest, &input),
99 _ => Err(StridedError::UnsupportedDType {
100 dtype: input_dtype.label(),
101 }),
102 });
103 result
104}
105
106pub fn erased_zip_into(
118 dtype: KernelDType,
119 op: ErasedZipOp,
120 ctx: &ExecContext,
121 dest: &mut ErasedRawStridedMut<'_>,
122 lhs: &ErasedRawStridedPtr<'_>,
123 rhs: &ErasedRawStridedPtr<'_>,
124) -> Result<()> {
125 check_dtype(dtype, dest.dtype())?;
126 check_dtype(dtype, lhs.dtype())?;
127 check_dtype(dtype, rhs.dtype())?;
128 validate_no_overlap(dest, lhs, 0)?;
129 validate_no_overlap(dest, rhs, 1)?;
130 let lhs = unsafe { lhs.try_as_ref_after_no_overlap() }?;
132 let rhs = unsafe { rhs.try_as_ref_after_no_overlap() }?;
134
135 let result = ctx.run(|| match dtype {
136 KernelDType::F32 => execute_one_shot_zip::<f32>(op, dest, &lhs, &rhs),
137 KernelDType::F64 => execute_one_shot_zip::<f64>(op, dest, &lhs, &rhs),
138 KernelDType::I32 => execute_one_shot_zip::<i32>(op, dest, &lhs, &rhs),
139 KernelDType::I64 => execute_one_shot_zip::<i64>(op, dest, &lhs, &rhs),
140 KernelDType::Bool => execute_one_shot_zip::<bool>(op, dest, &lhs, &rhs),
141 KernelDType::C32 => execute_one_shot_zip::<Complex32>(op, dest, &lhs, &rhs),
142 KernelDType::C64 => execute_one_shot_zip::<Complex64>(op, dest, &lhs, &rhs),
143 _ => Err(StridedError::UnsupportedDType {
144 dtype: dtype.label(),
145 }),
146 });
147 result
148}
149
150#[derive(Clone, Debug)]
152pub struct ErasedSlicePlan {
153 dtype: KernelDType,
154 plan: SlicePlan,
155}
156
157#[derive(Clone, Debug)]
159pub struct ErasedReversePlan {
160 dtype: KernelDType,
161 plan: ReversePlan,
162}
163
164#[derive(Clone, Debug)]
166pub struct ErasedPadPlan {
167 dtype: KernelDType,
168 plan: PadPlan,
169}
170
171impl ErasedSlicePlan {
172 #[allow(clippy::too_many_arguments)]
174 pub fn compile(
175 dtype: KernelDType,
176 operand_dims: &[usize],
177 operand_strides: &[isize],
178 dest_dims: &[usize],
179 dest_strides: &[isize],
180 starts: &[usize],
181 limits: &[usize],
182 slice_strides: &[usize],
183 ) -> Result<Self> {
184 check_static_indexing_dtype(dtype)?;
185 Ok(Self {
186 dtype,
187 plan: SlicePlan::compile(
188 operand_dims,
189 operand_strides,
190 dest_dims,
191 dest_strides,
192 starts,
193 limits,
194 slice_strides,
195 )?,
196 })
197 }
198
199 #[inline]
200 pub fn dtype(&self) -> KernelDType {
201 self.dtype
202 }
203
204 #[inline]
205 pub fn plan(&self) -> &SlicePlan {
206 &self.plan
207 }
208
209 pub fn execute(
211 &self,
212 ctx: &ExecContext,
213 dest: &mut ErasedRawStridedMut<'_>,
214 operand: &ErasedRawStridedRef<'_>,
215 ) -> Result<()> {
216 check_dtype(self.dtype, dest.dtype())?;
217 check_dtype(self.dtype, operand.dtype())?;
218
219 let result = ctx.run(|| match self.dtype {
220 KernelDType::F32 => execute_slice::<f32>(&self.plan, dest, operand),
221 KernelDType::F64 => execute_slice::<f64>(&self.plan, dest, operand),
222 KernelDType::I32 => execute_slice::<i32>(&self.plan, dest, operand),
223 KernelDType::I64 => execute_slice::<i64>(&self.plan, dest, operand),
224 KernelDType::Bool => execute_slice::<bool>(&self.plan, dest, operand),
225 KernelDType::C32 => execute_slice::<Complex32>(&self.plan, dest, operand),
226 KernelDType::C64 => execute_slice::<Complex64>(&self.plan, dest, operand),
227 _ => Err(StridedError::UnsupportedDType {
228 dtype: self.dtype.label(),
229 }),
230 });
231 result
232 }
233
234 pub fn execute_uninit(
242 &self,
243 ctx: &ExecContext,
244 dest: &mut ErasedRawStridedUninitMut<'_>,
245 operand: &ErasedRawStridedPtr<'_>,
246 ) -> Result<()> {
247 check_dtype(self.dtype, dest.dtype())?;
248 check_dtype(self.dtype, operand.dtype())?;
249 validate_uninit_no_overlap(dest, operand, 0)?;
250 let operand = unsafe { operand.try_as_ref_after_no_overlap() }?;
252
253 ctx.run(|| match self.dtype {
254 KernelDType::F32 => execute_slice_uninit::<f32>(&self.plan, dest, &operand),
255 KernelDType::F64 => execute_slice_uninit::<f64>(&self.plan, dest, &operand),
256 KernelDType::I32 => execute_slice_uninit::<i32>(&self.plan, dest, &operand),
257 KernelDType::I64 => execute_slice_uninit::<i64>(&self.plan, dest, &operand),
258 KernelDType::Bool => execute_slice_uninit::<bool>(&self.plan, dest, &operand),
259 KernelDType::C32 => execute_slice_uninit::<Complex32>(&self.plan, dest, &operand),
260 KernelDType::C64 => execute_slice_uninit::<Complex64>(&self.plan, dest, &operand),
261 _ => Err(StridedError::UnsupportedDType {
262 dtype: self.dtype.label(),
263 }),
264 })
265 }
266}
267
268impl ErasedReversePlan {
269 pub fn compile(
271 dtype: KernelDType,
272 operand_dims: &[usize],
273 operand_strides: &[isize],
274 dest_strides: &[isize],
275 axes: &[usize],
276 ) -> Result<Self> {
277 check_static_indexing_dtype(dtype)?;
278 Ok(Self {
279 dtype,
280 plan: ReversePlan::compile(operand_dims, operand_strides, dest_strides, axes)?,
281 })
282 }
283
284 #[inline]
285 pub fn dtype(&self) -> KernelDType {
286 self.dtype
287 }
288
289 #[inline]
290 pub fn plan(&self) -> &ReversePlan {
291 &self.plan
292 }
293
294 pub fn execute(
296 &self,
297 ctx: &ExecContext,
298 dest: &mut ErasedRawStridedMut<'_>,
299 operand: &ErasedRawStridedRef<'_>,
300 ) -> Result<()> {
301 check_dtype(self.dtype, dest.dtype())?;
302 check_dtype(self.dtype, operand.dtype())?;
303
304 let result = ctx.run(|| match self.dtype {
305 KernelDType::F32 => execute_reverse::<f32>(&self.plan, dest, operand),
306 KernelDType::F64 => execute_reverse::<f64>(&self.plan, dest, operand),
307 KernelDType::I32 => execute_reverse::<i32>(&self.plan, dest, operand),
308 KernelDType::I64 => execute_reverse::<i64>(&self.plan, dest, operand),
309 KernelDType::Bool => execute_reverse::<bool>(&self.plan, dest, operand),
310 KernelDType::C32 => execute_reverse::<Complex32>(&self.plan, dest, operand),
311 KernelDType::C64 => execute_reverse::<Complex64>(&self.plan, dest, operand),
312 _ => Err(StridedError::UnsupportedDType {
313 dtype: self.dtype.label(),
314 }),
315 });
316 result
317 }
318
319 pub fn execute_uninit(
327 &self,
328 ctx: &ExecContext,
329 dest: &mut ErasedRawStridedUninitMut<'_>,
330 operand: &ErasedRawStridedPtr<'_>,
331 ) -> Result<()> {
332 check_dtype(self.dtype, dest.dtype())?;
333 check_dtype(self.dtype, operand.dtype())?;
334 validate_uninit_no_overlap(dest, operand, 0)?;
335 let operand = unsafe { operand.try_as_ref_after_no_overlap() }?;
337
338 ctx.run(|| match self.dtype {
339 KernelDType::F32 => execute_reverse_uninit::<f32>(&self.plan, dest, &operand),
340 KernelDType::F64 => execute_reverse_uninit::<f64>(&self.plan, dest, &operand),
341 KernelDType::I32 => execute_reverse_uninit::<i32>(&self.plan, dest, &operand),
342 KernelDType::I64 => execute_reverse_uninit::<i64>(&self.plan, dest, &operand),
343 KernelDType::Bool => execute_reverse_uninit::<bool>(&self.plan, dest, &operand),
344 KernelDType::C32 => execute_reverse_uninit::<Complex32>(&self.plan, dest, &operand),
345 KernelDType::C64 => execute_reverse_uninit::<Complex64>(&self.plan, dest, &operand),
346 _ => Err(StridedError::UnsupportedDType {
347 dtype: self.dtype.label(),
348 }),
349 })
350 }
351}
352
353impl ErasedPadPlan {
354 #[allow(clippy::too_many_arguments)]
356 pub fn compile(
357 dtype: KernelDType,
358 operand_dims: &[usize],
359 operand_strides: &[isize],
360 dest_dims: &[usize],
361 dest_strides: &[isize],
362 edge_padding_low: &[i64],
363 edge_padding_high: &[i64],
364 interior_padding: &[i64],
365 ) -> Result<Self> {
366 check_static_indexing_dtype(dtype)?;
367 Ok(Self {
368 dtype,
369 plan: PadPlan::compile(
370 operand_dims,
371 operand_strides,
372 dest_dims,
373 dest_strides,
374 edge_padding_low,
375 edge_padding_high,
376 interior_padding,
377 )?,
378 })
379 }
380
381 #[inline]
382 pub fn dtype(&self) -> KernelDType {
383 self.dtype
384 }
385
386 #[inline]
387 pub fn plan(&self) -> &PadPlan {
388 &self.plan
389 }
390
391 pub fn execute(
393 &self,
394 ctx: &ExecContext,
395 dest: &mut ErasedRawStridedMut<'_>,
396 operand: &ErasedRawStridedRef<'_>,
397 fill: &[u8],
398 ) -> Result<()> {
399 check_dtype(self.dtype, dest.dtype())?;
400 check_dtype(self.dtype, operand.dtype())?;
401 validate_scalar_bytes(self.dtype, fill)?;
402
403 let result = ctx.run(|| match self.dtype {
404 KernelDType::F32 => execute_pad::<f32>(&self.plan, dest, operand, fill),
405 KernelDType::F64 => execute_pad::<f64>(&self.plan, dest, operand, fill),
406 KernelDType::I32 => execute_pad::<i32>(&self.plan, dest, operand, fill),
407 KernelDType::I64 => execute_pad::<i64>(&self.plan, dest, operand, fill),
408 KernelDType::Bool => execute_pad::<bool>(&self.plan, dest, operand, fill),
409 KernelDType::C32 => execute_pad::<Complex32>(&self.plan, dest, operand, fill),
410 KernelDType::C64 => execute_pad::<Complex64>(&self.plan, dest, operand, fill),
411 _ => Err(StridedError::UnsupportedDType {
412 dtype: self.dtype.label(),
413 }),
414 });
415 result
416 }
417
418 pub fn execute_uninit(
426 &self,
427 ctx: &ExecContext,
428 dest: &mut ErasedRawStridedUninitMut<'_>,
429 operand: &ErasedRawStridedPtr<'_>,
430 fill: &[u8],
431 ) -> Result<()> {
432 check_dtype(self.dtype, dest.dtype())?;
433 check_dtype(self.dtype, operand.dtype())?;
434 validate_scalar_bytes(self.dtype, fill)?;
435 validate_uninit_no_overlap(dest, operand, 0)?;
436 let operand = unsafe { operand.try_as_ref_after_no_overlap() }?;
438
439 ctx.run(|| match self.dtype {
440 KernelDType::F32 => execute_pad_uninit::<f32>(&self.plan, dest, &operand, fill),
441 KernelDType::F64 => execute_pad_uninit::<f64>(&self.plan, dest, &operand, fill),
442 KernelDType::I32 => execute_pad_uninit::<i32>(&self.plan, dest, &operand, fill),
443 KernelDType::I64 => execute_pad_uninit::<i64>(&self.plan, dest, &operand, fill),
444 KernelDType::Bool => execute_pad_uninit::<bool>(&self.plan, dest, &operand, fill),
445 KernelDType::C32 => execute_pad_uninit::<Complex32>(&self.plan, dest, &operand, fill),
446 KernelDType::C64 => execute_pad_uninit::<Complex64>(&self.plan, dest, &operand, fill),
447 _ => Err(StridedError::UnsupportedDType {
448 dtype: self.dtype.label(),
449 }),
450 })
451 }
452}
453
454#[derive(Clone, Debug)]
459pub struct ErasedGatherPlan {
460 dtype: KernelDType,
461 index_dtype: KernelDType,
462 plan: GatherPlan,
463}
464
465#[derive(Clone, Debug)]
467pub struct ErasedDynamicSlicePlan {
468 dtype: KernelDType,
469 index_dtype: KernelDType,
470 plan: DynamicSlicePlan,
471}
472
473#[derive(Clone, Debug)]
475pub struct ErasedDynamicUpdateSlicePlan {
476 dtype: KernelDType,
477 index_dtype: KernelDType,
478 plan: DynamicUpdateSlicePlan,
479}
480
481#[derive(Clone, Debug)]
483pub struct ErasedScatterPlan {
484 dtype: KernelDType,
485 index_dtype: KernelDType,
486 plan: ScatterPlan,
487}
488
489impl ErasedGatherPlan {
490 #[allow(clippy::too_many_arguments)]
492 pub fn compile(
493 dtype: KernelDType,
494 index_dtype: KernelDType,
495 operand_dims: &[usize],
496 operand_strides: &[isize],
497 index_dims: &[usize],
498 index_strides: &[isize],
499 dest_dims: &[usize],
500 dest_strides: &[isize],
501 spec: GatherSpec,
502 ) -> Result<Self> {
503 check_index_dtype(index_dtype)?;
504 check_gather_value_dtype(dtype)?;
505 Ok(Self {
506 dtype,
507 index_dtype,
508 plan: GatherPlan::compile(
509 operand_dims,
510 operand_strides,
511 index_dims,
512 index_strides,
513 dest_dims,
514 dest_strides,
515 spec,
516 )?,
517 })
518 }
519
520 #[inline]
521 pub fn dtype(&self) -> KernelDType {
522 self.dtype
523 }
524
525 #[inline]
526 pub fn index_dtype(&self) -> KernelDType {
527 self.index_dtype
528 }
529
530 #[inline]
531 pub fn plan(&self) -> &GatherPlan {
532 &self.plan
533 }
534
535 pub fn execute(
537 &self,
538 ctx: &ExecContext,
539 dest: &mut ErasedRawStridedMut<'_>,
540 operand: &ErasedRawStridedRef<'_>,
541 start_indices: &ErasedRawStridedRef<'_>,
542 ) -> Result<()> {
543 check_dtype(self.dtype, dest.dtype())?;
544 check_dtype(self.dtype, operand.dtype())?;
545 check_dtype(self.index_dtype, start_indices.dtype())?;
546
547 let result = ctx.run(|| match self.dtype {
548 KernelDType::F32 => dispatch_gather_index::<f32>(
549 &self.plan,
550 self.index_dtype,
551 dest,
552 &operand,
553 &start_indices,
554 ),
555 KernelDType::F64 => dispatch_gather_index::<f64>(
556 &self.plan,
557 self.index_dtype,
558 dest,
559 operand,
560 start_indices,
561 ),
562 KernelDType::I32 => dispatch_gather_index::<i32>(
563 &self.plan,
564 self.index_dtype,
565 dest,
566 operand,
567 start_indices,
568 ),
569 KernelDType::I64 => dispatch_gather_index::<i64>(
570 &self.plan,
571 self.index_dtype,
572 dest,
573 operand,
574 start_indices,
575 ),
576 KernelDType::Bool => dispatch_gather_index::<bool>(
577 &self.plan,
578 self.index_dtype,
579 dest,
580 operand,
581 start_indices,
582 ),
583 KernelDType::C32 => dispatch_gather_index::<Complex32>(
584 &self.plan,
585 self.index_dtype,
586 dest,
587 operand,
588 start_indices,
589 ),
590 KernelDType::C64 => dispatch_gather_index::<Complex64>(
591 &self.plan,
592 self.index_dtype,
593 dest,
594 operand,
595 start_indices,
596 ),
597 _ => Err(StridedError::UnsupportedDType {
598 dtype: self.dtype.label(),
599 }),
600 });
601 result
602 }
603
604 pub fn execute_uninit(
613 &self,
614 ctx: &ExecContext,
615 dest: &mut ErasedRawStridedUninitMut<'_>,
616 operand: &ErasedRawStridedPtr<'_>,
617 start_indices: &ErasedRawStridedPtr<'_>,
618 ) -> Result<()> {
619 check_dtype(self.dtype, dest.dtype())?;
620 check_dtype(self.dtype, operand.dtype())?;
621 check_dtype(self.index_dtype, start_indices.dtype())?;
622 validate_uninit_no_overlap(dest, operand, 0)?;
623 validate_uninit_no_overlap(dest, start_indices, 1)?;
624 let operand = &unsafe { operand.try_as_ref_after_no_overlap() }?;
626 let start_indices = &unsafe { start_indices.try_as_ref_after_no_overlap() }?;
628 let run = |dest: &mut ErasedRawStridedUninitMut<'_>| match self.dtype {
629 KernelDType::F32 => execute_gather_uninit_dispatch::<f32>(
630 &self.plan,
631 self.index_dtype,
632 dest,
633 operand,
634 start_indices,
635 ),
636 KernelDType::F64 => execute_gather_uninit_dispatch::<f64>(
637 &self.plan,
638 self.index_dtype,
639 dest,
640 operand,
641 start_indices,
642 ),
643 KernelDType::I32 => execute_gather_uninit_dispatch::<i32>(
644 &self.plan,
645 self.index_dtype,
646 dest,
647 operand,
648 start_indices,
649 ),
650 KernelDType::I64 => execute_gather_uninit_dispatch::<i64>(
651 &self.plan,
652 self.index_dtype,
653 dest,
654 operand,
655 start_indices,
656 ),
657 KernelDType::Bool => execute_gather_uninit_dispatch::<bool>(
658 &self.plan,
659 self.index_dtype,
660 dest,
661 operand,
662 start_indices,
663 ),
664 KernelDType::C32 => execute_gather_uninit_dispatch::<Complex32>(
665 &self.plan,
666 self.index_dtype,
667 dest,
668 operand,
669 start_indices,
670 ),
671 KernelDType::C64 => execute_gather_uninit_dispatch::<Complex64>(
672 &self.plan,
673 self.index_dtype,
674 dest,
675 operand,
676 start_indices,
677 ),
678 _ => Err(StridedError::UnsupportedDType {
679 dtype: self.dtype.label(),
680 }),
681 };
682 if ctx.is_serial() {
683 run(dest)
684 } else {
685 ctx.run(|| run(dest))
686 }
687 }
688}
689
690impl ErasedDynamicSlicePlan {
691 #[allow(clippy::too_many_arguments)]
693 pub fn compile(
694 dtype: KernelDType,
695 index_dtype: KernelDType,
696 operand_dims: &[usize],
697 operand_strides: &[isize],
698 start_dims: &[usize],
699 start_strides: &[isize],
700 dest_dims: &[usize],
701 dest_strides: &[isize],
702 slice_sizes: &[usize],
703 ) -> Result<Self> {
704 check_index_dtype(index_dtype)?;
705 check_gather_value_dtype(dtype)?;
706 Ok(Self {
707 dtype,
708 index_dtype,
709 plan: DynamicSlicePlan::compile(
710 operand_dims,
711 operand_strides,
712 start_dims,
713 start_strides,
714 dest_dims,
715 dest_strides,
716 slice_sizes,
717 )?,
718 })
719 }
720
721 #[inline]
722 pub fn dtype(&self) -> KernelDType {
723 self.dtype
724 }
725
726 #[inline]
727 pub fn index_dtype(&self) -> KernelDType {
728 self.index_dtype
729 }
730
731 #[inline]
732 pub fn plan(&self) -> &DynamicSlicePlan {
733 &self.plan
734 }
735
736 pub fn execute(
738 &self,
739 ctx: &ExecContext,
740 dest: &mut ErasedRawStridedMut<'_>,
741 operand: &ErasedRawStridedRef<'_>,
742 starts: &ErasedRawStridedRef<'_>,
743 ) -> Result<()> {
744 check_dtype(self.dtype, dest.dtype())?;
745 check_dtype(self.dtype, operand.dtype())?;
746 check_dtype(self.index_dtype, starts.dtype())?;
747
748 let result = ctx.run(|| match self.dtype {
749 KernelDType::F32 => dispatch_dynamic_slice_index::<f32>(
750 &self.plan,
751 self.index_dtype,
752 dest,
753 &operand,
754 &starts,
755 ),
756 KernelDType::F64 => dispatch_dynamic_slice_index::<f64>(
757 &self.plan,
758 self.index_dtype,
759 dest,
760 operand,
761 starts,
762 ),
763 KernelDType::I32 => dispatch_dynamic_slice_index::<i32>(
764 &self.plan,
765 self.index_dtype,
766 dest,
767 operand,
768 starts,
769 ),
770 KernelDType::I64 => dispatch_dynamic_slice_index::<i64>(
771 &self.plan,
772 self.index_dtype,
773 dest,
774 operand,
775 starts,
776 ),
777 KernelDType::Bool => dispatch_dynamic_slice_index::<bool>(
778 &self.plan,
779 self.index_dtype,
780 dest,
781 operand,
782 starts,
783 ),
784 KernelDType::C32 => dispatch_dynamic_slice_index::<Complex32>(
785 &self.plan,
786 self.index_dtype,
787 dest,
788 operand,
789 starts,
790 ),
791 KernelDType::C64 => dispatch_dynamic_slice_index::<Complex64>(
792 &self.plan,
793 self.index_dtype,
794 dest,
795 operand,
796 starts,
797 ),
798 _ => Err(StridedError::UnsupportedDType {
799 dtype: self.dtype.label(),
800 }),
801 });
802 result
803 }
804
805 pub fn execute_uninit(
813 &self,
814 ctx: &ExecContext,
815 dest: &mut ErasedRawStridedUninitMut<'_>,
816 operand: &ErasedRawStridedPtr<'_>,
817 starts: &ErasedRawStridedPtr<'_>,
818 ) -> Result<()> {
819 check_dtype(self.dtype, dest.dtype())?;
820 check_dtype(self.dtype, operand.dtype())?;
821 check_dtype(self.index_dtype, starts.dtype())?;
822 validate_uninit_no_overlap(dest, operand, 0)?;
823 validate_uninit_no_overlap(dest, starts, 1)?;
824 let operand = &unsafe { operand.try_as_ref_after_no_overlap() }?;
826 let starts = &unsafe { starts.try_as_ref_after_no_overlap() }?;
828 let run = |dest: &mut ErasedRawStridedUninitMut<'_>| match self.dtype {
829 KernelDType::F32 => execute_dynamic_slice_uninit_dispatch::<f32>(
830 &self.plan,
831 self.index_dtype,
832 dest,
833 operand,
834 starts,
835 ),
836 KernelDType::F64 => execute_dynamic_slice_uninit_dispatch::<f64>(
837 &self.plan,
838 self.index_dtype,
839 dest,
840 operand,
841 starts,
842 ),
843 KernelDType::I32 => execute_dynamic_slice_uninit_dispatch::<i32>(
844 &self.plan,
845 self.index_dtype,
846 dest,
847 operand,
848 starts,
849 ),
850 KernelDType::I64 => execute_dynamic_slice_uninit_dispatch::<i64>(
851 &self.plan,
852 self.index_dtype,
853 dest,
854 operand,
855 starts,
856 ),
857 KernelDType::Bool => execute_dynamic_slice_uninit_dispatch::<bool>(
858 &self.plan,
859 self.index_dtype,
860 dest,
861 operand,
862 starts,
863 ),
864 KernelDType::C32 => execute_dynamic_slice_uninit_dispatch::<Complex32>(
865 &self.plan,
866 self.index_dtype,
867 dest,
868 operand,
869 starts,
870 ),
871 KernelDType::C64 => execute_dynamic_slice_uninit_dispatch::<Complex64>(
872 &self.plan,
873 self.index_dtype,
874 dest,
875 operand,
876 starts,
877 ),
878 _ => Err(StridedError::UnsupportedDType {
879 dtype: self.dtype.label(),
880 }),
881 };
882 if ctx.is_serial() {
883 run(dest)
884 } else {
885 ctx.run(|| run(dest))
886 }
887 }
888}
889
890impl ErasedDynamicUpdateSlicePlan {
891 #[allow(clippy::too_many_arguments)]
893 pub fn compile(
894 dtype: KernelDType,
895 index_dtype: KernelDType,
896 operand_dims: &[usize],
897 operand_strides: &[isize],
898 start_dims: &[usize],
899 start_strides: &[isize],
900 update_dims: &[usize],
901 update_strides: &[isize],
902 dest_dims: &[usize],
903 dest_strides: &[isize],
904 ) -> Result<Self> {
905 check_index_dtype(index_dtype)?;
906 check_gather_value_dtype(dtype)?;
907 Ok(Self {
908 dtype,
909 index_dtype,
910 plan: DynamicUpdateSlicePlan::compile(
911 operand_dims,
912 operand_strides,
913 start_dims,
914 start_strides,
915 update_dims,
916 update_strides,
917 dest_dims,
918 dest_strides,
919 )?,
920 })
921 }
922
923 #[inline]
924 pub fn dtype(&self) -> KernelDType {
925 self.dtype
926 }
927
928 #[inline]
929 pub fn index_dtype(&self) -> KernelDType {
930 self.index_dtype
931 }
932
933 #[inline]
934 pub fn plan(&self) -> &DynamicUpdateSlicePlan {
935 &self.plan
936 }
937
938 pub fn execute(
940 &self,
941 ctx: &ExecContext,
942 dest: &mut ErasedRawStridedMut<'_>,
943 operand: &ErasedRawStridedRef<'_>,
944 update: &ErasedRawStridedRef<'_>,
945 starts: &ErasedRawStridedRef<'_>,
946 ) -> Result<()> {
947 check_dtype(self.dtype, dest.dtype())?;
948 check_dtype(self.dtype, operand.dtype())?;
949 check_dtype(self.dtype, update.dtype())?;
950 check_dtype(self.index_dtype, starts.dtype())?;
951
952 let result = ctx.run(|| match self.dtype {
953 KernelDType::F32 => dispatch_dynamic_update_slice_index::<f32>(
954 &self.plan,
955 self.index_dtype,
956 dest,
957 &operand,
958 &update,
959 &starts,
960 ),
961 KernelDType::F64 => dispatch_dynamic_update_slice_index::<f64>(
962 &self.plan,
963 self.index_dtype,
964 dest,
965 operand,
966 update,
967 starts,
968 ),
969 KernelDType::I32 => dispatch_dynamic_update_slice_index::<i32>(
970 &self.plan,
971 self.index_dtype,
972 dest,
973 operand,
974 update,
975 starts,
976 ),
977 KernelDType::I64 => dispatch_dynamic_update_slice_index::<i64>(
978 &self.plan,
979 self.index_dtype,
980 dest,
981 operand,
982 update,
983 starts,
984 ),
985 KernelDType::Bool => dispatch_dynamic_update_slice_index::<bool>(
986 &self.plan,
987 self.index_dtype,
988 dest,
989 operand,
990 update,
991 starts,
992 ),
993 KernelDType::C32 => dispatch_dynamic_update_slice_index::<Complex32>(
994 &self.plan,
995 self.index_dtype,
996 dest,
997 operand,
998 update,
999 starts,
1000 ),
1001 KernelDType::C64 => dispatch_dynamic_update_slice_index::<Complex64>(
1002 &self.plan,
1003 self.index_dtype,
1004 dest,
1005 operand,
1006 update,
1007 starts,
1008 ),
1009 _ => Err(StridedError::UnsupportedDType {
1010 dtype: self.dtype.label(),
1011 }),
1012 });
1013 result
1014 }
1015
1016 pub fn execute_uninit(
1023 &self,
1024 ctx: &ExecContext,
1025 dest: &mut ErasedRawStridedUninitMut<'_>,
1026 operand: &ErasedRawStridedPtr<'_>,
1027 update: &ErasedRawStridedPtr<'_>,
1028 starts: &ErasedRawStridedPtr<'_>,
1029 ) -> Result<()> {
1030 check_dtype(self.dtype, dest.dtype())?;
1031 check_dtype(self.dtype, operand.dtype())?;
1032 check_dtype(self.dtype, update.dtype())?;
1033 check_dtype(self.index_dtype, starts.dtype())?;
1034 validate_uninit_no_overlap(dest, operand, 0)?;
1035 validate_uninit_no_overlap(dest, update, 1)?;
1036 validate_uninit_no_overlap(dest, starts, 2)?;
1037 let operand = &unsafe { operand.try_as_ref_after_no_overlap() }?;
1039 let update = &unsafe { update.try_as_ref_after_no_overlap() }?;
1041 let starts = &unsafe { starts.try_as_ref_after_no_overlap() }?;
1043 let run = |dest: &mut ErasedRawStridedUninitMut<'_>| match self.dtype {
1044 KernelDType::F32 => execute_dynamic_update_uninit_dispatch::<f32>(
1045 &self.plan,
1046 self.index_dtype,
1047 dest,
1048 operand,
1049 update,
1050 starts,
1051 ),
1052 KernelDType::F64 => execute_dynamic_update_uninit_dispatch::<f64>(
1053 &self.plan,
1054 self.index_dtype,
1055 dest,
1056 operand,
1057 update,
1058 starts,
1059 ),
1060 KernelDType::I32 => execute_dynamic_update_uninit_dispatch::<i32>(
1061 &self.plan,
1062 self.index_dtype,
1063 dest,
1064 operand,
1065 update,
1066 starts,
1067 ),
1068 KernelDType::I64 => execute_dynamic_update_uninit_dispatch::<i64>(
1069 &self.plan,
1070 self.index_dtype,
1071 dest,
1072 operand,
1073 update,
1074 starts,
1075 ),
1076 KernelDType::Bool => execute_dynamic_update_uninit_dispatch::<bool>(
1077 &self.plan,
1078 self.index_dtype,
1079 dest,
1080 operand,
1081 update,
1082 starts,
1083 ),
1084 KernelDType::C32 => execute_dynamic_update_uninit_dispatch::<Complex32>(
1085 &self.plan,
1086 self.index_dtype,
1087 dest,
1088 operand,
1089 update,
1090 starts,
1091 ),
1092 KernelDType::C64 => execute_dynamic_update_uninit_dispatch::<Complex64>(
1093 &self.plan,
1094 self.index_dtype,
1095 dest,
1096 operand,
1097 update,
1098 starts,
1099 ),
1100 _ => Err(StridedError::UnsupportedDType {
1101 dtype: self.dtype.label(),
1102 }),
1103 };
1104 if ctx.is_serial() {
1105 run(dest)
1106 } else {
1107 ctx.run(|| run(dest))
1108 }
1109 }
1110}
1111
1112impl ErasedScatterPlan {
1113 #[allow(clippy::too_many_arguments)]
1115 pub fn compile(
1116 dtype: KernelDType,
1117 index_dtype: KernelDType,
1118 operand_dims: &[usize],
1119 operand_strides: &[isize],
1120 index_dims: &[usize],
1121 index_strides: &[isize],
1122 update_dims: &[usize],
1123 update_strides: &[isize],
1124 dest_dims: &[usize],
1125 dest_strides: &[isize],
1126 spec: ScatterSpec,
1127 ) -> Result<Self> {
1128 check_index_dtype(index_dtype)?;
1129 check_scatter_value_dtype(dtype)?;
1130 Ok(Self {
1131 dtype,
1132 index_dtype,
1133 plan: ScatterPlan::compile(
1134 operand_dims,
1135 operand_strides,
1136 index_dims,
1137 index_strides,
1138 update_dims,
1139 update_strides,
1140 dest_dims,
1141 dest_strides,
1142 spec,
1143 )?,
1144 })
1145 }
1146
1147 #[inline]
1148 pub fn dtype(&self) -> KernelDType {
1149 self.dtype
1150 }
1151
1152 #[inline]
1153 pub fn index_dtype(&self) -> KernelDType {
1154 self.index_dtype
1155 }
1156
1157 #[inline]
1158 pub fn plan(&self) -> &ScatterPlan {
1159 &self.plan
1160 }
1161
1162 pub fn execute(
1164 &self,
1165 ctx: &ExecContext,
1166 dest: &mut ErasedRawStridedMut<'_>,
1167 operand: &ErasedRawStridedRef<'_>,
1168 scatter_indices: &ErasedRawStridedRef<'_>,
1169 updates: &ErasedRawStridedRef<'_>,
1170 ) -> Result<()> {
1171 check_dtype(self.dtype, dest.dtype())?;
1172 check_dtype(self.dtype, operand.dtype())?;
1173 check_dtype(self.dtype, updates.dtype())?;
1174 check_dtype(self.index_dtype, scatter_indices.dtype())?;
1175
1176 let result = ctx.run(|| match self.dtype {
1177 KernelDType::F32 => dispatch_scatter_index::<f32>(
1178 &self.plan,
1179 self.index_dtype,
1180 dest,
1181 &operand,
1182 &scatter_indices,
1183 &updates,
1184 ),
1185 KernelDType::F64 => dispatch_scatter_index::<f64>(
1186 &self.plan,
1187 self.index_dtype,
1188 dest,
1189 operand,
1190 scatter_indices,
1191 updates,
1192 ),
1193 KernelDType::I32 => dispatch_scatter_index::<i32>(
1194 &self.plan,
1195 self.index_dtype,
1196 dest,
1197 operand,
1198 scatter_indices,
1199 updates,
1200 ),
1201 KernelDType::I64 => dispatch_scatter_index::<i64>(
1202 &self.plan,
1203 self.index_dtype,
1204 dest,
1205 operand,
1206 scatter_indices,
1207 updates,
1208 ),
1209 KernelDType::C32 => dispatch_scatter_index::<Complex32>(
1210 &self.plan,
1211 self.index_dtype,
1212 dest,
1213 operand,
1214 scatter_indices,
1215 updates,
1216 ),
1217 KernelDType::C64 => dispatch_scatter_index::<Complex64>(
1218 &self.plan,
1219 self.index_dtype,
1220 dest,
1221 operand,
1222 scatter_indices,
1223 updates,
1224 ),
1225 _ => Err(StridedError::UnsupportedDType {
1226 dtype: self.dtype.label(),
1227 }),
1228 });
1229 result
1230 }
1231
1232 pub fn execute_uninit(
1239 &self,
1240 ctx: &ExecContext,
1241 dest: &mut ErasedRawStridedUninitMut<'_>,
1242 operand: &ErasedRawStridedPtr<'_>,
1243 scatter_indices: &ErasedRawStridedPtr<'_>,
1244 updates: &ErasedRawStridedPtr<'_>,
1245 ) -> Result<()> {
1246 check_dtype(self.dtype, dest.dtype())?;
1247 check_dtype(self.dtype, operand.dtype())?;
1248 check_dtype(self.dtype, updates.dtype())?;
1249 check_dtype(self.index_dtype, scatter_indices.dtype())?;
1250 validate_uninit_no_overlap(dest, operand, 0)?;
1251 validate_uninit_no_overlap(dest, scatter_indices, 1)?;
1252 validate_uninit_no_overlap(dest, updates, 2)?;
1253 let operand = &unsafe { operand.try_as_ref_after_no_overlap() }?;
1255 let scatter_indices = &unsafe { scatter_indices.try_as_ref_after_no_overlap() }?;
1257 let updates = &unsafe { updates.try_as_ref_after_no_overlap() }?;
1259 let run = |dest: &mut ErasedRawStridedUninitMut<'_>| match self.dtype {
1260 KernelDType::F32 => execute_scatter_uninit_dispatch::<f32>(
1261 &self.plan,
1262 self.index_dtype,
1263 dest,
1264 operand,
1265 scatter_indices,
1266 updates,
1267 add_values::<f32>,
1268 ),
1269 KernelDType::F64 => execute_scatter_uninit_dispatch::<f64>(
1270 &self.plan,
1271 self.index_dtype,
1272 dest,
1273 operand,
1274 scatter_indices,
1275 updates,
1276 add_values::<f64>,
1277 ),
1278 KernelDType::I32 => execute_scatter_uninit_dispatch::<i32>(
1279 &self.plan,
1280 self.index_dtype,
1281 dest,
1282 operand,
1283 scatter_indices,
1284 updates,
1285 i32::wrapping_add,
1286 ),
1287 KernelDType::I64 => execute_scatter_uninit_dispatch::<i64>(
1288 &self.plan,
1289 self.index_dtype,
1290 dest,
1291 operand,
1292 scatter_indices,
1293 updates,
1294 i64::wrapping_add,
1295 ),
1296 KernelDType::C32 => execute_scatter_uninit_dispatch::<Complex32>(
1297 &self.plan,
1298 self.index_dtype,
1299 dest,
1300 operand,
1301 scatter_indices,
1302 updates,
1303 add_values::<Complex32>,
1304 ),
1305 KernelDType::C64 => execute_scatter_uninit_dispatch::<Complex64>(
1306 &self.plan,
1307 self.index_dtype,
1308 dest,
1309 operand,
1310 scatter_indices,
1311 updates,
1312 add_values::<Complex64>,
1313 ),
1314 _ => Err(StridedError::UnsupportedDType {
1315 dtype: self.dtype.label(),
1316 }),
1317 };
1318 if ctx.is_serial() {
1319 run(dest)
1320 } else {
1321 ctx.run(|| run(dest))
1322 }
1323 }
1324}
1325
1326fn add_values<T: Add<Output = T>>(lhs: T, rhs: T) -> T {
1327 lhs + rhs
1328}
1329
1330fn execute_one_shot_map<T: OneShotScalar>(
1331 op: ErasedMapOp,
1332 dest: &mut ErasedRawStridedMut<'_>,
1333 input: &ErasedRawStridedRef<'_>,
1334) -> Result<()> {
1335 if !T::supports_map(op) {
1336 return Err(StridedError::UnsupportedOp {
1337 op: op.label(),
1338 dtype: T::one_shot_dtype_label(),
1339 });
1340 }
1341 let validated = strided_basic::execution::validate_destination_layout_without_alloc(
1342 dest.dims(),
1343 dest.strides(),
1344 )?;
1345 strided_basic::execution::ensure_same_shape(dest.dims(), input.dims())?;
1346 if dest.dims().contains(&0) {
1347 return Ok(());
1348 }
1349
1350 let dest_dims = dest.dims();
1351 let dest_strides = dest.strides();
1352 let dest_offset = dest.offset();
1353 let dest_data = dest.data_as_mut::<T>()?;
1354 let mut dest =
1355 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
1356 let input = erased_raw_ref::<T>(input)?;
1357
1358 macro_rules! replay {
1362 ($op:ident) => {
1363 unsafe {
1365 strided_basic::execution::map_raw_into_validated::<T, T, Identity>(
1366 &mut dest,
1367 &input,
1368 |value| T::map(ErasedMapOp::$op, value),
1369 validated,
1370 )
1371 }
1372 };
1373 }
1374 match op {
1375 ErasedMapOp::Negate => replay!(Negate),
1376 ErasedMapOp::Conj => replay!(Conj),
1377 ErasedMapOp::Abs => replay!(Abs),
1378 ErasedMapOp::Sign => replay!(Sign),
1379 }
1380}
1381
1382fn execute_one_shot_map_with<D, A>(
1383 dest: &mut ErasedRawStridedMut<'_>,
1384 input: &ErasedRawStridedRef<'_>,
1385 map: impl Fn(A) -> D + crate::MaybeSync,
1386) -> Result<()>
1387where
1388 D: Copy + crate::MaybeSendSync + KernelStorageElement,
1389 A: Copy + crate::MaybeSendSync + KernelStorageElement,
1390{
1391 let validated = strided_basic::execution::validate_destination_layout_without_alloc(
1392 dest.dims(),
1393 dest.strides(),
1394 )?;
1395 strided_basic::execution::ensure_same_shape(dest.dims(), input.dims())?;
1396 if dest.dims().contains(&0) {
1397 return Ok(());
1398 }
1399 let dest_dims = dest.dims();
1400 let dest_strides = dest.strides();
1401 let dest_offset = dest.offset();
1402 let dest_data = dest.data_as_mut::<D>()?;
1403 let mut dest =
1404 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
1405 let input = erased_raw_ref::<A>(input)?;
1406 unsafe {
1408 strided_basic::execution::map_raw_into_validated::<D, A, Identity>(
1409 &mut dest, &input, map, validated,
1410 )
1411 }
1412}
1413
1414fn execute_one_shot_zip<T: OneShotScalar>(
1415 op: ErasedZipOp,
1416 dest: &mut ErasedRawStridedMut<'_>,
1417 lhs: &ErasedRawStridedRef<'_>,
1418 rhs: &ErasedRawStridedRef<'_>,
1419) -> Result<()> {
1420 if !T::supports_zip(op) {
1421 return Err(StridedError::UnsupportedOp {
1422 op: op.label(),
1423 dtype: T::one_shot_dtype_label(),
1424 });
1425 }
1426 let validated = strided_basic::execution::validate_destination_layout_without_alloc(
1427 dest.dims(),
1428 dest.strides(),
1429 )?;
1430 strided_basic::execution::ensure_same_shape(dest.dims(), lhs.dims())?;
1431 strided_basic::execution::ensure_same_shape(dest.dims(), rhs.dims())?;
1432 if dest.dims().contains(&0) {
1433 return Ok(());
1434 }
1435
1436 let lhs = erased_raw_ref::<T>(lhs)?;
1437 let rhs = erased_raw_ref::<T>(rhs)?;
1438 if matches!(op, ErasedZipOp::Divide | ErasedZipOp::Remainder)
1439 && T::INTEGER
1440 && raw_any(&rhs, T::is_zero)?
1441 {
1442 return Err(StridedError::IntegerDivisionByZero { op: op.label() });
1443 }
1444 let dest_dims = dest.dims();
1445 let dest_strides = dest.strides();
1446 let dest_offset = dest.offset();
1447 let dest_data = dest.data_as_mut::<T>()?;
1448 let mut dest =
1449 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
1450 macro_rules! replay {
1453 ($op:ident) => {
1454 unsafe {
1456 strided_basic::execution::zip_map2_raw_into_validated::<
1457 T,
1458 T,
1459 T,
1460 Identity,
1461 Identity,
1462 >(
1463 &mut dest,
1464 &lhs,
1465 &rhs,
1466 |lhs, rhs| T::zip(ErasedZipOp::$op, lhs, rhs),
1467 validated,
1468 )
1469 }
1470 };
1471 }
1472 match op {
1473 ErasedZipOp::Add => replay!(Add),
1474 ErasedZipOp::Subtract => replay!(Subtract),
1475 ErasedZipOp::Multiply => replay!(Multiply),
1476 ErasedZipOp::Divide => replay!(Divide),
1477 ErasedZipOp::Remainder => replay!(Remainder),
1478 ErasedZipOp::Maximum => replay!(Maximum),
1479 ErasedZipOp::Minimum => replay!(Minimum),
1480 }
1481}
1482
1483fn validate_no_overlap(
1484 dest: &ErasedRawStridedMut<'_>,
1485 input: &ErasedRawStridedPtr<'_>,
1486 input_index: usize,
1487) -> Result<()> {
1488 if input.overlaps_mut(dest)? {
1489 Err(StridedError::OverlappingInputOutput { input: input_index })
1490 } else {
1491 Ok(())
1492 }
1493}
1494
1495trait OneShotScalar: Copy + crate::MaybeSendSync + KernelStorageElement + 'static {
1496 const INTEGER: bool = false;
1497 fn is_zero(_value: Self) -> bool {
1498 false
1499 }
1500 fn one_shot_dtype_label() -> &'static str;
1501 fn supports_map(op: ErasedMapOp) -> bool;
1502 fn supports_zip(op: ErasedZipOp) -> bool;
1503 fn map(op: ErasedMapOp, value: Self) -> Self;
1504 fn zip(op: ErasedZipOp, lhs: Self, rhs: Self) -> Self;
1505}
1506
1507macro_rules! impl_real_one_shot_scalar {
1508 ($ty:ty, $label:literal) => {
1509 impl OneShotScalar for $ty {
1510 fn one_shot_dtype_label() -> &'static str {
1511 $label
1512 }
1513
1514 fn supports_map(_op: ErasedMapOp) -> bool {
1515 true
1516 }
1517
1518 fn supports_zip(_op: ErasedZipOp) -> bool {
1519 true
1520 }
1521
1522 #[inline(always)]
1523 fn map(op: ErasedMapOp, value: Self) -> Self {
1524 match op {
1525 ErasedMapOp::Negate => -value,
1526 ErasedMapOp::Conj => value,
1527 ErasedMapOp::Abs => value.abs(),
1528 ErasedMapOp::Sign => {
1529 if value == 0.0 {
1530 0.0
1531 } else {
1532 value.signum()
1533 }
1534 }
1535 }
1536 }
1537
1538 #[inline(always)]
1539 fn zip(op: ErasedZipOp, lhs: Self, rhs: Self) -> Self {
1540 match op {
1541 ErasedZipOp::Add => lhs + rhs,
1542 ErasedZipOp::Subtract => lhs - rhs,
1543 ErasedZipOp::Multiply => lhs * rhs,
1544 ErasedZipOp::Divide => lhs / rhs,
1545 ErasedZipOp::Remainder => lhs % rhs,
1546 ErasedZipOp::Maximum => {
1551 let picked = if lhs >= rhs { lhs } else { rhs };
1552 if lhs.is_nan() | rhs.is_nan() {
1553 <$ty>::NAN
1554 } else {
1555 picked
1556 }
1557 }
1558 ErasedZipOp::Minimum => {
1559 let picked = if lhs <= rhs { lhs } else { rhs };
1560 if lhs.is_nan() | rhs.is_nan() {
1561 <$ty>::NAN
1562 } else {
1563 picked
1564 }
1565 }
1566 }
1567 }
1568 }
1569 };
1570}
1571
1572macro_rules! impl_integer_one_shot_scalar {
1573 ($ty:ty, $label:literal) => {
1574 impl OneShotScalar for $ty {
1575 const INTEGER: bool = true;
1576
1577 fn is_zero(value: Self) -> bool {
1578 value == 0
1579 }
1580 fn one_shot_dtype_label() -> &'static str {
1581 $label
1582 }
1583
1584 fn supports_map(_op: ErasedMapOp) -> bool {
1585 true
1586 }
1587
1588 fn supports_zip(_op: ErasedZipOp) -> bool {
1589 true
1590 }
1591
1592 #[inline(always)]
1593 fn map(op: ErasedMapOp, value: Self) -> Self {
1594 match op {
1595 ErasedMapOp::Negate => value.wrapping_neg(),
1596 ErasedMapOp::Conj => value,
1597 ErasedMapOp::Abs => value.wrapping_abs(),
1598 ErasedMapOp::Sign => value.signum(),
1599 }
1600 }
1601
1602 #[inline(always)]
1603 fn zip(op: ErasedZipOp, lhs: Self, rhs: Self) -> Self {
1604 match op {
1605 ErasedZipOp::Add => lhs.wrapping_add(rhs),
1606 ErasedZipOp::Subtract => lhs.wrapping_sub(rhs),
1607 ErasedZipOp::Multiply => lhs.wrapping_mul(rhs),
1608 ErasedZipOp::Maximum => lhs.max(rhs),
1609 ErasedZipOp::Minimum => lhs.min(rhs),
1610 ErasedZipOp::Divide => lhs.wrapping_div(rhs),
1611 ErasedZipOp::Remainder => lhs.wrapping_rem(rhs),
1612 }
1613 }
1614 }
1615 };
1616}
1617
1618macro_rules! impl_complex_one_shot_scalar {
1619 ($ty:ty, $label:literal) => {
1620 impl OneShotScalar for $ty {
1621 fn one_shot_dtype_label() -> &'static str {
1622 $label
1623 }
1624
1625 fn supports_map(_op: ErasedMapOp) -> bool {
1626 true
1627 }
1628
1629 fn supports_zip(op: ErasedZipOp) -> bool {
1630 !matches!(
1631 op,
1632 ErasedZipOp::Remainder | ErasedZipOp::Maximum | ErasedZipOp::Minimum
1633 )
1634 }
1635
1636 #[inline(always)]
1637 fn map(op: ErasedMapOp, value: Self) -> Self {
1638 match op {
1639 ErasedMapOp::Negate => -value,
1640 ErasedMapOp::Conj => value.conj(),
1641 ErasedMapOp::Abs => Self::new(value.norm(), 0.0),
1642 ErasedMapOp::Sign => {
1643 if value.re == 0.0 && value.im == 0.0 {
1644 Self::new(0.0, 0.0)
1645 } else {
1646 let norm = value.norm();
1651 Self::new(value.re / norm, value.im / norm)
1652 }
1653 }
1654 }
1655 }
1656
1657 #[inline(always)]
1658 fn zip(op: ErasedZipOp, lhs: Self, rhs: Self) -> Self {
1659 match op {
1660 ErasedZipOp::Add => lhs + rhs,
1661 ErasedZipOp::Subtract => lhs - rhs,
1662 ErasedZipOp::Multiply => lhs * rhs,
1663 ErasedZipOp::Divide => lhs / rhs,
1664 ErasedZipOp::Remainder | ErasedZipOp::Maximum | ErasedZipOp::Minimum => {
1665 unreachable!("unsupported complex one-shot op")
1666 }
1667 }
1668 }
1669 }
1670 };
1671}
1672
1673impl_real_one_shot_scalar!(f32, "f32");
1674
1675impl_real_one_shot_scalar!(f64, "f64");
1676
1677impl_integer_one_shot_scalar!(i32, "i32");
1678
1679impl_integer_one_shot_scalar!(i64, "i64");
1680
1681impl_complex_one_shot_scalar!(Complex32, "c32");
1682
1683impl_complex_one_shot_scalar!(Complex64, "c64");
1684
1685impl OneShotScalar for bool {
1686 fn one_shot_dtype_label() -> &'static str {
1687 "bool"
1688 }
1689
1690 fn supports_map(op: ErasedMapOp) -> bool {
1691 matches!(op, ErasedMapOp::Conj)
1692 }
1693
1694 fn supports_zip(_op: ErasedZipOp) -> bool {
1695 false
1696 }
1697
1698 fn map(op: ErasedMapOp, value: Self) -> Self {
1699 match op {
1700 ErasedMapOp::Conj => value,
1701 _ => unreachable!("unsupported bool one-shot op"),
1702 }
1703 }
1704
1705 fn zip(_op: ErasedZipOp, _lhs: Self, _rhs: Self) -> Self {
1706 unreachable!("unsupported bool one-shot op")
1707 }
1708}
1709
1710fn erased_raw_ref<'a, T: KernelStorageElement>(
1711 src: &'a ErasedRawStridedRef<'a>,
1712) -> Result<RawStridedRef<'a, T>> {
1713 let data = src.data_as::<T>()?;
1714 Ok(unsafe { RawStridedRef::new_unchecked(data, src.dims(), src.strides(), src.offset()) })
1715}
1716
1717fn map_output_dtype(dtype: KernelDType, op: ErasedMapOp) -> Result<KernelDType> {
1718 match (dtype, op) {
1719 (KernelDType::C32, ErasedMapOp::Abs) => Ok(KernelDType::F32),
1720 (KernelDType::C64, ErasedMapOp::Abs) => Ok(KernelDType::F64),
1721 (KernelDType::Bool, ErasedMapOp::Conj) => Ok(KernelDType::Bool),
1722 (KernelDType::Bool, _) => Err(StridedError::UnsupportedOp {
1723 op: op.label(),
1724 dtype: dtype.label(),
1725 }),
1726 _ => Ok(dtype),
1727 }
1728}
1729
1730fn raw_any<T: Copy>(
1731 src: &RawStridedRef<'_, T>,
1732 predicate: impl Fn(T) -> bool + Copy,
1733) -> Result<bool> {
1734 let total = src
1735 .dims()
1736 .iter()
1737 .try_fold(1usize, |total, &dim| total.checked_mul(dim))
1738 .ok_or(StridedError::OffsetOverflow)?;
1739 if total == 0 {
1740 return Ok(false);
1741 }
1742
1743 let rank = src.dims().len();
1744 if rank <= RAW_FUSED_RANK_LIMIT {
1745 let mut coordinates = [0usize; RAW_FUSED_RANK_LIMIT];
1746 let mut resets = [0isize; RAW_FUSED_RANK_LIMIT];
1747 raw_any_odometer(
1748 src,
1749 predicate,
1750 &mut coordinates[..rank],
1751 &mut resets[..rank],
1752 )
1753 } else {
1754 let mut coordinates = vec![0usize; rank];
1757 let mut resets = vec![0isize; rank];
1758 raw_any_odometer(src, predicate, &mut coordinates, &mut resets)
1759 }
1760}
1761
1762fn raw_any_odometer<T: Copy>(
1763 src: &RawStridedRef<'_, T>,
1764 predicate: impl Fn(T) -> bool + Copy,
1765 coordinates: &mut [usize],
1766 resets: &mut [isize],
1767) -> Result<bool> {
1768 let dims = src.dims();
1769 let strides = src.strides();
1770 for axis in 0..dims.len() {
1772 let last = isize::try_from(dims[axis] - 1).map_err(|_| StridedError::OffsetOverflow)?;
1773 resets[axis] = strides[axis]
1774 .checked_mul(last)
1775 .and_then(isize::checked_neg)
1776 .ok_or(StridedError::OffsetOverflow)?;
1777 }
1778
1779 let mut offset = src.offset();
1780 loop {
1781 if predicate(unsafe { *src.data().as_ptr().offset(offset) }) {
1783 return Ok(true);
1784 }
1785
1786 let mut axis = 0;
1787 while axis < dims.len() && coordinates[axis] == dims[axis] - 1 {
1788 axis += 1;
1789 }
1790 if axis == dims.len() {
1791 return Ok(false);
1792 }
1793 for reset_axis in 0..axis {
1794 coordinates[reset_axis] = 0;
1795 offset = offset
1796 .checked_add(resets[reset_axis])
1797 .ok_or(StridedError::OffsetOverflow)?;
1798 }
1799 coordinates[axis] = coordinates[axis]
1800 .checked_add(1)
1801 .ok_or(StridedError::OffsetOverflow)?;
1802 offset = offset
1803 .checked_add(strides[axis])
1804 .ok_or(StridedError::OffsetOverflow)?;
1805 }
1806}
1807
1808fn check_index_dtype(dtype: KernelDType) -> Result<()> {
1809 match dtype {
1810 KernelDType::I32 | KernelDType::I64 => Ok(()),
1811 _ => Err(StridedError::UnsupportedDType {
1812 dtype: dtype.label(),
1813 }),
1814 }
1815}
1816
1817fn check_gather_value_dtype(dtype: KernelDType) -> Result<()> {
1818 match dtype {
1819 KernelDType::F32
1820 | KernelDType::F64
1821 | KernelDType::I32
1822 | KernelDType::I64
1823 | KernelDType::Bool
1824 | KernelDType::C32
1825 | KernelDType::C64 => Ok(()),
1826 _ => Err(StridedError::UnsupportedDType {
1827 dtype: dtype.label(),
1828 }),
1829 }
1830}
1831
1832fn check_scatter_value_dtype(dtype: KernelDType) -> Result<()> {
1833 match dtype {
1834 KernelDType::F32
1835 | KernelDType::F64
1836 | KernelDType::I32
1837 | KernelDType::I64
1838 | KernelDType::C32
1839 | KernelDType::C64 => Ok(()),
1840 _ => Err(StridedError::UnsupportedDType {
1841 dtype: dtype.label(),
1842 }),
1843 }
1844}
1845
1846fn validate_scalar_bytes(dtype: KernelDType, bytes: &[u8]) -> Result<()> {
1847 let element_size = dtype.size_of();
1848 if bytes.len() != element_size {
1849 return Err(StridedError::ByteLengthMismatch {
1850 dtype: dtype.label(),
1851 byte_len: bytes.len(),
1852 element_size,
1853 });
1854 }
1855 if dtype.requires_valid_byte_values() {
1856 if let Some(&value) = bytes.iter().find(|&&value| value > 1) {
1857 return Err(StridedError::InvalidBoolByte { value });
1858 }
1859 }
1860 Ok(())
1861}
1862
1863fn execute_slice<T>(
1864 plan: &SlicePlan,
1865 dest: &mut ErasedRawStridedMut<'_>,
1866 operand: &ErasedRawStridedRef<'_>,
1867) -> Result<()>
1868where
1869 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1870{
1871 let operand_data = operand.data_as::<T>()?;
1872 let dest_dims = dest.dims();
1873 let dest_strides = dest.strides();
1874 let dest_offset = dest.offset();
1875 let dest_data = dest.data_as_mut::<T>()?;
1876 let operand_ref = unsafe {
1877 RawStridedRef::new_unchecked(
1878 operand_data,
1879 operand.dims(),
1880 operand.strides(),
1881 operand.offset(),
1882 )
1883 };
1884 let mut dest_ref =
1885 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
1886 plan.execute(&mut dest_ref, &operand_ref)
1887}
1888
1889fn execute_gather_uninit_dispatch<T>(
1890 plan: &GatherPlan,
1891 index_dtype: KernelDType,
1892 dest: &mut ErasedRawStridedUninitMut<'_>,
1893 operand: &ErasedRawStridedRef<'_>,
1894 start_indices: &ErasedRawStridedRef<'_>,
1895) -> Result<()>
1896where
1897 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1898{
1899 match index_dtype {
1900 KernelDType::I32 => {
1901 execute_gather_uninit::<T, i32>(plan, index_dtype, dest, operand, start_indices)
1902 }
1903 KernelDType::I64 => {
1904 execute_gather_uninit::<T, i64>(plan, index_dtype, dest, operand, start_indices)
1905 }
1906 _ => Err(StridedError::UnsupportedDType {
1907 dtype: index_dtype.label(),
1908 }),
1909 }
1910}
1911
1912fn execute_gather_uninit<T, I>(
1913 plan: &GatherPlan,
1914 _index_dtype: KernelDType,
1915 dest: &mut ErasedRawStridedUninitMut<'_>,
1916 operand: &ErasedRawStridedRef<'_>,
1917 start_indices: &ErasedRawStridedRef<'_>,
1918) -> Result<()>
1919where
1920 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1921 I: GatherIndex + KernelStorageElement,
1922{
1923 let operand_data = operand.data_as::<T>()?;
1924 let index_data = start_indices.data_as::<I>()?;
1925 let dest_dims = dest.dims();
1926 let dest_strides = dest.strides();
1927 let dest_offset = dest.offset();
1928 let dest_data = dest.data_as_uninit_mut::<T>()?;
1929 let operand_ref = unsafe {
1930 RawStridedRef::new_unchecked(
1931 operand_data,
1932 operand.dims(),
1933 operand.strides(),
1934 operand.offset(),
1935 )
1936 };
1937 let index_ref = unsafe {
1938 RawStridedRef::new_unchecked(
1939 index_data,
1940 start_indices.dims(),
1941 start_indices.strides(),
1942 start_indices.offset(),
1943 )
1944 };
1945 let mut dest_ref =
1946 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
1947 unsafe {
1949 strided_basic::execution::gather_into_uninit(plan, &mut dest_ref, &operand_ref, &index_ref)
1950 }
1951}
1952
1953fn execute_dynamic_slice_uninit_dispatch<T>(
1954 plan: &DynamicSlicePlan,
1955 index_dtype: KernelDType,
1956 dest: &mut ErasedRawStridedUninitMut<'_>,
1957 operand: &ErasedRawStridedRef<'_>,
1958 starts: &ErasedRawStridedRef<'_>,
1959) -> Result<()>
1960where
1961 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1962{
1963 match index_dtype {
1964 KernelDType::I32 => execute_dynamic_slice_uninit::<T, i32>(plan, dest, operand, starts),
1965 KernelDType::I64 => execute_dynamic_slice_uninit::<T, i64>(plan, dest, operand, starts),
1966 _ => Err(StridedError::UnsupportedDType {
1967 dtype: index_dtype.label(),
1968 }),
1969 }
1970}
1971
1972fn execute_dynamic_slice_uninit<T, I>(
1973 plan: &DynamicSlicePlan,
1974 dest: &mut ErasedRawStridedUninitMut<'_>,
1975 operand: &ErasedRawStridedRef<'_>,
1976 starts: &ErasedRawStridedRef<'_>,
1977) -> Result<()>
1978where
1979 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1980 I: GatherIndex + KernelStorageElement,
1981{
1982 let operand_data = operand.data_as::<T>()?;
1983 let starts_data = starts.data_as::<I>()?;
1984 let dest_dims = dest.dims();
1985 let dest_strides = dest.strides();
1986 let dest_offset = dest.offset();
1987 let dest_data = dest.data_as_uninit_mut::<T>()?;
1988 let operand_ref = unsafe {
1989 RawStridedRef::new_unchecked(
1990 operand_data,
1991 operand.dims(),
1992 operand.strides(),
1993 operand.offset(),
1994 )
1995 };
1996 let starts_ref = unsafe {
1997 RawStridedRef::new_unchecked(
1998 starts_data,
1999 starts.dims(),
2000 starts.strides(),
2001 starts.offset(),
2002 )
2003 };
2004 let mut dest_ref =
2005 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2006 unsafe {
2008 strided_basic::execution::dynamic_slice_into_uninit(
2009 plan,
2010 &mut dest_ref,
2011 &operand_ref,
2012 &starts_ref,
2013 )
2014 }
2015}
2016
2017fn execute_dynamic_update_uninit_dispatch<T>(
2018 plan: &DynamicUpdateSlicePlan,
2019 index_dtype: KernelDType,
2020 dest: &mut ErasedRawStridedUninitMut<'_>,
2021 operand: &ErasedRawStridedRef<'_>,
2022 update: &ErasedRawStridedRef<'_>,
2023 starts: &ErasedRawStridedRef<'_>,
2024) -> Result<()>
2025where
2026 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2027{
2028 match index_dtype {
2029 KernelDType::I32 => {
2030 execute_dynamic_update_uninit::<T, i32>(plan, dest, operand, update, starts)
2031 }
2032 KernelDType::I64 => {
2033 execute_dynamic_update_uninit::<T, i64>(plan, dest, operand, update, starts)
2034 }
2035 _ => Err(StridedError::UnsupportedDType {
2036 dtype: index_dtype.label(),
2037 }),
2038 }
2039}
2040
2041fn execute_dynamic_update_uninit<T, I>(
2042 plan: &DynamicUpdateSlicePlan,
2043 dest: &mut ErasedRawStridedUninitMut<'_>,
2044 operand: &ErasedRawStridedRef<'_>,
2045 update: &ErasedRawStridedRef<'_>,
2046 starts: &ErasedRawStridedRef<'_>,
2047) -> Result<()>
2048where
2049 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2050 I: GatherIndex + KernelStorageElement,
2051{
2052 let operand_data = operand.data_as::<T>()?;
2053 let update_data = update.data_as::<T>()?;
2054 let starts_data = starts.data_as::<I>()?;
2055 let dest_dims = dest.dims();
2056 let dest_strides = dest.strides();
2057 let dest_offset = dest.offset();
2058 let dest_data = dest.data_as_uninit_mut::<T>()?;
2059 let operand_ref = unsafe {
2060 RawStridedRef::new_unchecked(
2061 operand_data,
2062 operand.dims(),
2063 operand.strides(),
2064 operand.offset(),
2065 )
2066 };
2067 let update_ref = unsafe {
2068 RawStridedRef::new_unchecked(
2069 update_data,
2070 update.dims(),
2071 update.strides(),
2072 update.offset(),
2073 )
2074 };
2075 let starts_ref = unsafe {
2076 RawStridedRef::new_unchecked(
2077 starts_data,
2078 starts.dims(),
2079 starts.strides(),
2080 starts.offset(),
2081 )
2082 };
2083 let mut dest_ref =
2084 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2085 unsafe {
2087 strided_basic::execution::dynamic_update_into_uninit(
2088 plan,
2089 &mut dest_ref,
2090 &operand_ref,
2091 &update_ref,
2092 &starts_ref,
2093 )
2094 }
2095}
2096
2097fn execute_scatter_uninit_dispatch<T>(
2098 plan: &ScatterPlan,
2099 index_dtype: KernelDType,
2100 dest: &mut ErasedRawStridedUninitMut<'_>,
2101 operand: &ErasedRawStridedRef<'_>,
2102 scatter_indices: &ErasedRawStridedRef<'_>,
2103 updates: &ErasedRawStridedRef<'_>,
2104 combine: fn(T, T) -> T,
2105) -> Result<()>
2106where
2107 T: Copy + Add<Output = T> + crate::MaybeSendSync + KernelStorageElement,
2108{
2109 match index_dtype {
2110 KernelDType::I32 => {
2111 execute_scatter_uninit::<T, i32>(plan, dest, operand, scatter_indices, updates, combine)
2112 }
2113 KernelDType::I64 => {
2114 execute_scatter_uninit::<T, i64>(plan, dest, operand, scatter_indices, updates, combine)
2115 }
2116 _ => Err(StridedError::UnsupportedDType {
2117 dtype: index_dtype.label(),
2118 }),
2119 }
2120}
2121
2122fn execute_scatter_uninit<T, I>(
2123 plan: &ScatterPlan,
2124 dest: &mut ErasedRawStridedUninitMut<'_>,
2125 operand: &ErasedRawStridedRef<'_>,
2126 scatter_indices: &ErasedRawStridedRef<'_>,
2127 updates: &ErasedRawStridedRef<'_>,
2128 combine: fn(T, T) -> T,
2129) -> Result<()>
2130where
2131 T: Copy + Add<Output = T> + crate::MaybeSendSync + KernelStorageElement,
2132 I: GatherIndex + KernelStorageElement,
2133{
2134 let indices = scatter_indices;
2135 let operand_data = operand.data_as::<T>()?;
2136 let index_data = indices.data_as::<I>()?;
2137 let update_data = updates.data_as::<T>()?;
2138 let dest_dims = dest.dims();
2139 let dest_strides = dest.strides();
2140 let dest_offset = dest.offset();
2141 let dest_data = dest.data_as_uninit_mut::<T>()?;
2142 let operand_ref = unsafe {
2143 RawStridedRef::new_unchecked(
2144 operand_data,
2145 operand.dims(),
2146 operand.strides(),
2147 operand.offset(),
2148 )
2149 };
2150 let index_ref = unsafe {
2151 RawStridedRef::new_unchecked(
2152 index_data,
2153 indices.dims(),
2154 indices.strides(),
2155 indices.offset(),
2156 )
2157 };
2158 let update_ref = unsafe {
2159 RawStridedRef::new_unchecked(
2160 update_data,
2161 updates.dims(),
2162 updates.strides(),
2163 updates.offset(),
2164 )
2165 };
2166 let mut dest_ref =
2167 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2168 unsafe {
2170 strided_basic::execution::scatter_into_uninit(
2171 plan,
2172 &mut dest_ref,
2173 &operand_ref,
2174 &index_ref,
2175 &update_ref,
2176 combine,
2177 )
2178 }
2179}
2180
2181fn execute_slice_uninit<T>(
2182 plan: &SlicePlan,
2183 dest: &mut ErasedRawStridedUninitMut<'_>,
2184 operand: &ErasedRawStridedRef<'_>,
2185) -> Result<()>
2186where
2187 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2188{
2189 let operand_data = operand.data_as::<T>()?;
2190 let dest_dims = dest.dims();
2191 let dest_strides = dest.strides();
2192 let dest_offset = dest.offset();
2193 let dest_data = dest.data_as_uninit_mut::<T>()?;
2194 let operand_ref = unsafe {
2195 RawStridedRef::new_unchecked(
2196 operand_data,
2197 operand.dims(),
2198 operand.strides(),
2199 operand.offset(),
2200 )
2201 };
2202 let mut dest_ref =
2203 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2204 plan.execute_uninit(&mut dest_ref, &operand_ref)
2205}
2206
2207fn execute_reverse<T>(
2208 plan: &ReversePlan,
2209 dest: &mut ErasedRawStridedMut<'_>,
2210 operand: &ErasedRawStridedRef<'_>,
2211) -> Result<()>
2212where
2213 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2214{
2215 let operand_data = operand.data_as::<T>()?;
2216 let dest_dims = dest.dims();
2217 let dest_strides = dest.strides();
2218 let dest_offset = dest.offset();
2219 let dest_data = dest.data_as_mut::<T>()?;
2220 let operand_ref = unsafe {
2221 RawStridedRef::new_unchecked(
2222 operand_data,
2223 operand.dims(),
2224 operand.strides(),
2225 operand.offset(),
2226 )
2227 };
2228 let mut dest_ref =
2229 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2230 plan.execute(&mut dest_ref, &operand_ref)
2231}
2232
2233fn execute_reverse_uninit<T>(
2234 plan: &ReversePlan,
2235 dest: &mut ErasedRawStridedUninitMut<'_>,
2236 operand: &ErasedRawStridedRef<'_>,
2237) -> Result<()>
2238where
2239 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2240{
2241 let operand_data = operand.data_as::<T>()?;
2242 let dest_dims = dest.dims();
2243 let dest_strides = dest.strides();
2244 let dest_offset = dest.offset();
2245 let dest_data = dest.data_as_uninit_mut::<T>()?;
2246 let operand_ref = unsafe {
2247 RawStridedRef::new_unchecked(
2248 operand_data,
2249 operand.dims(),
2250 operand.strides(),
2251 operand.offset(),
2252 )
2253 };
2254 let mut dest_ref =
2255 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2256 plan.execute_uninit(&mut dest_ref, &operand_ref)
2257}
2258
2259fn execute_pad<T>(
2260 plan: &PadPlan,
2261 dest: &mut ErasedRawStridedMut<'_>,
2262 operand: &ErasedRawStridedRef<'_>,
2263 fill: &[u8],
2264) -> Result<()>
2265where
2266 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2267{
2268 let fill = read_unaligned_scalar::<T>(fill);
2269 let operand_data = operand.data_as::<T>()?;
2270 let dest_dims = dest.dims();
2271 let dest_strides = dest.strides();
2272 let dest_offset = dest.offset();
2273 let dest_data = dest.data_as_mut::<T>()?;
2274 let operand_ref = unsafe {
2275 RawStridedRef::new_unchecked(
2276 operand_data,
2277 operand.dims(),
2278 operand.strides(),
2279 operand.offset(),
2280 )
2281 };
2282 let mut dest_ref =
2283 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2284 plan.execute(&mut dest_ref, &operand_ref, fill)
2285}
2286
2287fn execute_pad_uninit<T>(
2288 plan: &PadPlan,
2289 dest: &mut ErasedRawStridedUninitMut<'_>,
2290 operand: &ErasedRawStridedRef<'_>,
2291 fill: &[u8],
2292) -> Result<()>
2293where
2294 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2295{
2296 let fill = read_unaligned_scalar::<T>(fill);
2297 let operand_data = operand.data_as::<T>()?;
2298 let dest_dims = dest.dims();
2299 let dest_strides = dest.strides();
2300 let dest_offset = dest.offset();
2301 let dest_data = dest.data_as_uninit_mut::<T>()?;
2302 let operand_ref = unsafe {
2303 RawStridedRef::new_unchecked(
2304 operand_data,
2305 operand.dims(),
2306 operand.strides(),
2307 operand.offset(),
2308 )
2309 };
2310 let mut dest_ref =
2311 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2312 plan.execute_uninit(&mut dest_ref, &operand_ref, fill)
2313}
2314
2315fn dispatch_gather_index<T>(
2316 plan: &GatherPlan,
2317 index_dtype: KernelDType,
2318 dest: &mut ErasedRawStridedMut<'_>,
2319 operand: &ErasedRawStridedRef<'_>,
2320 start_indices: &ErasedRawStridedRef<'_>,
2321) -> Result<()>
2322where
2323 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2324{
2325 match index_dtype {
2326 KernelDType::I32 => execute_gather::<T, i32>(plan, dest, operand, start_indices),
2327 KernelDType::I64 => execute_gather::<T, i64>(plan, dest, operand, start_indices),
2328 _ => Err(StridedError::UnsupportedDType {
2329 dtype: index_dtype.label(),
2330 }),
2331 }
2332}
2333
2334fn execute_gather<T, I>(
2335 plan: &GatherPlan,
2336 dest: &mut ErasedRawStridedMut<'_>,
2337 operand: &ErasedRawStridedRef<'_>,
2338 start_indices: &ErasedRawStridedRef<'_>,
2339) -> Result<()>
2340where
2341 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2342 I: GatherIndex + KernelStorageElement,
2343{
2344 let operand_data = operand.data_as::<T>()?;
2345 let index_data = start_indices.data_as::<I>()?;
2346 let dest_dims = dest.dims();
2347 let dest_strides = dest.strides();
2348 let dest_offset = dest.offset();
2349 let dest_data = dest.data_as_mut::<T>()?;
2350 let operand_ref = unsafe {
2351 RawStridedRef::new_unchecked(
2352 operand_data,
2353 operand.dims(),
2354 operand.strides(),
2355 operand.offset(),
2356 )
2357 };
2358 let index_ref = unsafe {
2359 RawStridedRef::new_unchecked(
2360 index_data,
2361 start_indices.dims(),
2362 start_indices.strides(),
2363 start_indices.offset(),
2364 )
2365 };
2366 let mut dest_ref =
2367 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2368 plan.execute(&mut dest_ref, &operand_ref, &index_ref)
2369}
2370
2371fn dispatch_dynamic_slice_index<T>(
2372 plan: &DynamicSlicePlan,
2373 index_dtype: KernelDType,
2374 dest: &mut ErasedRawStridedMut<'_>,
2375 operand: &ErasedRawStridedRef<'_>,
2376 starts: &ErasedRawStridedRef<'_>,
2377) -> Result<()>
2378where
2379 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2380{
2381 match index_dtype {
2382 KernelDType::I32 => execute_dynamic_slice::<T, i32>(plan, dest, operand, starts),
2383 KernelDType::I64 => execute_dynamic_slice::<T, i64>(plan, dest, operand, starts),
2384 _ => Err(StridedError::UnsupportedDType {
2385 dtype: index_dtype.label(),
2386 }),
2387 }
2388}
2389
2390fn execute_dynamic_slice<T, I>(
2391 plan: &DynamicSlicePlan,
2392 dest: &mut ErasedRawStridedMut<'_>,
2393 operand: &ErasedRawStridedRef<'_>,
2394 starts: &ErasedRawStridedRef<'_>,
2395) -> Result<()>
2396where
2397 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2398 I: GatherIndex + KernelStorageElement,
2399{
2400 let operand_data = operand.data_as::<T>()?;
2401 let start_data = starts.data_as::<I>()?;
2402 let dest_dims = dest.dims();
2403 let dest_strides = dest.strides();
2404 let dest_offset = dest.offset();
2405 let dest_data = dest.data_as_mut::<T>()?;
2406 let operand_ref = unsafe {
2407 RawStridedRef::new_unchecked(
2408 operand_data,
2409 operand.dims(),
2410 operand.strides(),
2411 operand.offset(),
2412 )
2413 };
2414 let start_ref = unsafe {
2415 RawStridedRef::new_unchecked(start_data, starts.dims(), starts.strides(), starts.offset())
2416 };
2417 let mut dest_ref =
2418 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2419 plan.execute(&mut dest_ref, &operand_ref, &start_ref)
2420}
2421
2422fn dispatch_dynamic_update_slice_index<T>(
2423 plan: &DynamicUpdateSlicePlan,
2424 index_dtype: KernelDType,
2425 dest: &mut ErasedRawStridedMut<'_>,
2426 operand: &ErasedRawStridedRef<'_>,
2427 update: &ErasedRawStridedRef<'_>,
2428 starts: &ErasedRawStridedRef<'_>,
2429) -> Result<()>
2430where
2431 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2432{
2433 match index_dtype {
2434 KernelDType::I32 => {
2435 execute_dynamic_update_slice::<T, i32>(plan, dest, operand, update, starts)
2436 }
2437 KernelDType::I64 => {
2438 execute_dynamic_update_slice::<T, i64>(plan, dest, operand, update, starts)
2439 }
2440 _ => Err(StridedError::UnsupportedDType {
2441 dtype: index_dtype.label(),
2442 }),
2443 }
2444}
2445
2446fn execute_dynamic_update_slice<T, I>(
2447 plan: &DynamicUpdateSlicePlan,
2448 dest: &mut ErasedRawStridedMut<'_>,
2449 operand: &ErasedRawStridedRef<'_>,
2450 update: &ErasedRawStridedRef<'_>,
2451 starts: &ErasedRawStridedRef<'_>,
2452) -> Result<()>
2453where
2454 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2455 I: GatherIndex + KernelStorageElement,
2456{
2457 let operand_data = operand.data_as::<T>()?;
2458 let update_data = update.data_as::<T>()?;
2459 let start_data = starts.data_as::<I>()?;
2460 let dest_dims = dest.dims();
2461 let dest_strides = dest.strides();
2462 let dest_offset = dest.offset();
2463 let dest_data = dest.data_as_mut::<T>()?;
2464 let operand_ref = unsafe {
2465 RawStridedRef::new_unchecked(
2466 operand_data,
2467 operand.dims(),
2468 operand.strides(),
2469 operand.offset(),
2470 )
2471 };
2472 let update_ref = unsafe {
2473 RawStridedRef::new_unchecked(
2474 update_data,
2475 update.dims(),
2476 update.strides(),
2477 update.offset(),
2478 )
2479 };
2480 let start_ref = unsafe {
2481 RawStridedRef::new_unchecked(start_data, starts.dims(), starts.strides(), starts.offset())
2482 };
2483 let mut dest_ref =
2484 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2485 plan.execute(&mut dest_ref, &operand_ref, &update_ref, &start_ref)
2486}
2487
2488fn dispatch_scatter_index<T>(
2489 plan: &ScatterPlan,
2490 index_dtype: KernelDType,
2491 dest: &mut ErasedRawStridedMut<'_>,
2492 operand: &ErasedRawStridedRef<'_>,
2493 scatter_indices: &ErasedRawStridedRef<'_>,
2494 updates: &ErasedRawStridedRef<'_>,
2495) -> Result<()>
2496where
2497 T: Copy + Add<Output = T> + crate::MaybeSendSync + KernelStorageElement,
2498{
2499 match index_dtype {
2500 KernelDType::I32 => {
2501 execute_scatter::<T, i32>(plan, dest, operand, scatter_indices, updates)
2502 }
2503 KernelDType::I64 => {
2504 execute_scatter::<T, i64>(plan, dest, operand, scatter_indices, updates)
2505 }
2506 _ => Err(StridedError::UnsupportedDType {
2507 dtype: index_dtype.label(),
2508 }),
2509 }
2510}
2511
2512fn execute_scatter<T, I>(
2513 plan: &ScatterPlan,
2514 dest: &mut ErasedRawStridedMut<'_>,
2515 operand: &ErasedRawStridedRef<'_>,
2516 scatter_indices: &ErasedRawStridedRef<'_>,
2517 updates: &ErasedRawStridedRef<'_>,
2518) -> Result<()>
2519where
2520 T: Copy + Add<Output = T> + crate::MaybeSendSync + KernelStorageElement,
2521 I: GatherIndex + KernelStorageElement,
2522{
2523 let operand_data = operand.data_as::<T>()?;
2524 let index_data = scatter_indices.data_as::<I>()?;
2525 let update_data = updates.data_as::<T>()?;
2526 let dest_dims = dest.dims();
2527 let dest_strides = dest.strides();
2528 let dest_offset = dest.offset();
2529 let dest_data = dest.data_as_mut::<T>()?;
2530 let operand_ref = unsafe {
2531 RawStridedRef::new_unchecked(
2532 operand_data,
2533 operand.dims(),
2534 operand.strides(),
2535 operand.offset(),
2536 )
2537 };
2538 let index_ref = unsafe {
2539 RawStridedRef::new_unchecked(
2540 index_data,
2541 scatter_indices.dims(),
2542 scatter_indices.strides(),
2543 scatter_indices.offset(),
2544 )
2545 };
2546 let update_ref = unsafe {
2547 RawStridedRef::new_unchecked(
2548 update_data,
2549 updates.dims(),
2550 updates.strides(),
2551 updates.offset(),
2552 )
2553 };
2554 let mut dest_ref =
2555 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2556 plan.execute(&mut dest_ref, &operand_ref, &index_ref, &update_ref)
2557}
2558
2559fn read_unaligned_scalar<T>(bytes: &[u8]) -> T
2560where
2561 T: Copy,
2562{
2563 unsafe { core::ptr::read_unaligned(bytes.as_ptr().cast::<T>()) }
2564}