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 #[inline(always)]
1507 fn clamp(x: Self, lo: Self, hi: Self) -> Self {
1508 let raised = Self::zip(ErasedZipOp::Maximum, lo, x);
1509 Self::zip(ErasedZipOp::Minimum, hi, raised)
1510 }
1511}
1512
1513macro_rules! impl_real_one_shot_scalar {
1514 ($ty:ty, $label:literal) => {
1515 impl OneShotScalar for $ty {
1516 fn one_shot_dtype_label() -> &'static str {
1517 $label
1518 }
1519
1520 fn supports_map(_op: ErasedMapOp) -> bool {
1521 true
1522 }
1523
1524 fn supports_zip(_op: ErasedZipOp) -> bool {
1525 true
1526 }
1527
1528 #[inline(always)]
1529 fn map(op: ErasedMapOp, value: Self) -> Self {
1530 match op {
1531 ErasedMapOp::Negate => -value,
1532 ErasedMapOp::Conj => value,
1533 ErasedMapOp::Abs => value.abs(),
1534 ErasedMapOp::Sign => {
1535 if value == 0.0 {
1536 0.0
1537 } else {
1538 value.signum()
1539 }
1540 }
1541 }
1542 }
1543
1544 #[inline(always)]
1545 fn zip(op: ErasedZipOp, lhs: Self, rhs: Self) -> Self {
1546 match op {
1547 ErasedZipOp::Add => lhs + rhs,
1548 ErasedZipOp::Subtract => lhs - rhs,
1549 ErasedZipOp::Multiply => lhs * rhs,
1550 ErasedZipOp::Divide => lhs / rhs,
1551 ErasedZipOp::Remainder => lhs % rhs,
1552 ErasedZipOp::Maximum => {
1557 let picked = if lhs >= rhs { lhs } else { rhs };
1558 if lhs.is_nan() | rhs.is_nan() {
1559 <$ty>::NAN
1560 } else {
1561 picked
1562 }
1563 }
1564 ErasedZipOp::Minimum => {
1565 let picked = if lhs <= rhs { lhs } else { rhs };
1566 if lhs.is_nan() | rhs.is_nan() {
1567 <$ty>::NAN
1568 } else {
1569 picked
1570 }
1571 }
1572 }
1573 }
1574
1575 #[inline(always)]
1580 fn clamp(x: Self, lo: Self, hi: Self) -> Self {
1581 let raised = if lo >= x { lo } else { x };
1582 let lowered = if hi <= raised { hi } else { raised };
1583 if x.is_nan() | lo.is_nan() | hi.is_nan() {
1584 <$ty>::NAN
1585 } else {
1586 lowered
1587 }
1588 }
1589 }
1590 };
1591}
1592
1593macro_rules! impl_integer_one_shot_scalar {
1594 ($ty:ty, $label:literal) => {
1595 impl OneShotScalar for $ty {
1596 const INTEGER: bool = true;
1597
1598 fn is_zero(value: Self) -> bool {
1599 value == 0
1600 }
1601 fn one_shot_dtype_label() -> &'static str {
1602 $label
1603 }
1604
1605 fn supports_map(_op: ErasedMapOp) -> bool {
1606 true
1607 }
1608
1609 fn supports_zip(_op: ErasedZipOp) -> bool {
1610 true
1611 }
1612
1613 #[inline(always)]
1614 fn map(op: ErasedMapOp, value: Self) -> Self {
1615 match op {
1616 ErasedMapOp::Negate => value.wrapping_neg(),
1617 ErasedMapOp::Conj => value,
1618 ErasedMapOp::Abs => value.wrapping_abs(),
1619 ErasedMapOp::Sign => value.signum(),
1620 }
1621 }
1622
1623 #[inline(always)]
1624 fn zip(op: ErasedZipOp, lhs: Self, rhs: Self) -> Self {
1625 match op {
1626 ErasedZipOp::Add => lhs.wrapping_add(rhs),
1627 ErasedZipOp::Subtract => lhs.wrapping_sub(rhs),
1628 ErasedZipOp::Multiply => lhs.wrapping_mul(rhs),
1629 ErasedZipOp::Maximum => lhs.max(rhs),
1630 ErasedZipOp::Minimum => lhs.min(rhs),
1631 ErasedZipOp::Divide => lhs.wrapping_div(rhs),
1632 ErasedZipOp::Remainder => lhs.wrapping_rem(rhs),
1633 }
1634 }
1635 }
1636 };
1637}
1638
1639macro_rules! impl_complex_one_shot_scalar {
1640 ($ty:ty, $label:literal, $div:path) => {
1641 impl OneShotScalar for $ty {
1642 fn one_shot_dtype_label() -> &'static str {
1643 $label
1644 }
1645
1646 fn supports_map(_op: ErasedMapOp) -> bool {
1647 true
1648 }
1649
1650 fn supports_zip(op: ErasedZipOp) -> bool {
1651 !matches!(
1652 op,
1653 ErasedZipOp::Remainder | ErasedZipOp::Maximum | ErasedZipOp::Minimum
1654 )
1655 }
1656
1657 #[inline(always)]
1658 fn map(op: ErasedMapOp, value: Self) -> Self {
1659 match op {
1660 ErasedMapOp::Negate => -value,
1661 ErasedMapOp::Conj => value.conj(),
1662 ErasedMapOp::Abs => Self::new(value.norm(), 0.0),
1663 ErasedMapOp::Sign => {
1664 if value.re == 0.0 && value.im == 0.0 {
1665 Self::new(0.0, 0.0)
1666 } else {
1667 let norm = value.norm();
1672 Self::new(value.re / norm, value.im / norm)
1673 }
1674 }
1675 }
1676 }
1677
1678 #[inline(always)]
1679 fn zip(op: ErasedZipOp, lhs: Self, rhs: Self) -> Self {
1680 match op {
1681 ErasedZipOp::Add => lhs + rhs,
1682 ErasedZipOp::Subtract => lhs - rhs,
1683 ErasedZipOp::Multiply => lhs * rhs,
1684 ErasedZipOp::Divide => $div(lhs, rhs),
1685 ErasedZipOp::Remainder | ErasedZipOp::Maximum | ErasedZipOp::Minimum => {
1686 unreachable!("unsupported complex one-shot op")
1687 }
1688 }
1689 }
1690 }
1691 };
1692}
1693
1694impl_real_one_shot_scalar!(f32, "f32");
1695
1696impl_real_one_shot_scalar!(f64, "f64");
1697
1698impl_integer_one_shot_scalar!(i32, "i32");
1699
1700impl_integer_one_shot_scalar!(i64, "i64");
1701
1702impl_complex_one_shot_scalar!(Complex32, "c32", strided_basic::robust_complex_divide_f32);
1703
1704impl_complex_one_shot_scalar!(Complex64, "c64", strided_basic::robust_complex_divide_f64);
1705
1706impl OneShotScalar for bool {
1707 fn one_shot_dtype_label() -> &'static str {
1708 "bool"
1709 }
1710
1711 fn supports_map(op: ErasedMapOp) -> bool {
1712 matches!(op, ErasedMapOp::Conj)
1713 }
1714
1715 fn supports_zip(_op: ErasedZipOp) -> bool {
1716 false
1717 }
1718
1719 fn map(op: ErasedMapOp, value: Self) -> Self {
1720 match op {
1721 ErasedMapOp::Conj => value,
1722 _ => unreachable!("unsupported bool one-shot op"),
1723 }
1724 }
1725
1726 fn zip(_op: ErasedZipOp, _lhs: Self, _rhs: Self) -> Self {
1727 unreachable!("unsupported bool one-shot op")
1728 }
1729}
1730
1731fn erased_raw_ref<'a, T: KernelStorageElement>(
1732 src: &'a ErasedRawStridedRef<'a>,
1733) -> Result<RawStridedRef<'a, T>> {
1734 let data = src.data_as::<T>()?;
1735 Ok(unsafe { RawStridedRef::new_unchecked(data, src.dims(), src.strides(), src.offset()) })
1736}
1737
1738fn map_output_dtype(dtype: KernelDType, op: ErasedMapOp) -> Result<KernelDType> {
1739 match (dtype, op) {
1740 (KernelDType::C32, ErasedMapOp::Abs) => Ok(KernelDType::F32),
1741 (KernelDType::C64, ErasedMapOp::Abs) => Ok(KernelDType::F64),
1742 (KernelDType::Bool, ErasedMapOp::Conj) => Ok(KernelDType::Bool),
1743 (KernelDType::Bool, _) => Err(StridedError::UnsupportedOp {
1744 op: op.label(),
1745 dtype: dtype.label(),
1746 }),
1747 _ => Ok(dtype),
1748 }
1749}
1750
1751fn raw_any<T: Copy>(
1752 src: &RawStridedRef<'_, T>,
1753 predicate: impl Fn(T) -> bool + Copy,
1754) -> Result<bool> {
1755 let total = src
1756 .dims()
1757 .iter()
1758 .try_fold(1usize, |total, &dim| total.checked_mul(dim))
1759 .ok_or(StridedError::OffsetOverflow)?;
1760 if total == 0 {
1761 return Ok(false);
1762 }
1763
1764 let rank = src.dims().len();
1765 if rank <= RAW_FUSED_RANK_LIMIT {
1766 let mut coordinates = [0usize; RAW_FUSED_RANK_LIMIT];
1767 let mut resets = [0isize; RAW_FUSED_RANK_LIMIT];
1768 raw_any_odometer(
1769 src,
1770 predicate,
1771 &mut coordinates[..rank],
1772 &mut resets[..rank],
1773 )
1774 } else {
1775 let mut coordinates = vec![0usize; rank];
1778 let mut resets = vec![0isize; rank];
1779 raw_any_odometer(src, predicate, &mut coordinates, &mut resets)
1780 }
1781}
1782
1783fn raw_any_odometer<T: Copy>(
1784 src: &RawStridedRef<'_, T>,
1785 predicate: impl Fn(T) -> bool + Copy,
1786 coordinates: &mut [usize],
1787 resets: &mut [isize],
1788) -> Result<bool> {
1789 let dims = src.dims();
1790 let strides = src.strides();
1791 for axis in 0..dims.len() {
1793 let last = isize::try_from(dims[axis] - 1).map_err(|_| StridedError::OffsetOverflow)?;
1794 resets[axis] = strides[axis]
1795 .checked_mul(last)
1796 .and_then(isize::checked_neg)
1797 .ok_or(StridedError::OffsetOverflow)?;
1798 }
1799
1800 let mut offset = src.offset();
1801 loop {
1802 if predicate(unsafe { *src.data().as_ptr().offset(offset) }) {
1804 return Ok(true);
1805 }
1806
1807 let mut axis = 0;
1808 while axis < dims.len() && coordinates[axis] == dims[axis] - 1 {
1809 axis += 1;
1810 }
1811 if axis == dims.len() {
1812 return Ok(false);
1813 }
1814 for reset_axis in 0..axis {
1815 coordinates[reset_axis] = 0;
1816 offset = offset
1817 .checked_add(resets[reset_axis])
1818 .ok_or(StridedError::OffsetOverflow)?;
1819 }
1820 coordinates[axis] = coordinates[axis]
1821 .checked_add(1)
1822 .ok_or(StridedError::OffsetOverflow)?;
1823 offset = offset
1824 .checked_add(strides[axis])
1825 .ok_or(StridedError::OffsetOverflow)?;
1826 }
1827}
1828
1829fn check_index_dtype(dtype: KernelDType) -> Result<()> {
1830 match dtype {
1831 KernelDType::I32 | KernelDType::I64 => Ok(()),
1832 _ => Err(StridedError::UnsupportedDType {
1833 dtype: dtype.label(),
1834 }),
1835 }
1836}
1837
1838fn check_gather_value_dtype(dtype: KernelDType) -> Result<()> {
1839 match dtype {
1840 KernelDType::F32
1841 | KernelDType::F64
1842 | KernelDType::I32
1843 | KernelDType::I64
1844 | KernelDType::Bool
1845 | KernelDType::C32
1846 | KernelDType::C64 => Ok(()),
1847 _ => Err(StridedError::UnsupportedDType {
1848 dtype: dtype.label(),
1849 }),
1850 }
1851}
1852
1853fn check_scatter_value_dtype(dtype: KernelDType) -> Result<()> {
1854 match dtype {
1855 KernelDType::F32
1856 | KernelDType::F64
1857 | KernelDType::I32
1858 | KernelDType::I64
1859 | KernelDType::C32
1860 | KernelDType::C64 => Ok(()),
1861 _ => Err(StridedError::UnsupportedDType {
1862 dtype: dtype.label(),
1863 }),
1864 }
1865}
1866
1867fn validate_scalar_bytes(dtype: KernelDType, bytes: &[u8]) -> Result<()> {
1868 let element_size = dtype.size_of();
1869 if bytes.len() != element_size {
1870 return Err(StridedError::ByteLengthMismatch {
1871 dtype: dtype.label(),
1872 byte_len: bytes.len(),
1873 element_size,
1874 });
1875 }
1876 if dtype.requires_valid_byte_values() {
1877 if let Some(&value) = bytes.iter().find(|&&value| value > 1) {
1878 return Err(StridedError::InvalidBoolByte { value });
1879 }
1880 }
1881 Ok(())
1882}
1883
1884fn execute_slice<T>(
1885 plan: &SlicePlan,
1886 dest: &mut ErasedRawStridedMut<'_>,
1887 operand: &ErasedRawStridedRef<'_>,
1888) -> Result<()>
1889where
1890 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1891{
1892 let operand_data = operand.data_as::<T>()?;
1893 let dest_dims = dest.dims();
1894 let dest_strides = dest.strides();
1895 let dest_offset = dest.offset();
1896 let dest_data = dest.data_as_mut::<T>()?;
1897 let operand_ref = unsafe {
1898 RawStridedRef::new_unchecked(
1899 operand_data,
1900 operand.dims(),
1901 operand.strides(),
1902 operand.offset(),
1903 )
1904 };
1905 let mut dest_ref =
1906 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
1907 plan.execute(&mut dest_ref, &operand_ref)
1908}
1909
1910fn execute_gather_uninit_dispatch<T>(
1911 plan: &GatherPlan,
1912 index_dtype: KernelDType,
1913 dest: &mut ErasedRawStridedUninitMut<'_>,
1914 operand: &ErasedRawStridedRef<'_>,
1915 start_indices: &ErasedRawStridedRef<'_>,
1916) -> Result<()>
1917where
1918 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1919{
1920 match index_dtype {
1921 KernelDType::I32 => {
1922 execute_gather_uninit::<T, i32>(plan, index_dtype, dest, operand, start_indices)
1923 }
1924 KernelDType::I64 => {
1925 execute_gather_uninit::<T, i64>(plan, index_dtype, dest, operand, start_indices)
1926 }
1927 _ => Err(StridedError::UnsupportedDType {
1928 dtype: index_dtype.label(),
1929 }),
1930 }
1931}
1932
1933fn execute_gather_uninit<T, I>(
1934 plan: &GatherPlan,
1935 _index_dtype: KernelDType,
1936 dest: &mut ErasedRawStridedUninitMut<'_>,
1937 operand: &ErasedRawStridedRef<'_>,
1938 start_indices: &ErasedRawStridedRef<'_>,
1939) -> Result<()>
1940where
1941 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1942 I: GatherIndex + KernelStorageElement,
1943{
1944 let operand_data = operand.data_as::<T>()?;
1945 let index_data = start_indices.data_as::<I>()?;
1946 let dest_dims = dest.dims();
1947 let dest_strides = dest.strides();
1948 let dest_offset = dest.offset();
1949 let dest_data = dest.data_as_uninit_mut::<T>()?;
1950 let operand_ref = unsafe {
1951 RawStridedRef::new_unchecked(
1952 operand_data,
1953 operand.dims(),
1954 operand.strides(),
1955 operand.offset(),
1956 )
1957 };
1958 let index_ref = unsafe {
1959 RawStridedRef::new_unchecked(
1960 index_data,
1961 start_indices.dims(),
1962 start_indices.strides(),
1963 start_indices.offset(),
1964 )
1965 };
1966 let mut dest_ref =
1967 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
1968 unsafe {
1970 strided_basic::execution::gather_into_uninit(plan, &mut dest_ref, &operand_ref, &index_ref)
1971 }
1972}
1973
1974fn execute_dynamic_slice_uninit_dispatch<T>(
1975 plan: &DynamicSlicePlan,
1976 index_dtype: KernelDType,
1977 dest: &mut ErasedRawStridedUninitMut<'_>,
1978 operand: &ErasedRawStridedRef<'_>,
1979 starts: &ErasedRawStridedRef<'_>,
1980) -> Result<()>
1981where
1982 T: Copy + crate::MaybeSendSync + KernelStorageElement,
1983{
1984 match index_dtype {
1985 KernelDType::I32 => execute_dynamic_slice_uninit::<T, i32>(plan, dest, operand, starts),
1986 KernelDType::I64 => execute_dynamic_slice_uninit::<T, i64>(plan, dest, operand, starts),
1987 _ => Err(StridedError::UnsupportedDType {
1988 dtype: index_dtype.label(),
1989 }),
1990 }
1991}
1992
1993fn execute_dynamic_slice_uninit<T, I>(
1994 plan: &DynamicSlicePlan,
1995 dest: &mut ErasedRawStridedUninitMut<'_>,
1996 operand: &ErasedRawStridedRef<'_>,
1997 starts: &ErasedRawStridedRef<'_>,
1998) -> Result<()>
1999where
2000 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2001 I: GatherIndex + KernelStorageElement,
2002{
2003 let operand_data = operand.data_as::<T>()?;
2004 let starts_data = starts.data_as::<I>()?;
2005 let dest_dims = dest.dims();
2006 let dest_strides = dest.strides();
2007 let dest_offset = dest.offset();
2008 let dest_data = dest.data_as_uninit_mut::<T>()?;
2009 let operand_ref = unsafe {
2010 RawStridedRef::new_unchecked(
2011 operand_data,
2012 operand.dims(),
2013 operand.strides(),
2014 operand.offset(),
2015 )
2016 };
2017 let starts_ref = unsafe {
2018 RawStridedRef::new_unchecked(
2019 starts_data,
2020 starts.dims(),
2021 starts.strides(),
2022 starts.offset(),
2023 )
2024 };
2025 let mut dest_ref =
2026 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2027 unsafe {
2029 strided_basic::execution::dynamic_slice_into_uninit(
2030 plan,
2031 &mut dest_ref,
2032 &operand_ref,
2033 &starts_ref,
2034 )
2035 }
2036}
2037
2038fn execute_dynamic_update_uninit_dispatch<T>(
2039 plan: &DynamicUpdateSlicePlan,
2040 index_dtype: KernelDType,
2041 dest: &mut ErasedRawStridedUninitMut<'_>,
2042 operand: &ErasedRawStridedRef<'_>,
2043 update: &ErasedRawStridedRef<'_>,
2044 starts: &ErasedRawStridedRef<'_>,
2045) -> Result<()>
2046where
2047 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2048{
2049 match index_dtype {
2050 KernelDType::I32 => {
2051 execute_dynamic_update_uninit::<T, i32>(plan, dest, operand, update, starts)
2052 }
2053 KernelDType::I64 => {
2054 execute_dynamic_update_uninit::<T, i64>(plan, dest, operand, update, starts)
2055 }
2056 _ => Err(StridedError::UnsupportedDType {
2057 dtype: index_dtype.label(),
2058 }),
2059 }
2060}
2061
2062fn execute_dynamic_update_uninit<T, I>(
2063 plan: &DynamicUpdateSlicePlan,
2064 dest: &mut ErasedRawStridedUninitMut<'_>,
2065 operand: &ErasedRawStridedRef<'_>,
2066 update: &ErasedRawStridedRef<'_>,
2067 starts: &ErasedRawStridedRef<'_>,
2068) -> Result<()>
2069where
2070 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2071 I: GatherIndex + KernelStorageElement,
2072{
2073 let operand_data = operand.data_as::<T>()?;
2074 let update_data = update.data_as::<T>()?;
2075 let starts_data = starts.data_as::<I>()?;
2076 let dest_dims = dest.dims();
2077 let dest_strides = dest.strides();
2078 let dest_offset = dest.offset();
2079 let dest_data = dest.data_as_uninit_mut::<T>()?;
2080 let operand_ref = unsafe {
2081 RawStridedRef::new_unchecked(
2082 operand_data,
2083 operand.dims(),
2084 operand.strides(),
2085 operand.offset(),
2086 )
2087 };
2088 let update_ref = unsafe {
2089 RawStridedRef::new_unchecked(
2090 update_data,
2091 update.dims(),
2092 update.strides(),
2093 update.offset(),
2094 )
2095 };
2096 let starts_ref = unsafe {
2097 RawStridedRef::new_unchecked(
2098 starts_data,
2099 starts.dims(),
2100 starts.strides(),
2101 starts.offset(),
2102 )
2103 };
2104 let mut dest_ref =
2105 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2106 unsafe {
2108 strided_basic::execution::dynamic_update_into_uninit(
2109 plan,
2110 &mut dest_ref,
2111 &operand_ref,
2112 &update_ref,
2113 &starts_ref,
2114 )
2115 }
2116}
2117
2118fn execute_scatter_uninit_dispatch<T>(
2119 plan: &ScatterPlan,
2120 index_dtype: KernelDType,
2121 dest: &mut ErasedRawStridedUninitMut<'_>,
2122 operand: &ErasedRawStridedRef<'_>,
2123 scatter_indices: &ErasedRawStridedRef<'_>,
2124 updates: &ErasedRawStridedRef<'_>,
2125 combine: fn(T, T) -> T,
2126) -> Result<()>
2127where
2128 T: Copy + Add<Output = T> + crate::MaybeSendSync + KernelStorageElement,
2129{
2130 match index_dtype {
2131 KernelDType::I32 => {
2132 execute_scatter_uninit::<T, i32>(plan, dest, operand, scatter_indices, updates, combine)
2133 }
2134 KernelDType::I64 => {
2135 execute_scatter_uninit::<T, i64>(plan, dest, operand, scatter_indices, updates, combine)
2136 }
2137 _ => Err(StridedError::UnsupportedDType {
2138 dtype: index_dtype.label(),
2139 }),
2140 }
2141}
2142
2143fn execute_scatter_uninit<T, I>(
2144 plan: &ScatterPlan,
2145 dest: &mut ErasedRawStridedUninitMut<'_>,
2146 operand: &ErasedRawStridedRef<'_>,
2147 scatter_indices: &ErasedRawStridedRef<'_>,
2148 updates: &ErasedRawStridedRef<'_>,
2149 combine: fn(T, T) -> T,
2150) -> Result<()>
2151where
2152 T: Copy + Add<Output = T> + crate::MaybeSendSync + KernelStorageElement,
2153 I: GatherIndex + KernelStorageElement,
2154{
2155 let indices = scatter_indices;
2156 let operand_data = operand.data_as::<T>()?;
2157 let index_data = indices.data_as::<I>()?;
2158 let update_data = updates.data_as::<T>()?;
2159 let dest_dims = dest.dims();
2160 let dest_strides = dest.strides();
2161 let dest_offset = dest.offset();
2162 let dest_data = dest.data_as_uninit_mut::<T>()?;
2163 let operand_ref = unsafe {
2164 RawStridedRef::new_unchecked(
2165 operand_data,
2166 operand.dims(),
2167 operand.strides(),
2168 operand.offset(),
2169 )
2170 };
2171 let index_ref = unsafe {
2172 RawStridedRef::new_unchecked(
2173 index_data,
2174 indices.dims(),
2175 indices.strides(),
2176 indices.offset(),
2177 )
2178 };
2179 let update_ref = unsafe {
2180 RawStridedRef::new_unchecked(
2181 update_data,
2182 updates.dims(),
2183 updates.strides(),
2184 updates.offset(),
2185 )
2186 };
2187 let mut dest_ref =
2188 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2189 unsafe {
2191 strided_basic::execution::scatter_into_uninit(
2192 plan,
2193 &mut dest_ref,
2194 &operand_ref,
2195 &index_ref,
2196 &update_ref,
2197 combine,
2198 )
2199 }
2200}
2201
2202fn execute_slice_uninit<T>(
2203 plan: &SlicePlan,
2204 dest: &mut ErasedRawStridedUninitMut<'_>,
2205 operand: &ErasedRawStridedRef<'_>,
2206) -> Result<()>
2207where
2208 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2209{
2210 let operand_data = operand.data_as::<T>()?;
2211 let dest_dims = dest.dims();
2212 let dest_strides = dest.strides();
2213 let dest_offset = dest.offset();
2214 let dest_data = dest.data_as_uninit_mut::<T>()?;
2215 let operand_ref = unsafe {
2216 RawStridedRef::new_unchecked(
2217 operand_data,
2218 operand.dims(),
2219 operand.strides(),
2220 operand.offset(),
2221 )
2222 };
2223 let mut dest_ref =
2224 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2225 plan.execute_uninit(&mut dest_ref, &operand_ref)
2226}
2227
2228fn execute_reverse<T>(
2229 plan: &ReversePlan,
2230 dest: &mut ErasedRawStridedMut<'_>,
2231 operand: &ErasedRawStridedRef<'_>,
2232) -> Result<()>
2233where
2234 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2235{
2236 let operand_data = operand.data_as::<T>()?;
2237 let dest_dims = dest.dims();
2238 let dest_strides = dest.strides();
2239 let dest_offset = dest.offset();
2240 let dest_data = dest.data_as_mut::<T>()?;
2241 let operand_ref = unsafe {
2242 RawStridedRef::new_unchecked(
2243 operand_data,
2244 operand.dims(),
2245 operand.strides(),
2246 operand.offset(),
2247 )
2248 };
2249 let mut dest_ref =
2250 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2251 plan.execute(&mut dest_ref, &operand_ref)
2252}
2253
2254fn execute_reverse_uninit<T>(
2255 plan: &ReversePlan,
2256 dest: &mut ErasedRawStridedUninitMut<'_>,
2257 operand: &ErasedRawStridedRef<'_>,
2258) -> Result<()>
2259where
2260 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2261{
2262 let operand_data = operand.data_as::<T>()?;
2263 let dest_dims = dest.dims();
2264 let dest_strides = dest.strides();
2265 let dest_offset = dest.offset();
2266 let dest_data = dest.data_as_uninit_mut::<T>()?;
2267 let operand_ref = unsafe {
2268 RawStridedRef::new_unchecked(
2269 operand_data,
2270 operand.dims(),
2271 operand.strides(),
2272 operand.offset(),
2273 )
2274 };
2275 let mut dest_ref =
2276 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2277 plan.execute_uninit(&mut dest_ref, &operand_ref)
2278}
2279
2280fn execute_pad<T>(
2281 plan: &PadPlan,
2282 dest: &mut ErasedRawStridedMut<'_>,
2283 operand: &ErasedRawStridedRef<'_>,
2284 fill: &[u8],
2285) -> Result<()>
2286where
2287 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2288{
2289 let fill = read_unaligned_scalar::<T>(fill);
2290 let operand_data = operand.data_as::<T>()?;
2291 let dest_dims = dest.dims();
2292 let dest_strides = dest.strides();
2293 let dest_offset = dest.offset();
2294 let dest_data = dest.data_as_mut::<T>()?;
2295 let operand_ref = unsafe {
2296 RawStridedRef::new_unchecked(
2297 operand_data,
2298 operand.dims(),
2299 operand.strides(),
2300 operand.offset(),
2301 )
2302 };
2303 let mut dest_ref =
2304 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2305 plan.execute(&mut dest_ref, &operand_ref, fill)
2306}
2307
2308fn execute_pad_uninit<T>(
2309 plan: &PadPlan,
2310 dest: &mut ErasedRawStridedUninitMut<'_>,
2311 operand: &ErasedRawStridedRef<'_>,
2312 fill: &[u8],
2313) -> Result<()>
2314where
2315 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2316{
2317 let fill = read_unaligned_scalar::<T>(fill);
2318 let operand_data = operand.data_as::<T>()?;
2319 let dest_dims = dest.dims();
2320 let dest_strides = dest.strides();
2321 let dest_offset = dest.offset();
2322 let dest_data = dest.data_as_uninit_mut::<T>()?;
2323 let operand_ref = unsafe {
2324 RawStridedRef::new_unchecked(
2325 operand_data,
2326 operand.dims(),
2327 operand.strides(),
2328 operand.offset(),
2329 )
2330 };
2331 let mut dest_ref =
2332 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2333 plan.execute_uninit(&mut dest_ref, &operand_ref, fill)
2334}
2335
2336fn dispatch_gather_index<T>(
2337 plan: &GatherPlan,
2338 index_dtype: KernelDType,
2339 dest: &mut ErasedRawStridedMut<'_>,
2340 operand: &ErasedRawStridedRef<'_>,
2341 start_indices: &ErasedRawStridedRef<'_>,
2342) -> Result<()>
2343where
2344 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2345{
2346 match index_dtype {
2347 KernelDType::I32 => execute_gather::<T, i32>(plan, dest, operand, start_indices),
2348 KernelDType::I64 => execute_gather::<T, i64>(plan, dest, operand, start_indices),
2349 _ => Err(StridedError::UnsupportedDType {
2350 dtype: index_dtype.label(),
2351 }),
2352 }
2353}
2354
2355fn execute_gather<T, I>(
2356 plan: &GatherPlan,
2357 dest: &mut ErasedRawStridedMut<'_>,
2358 operand: &ErasedRawStridedRef<'_>,
2359 start_indices: &ErasedRawStridedRef<'_>,
2360) -> Result<()>
2361where
2362 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2363 I: GatherIndex + KernelStorageElement,
2364{
2365 let operand_data = operand.data_as::<T>()?;
2366 let index_data = start_indices.data_as::<I>()?;
2367 let dest_dims = dest.dims();
2368 let dest_strides = dest.strides();
2369 let dest_offset = dest.offset();
2370 let dest_data = dest.data_as_mut::<T>()?;
2371 let operand_ref = unsafe {
2372 RawStridedRef::new_unchecked(
2373 operand_data,
2374 operand.dims(),
2375 operand.strides(),
2376 operand.offset(),
2377 )
2378 };
2379 let index_ref = unsafe {
2380 RawStridedRef::new_unchecked(
2381 index_data,
2382 start_indices.dims(),
2383 start_indices.strides(),
2384 start_indices.offset(),
2385 )
2386 };
2387 let mut dest_ref =
2388 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2389 plan.execute(&mut dest_ref, &operand_ref, &index_ref)
2390}
2391
2392fn dispatch_dynamic_slice_index<T>(
2393 plan: &DynamicSlicePlan,
2394 index_dtype: KernelDType,
2395 dest: &mut ErasedRawStridedMut<'_>,
2396 operand: &ErasedRawStridedRef<'_>,
2397 starts: &ErasedRawStridedRef<'_>,
2398) -> Result<()>
2399where
2400 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2401{
2402 match index_dtype {
2403 KernelDType::I32 => execute_dynamic_slice::<T, i32>(plan, dest, operand, starts),
2404 KernelDType::I64 => execute_dynamic_slice::<T, i64>(plan, dest, operand, starts),
2405 _ => Err(StridedError::UnsupportedDType {
2406 dtype: index_dtype.label(),
2407 }),
2408 }
2409}
2410
2411fn execute_dynamic_slice<T, I>(
2412 plan: &DynamicSlicePlan,
2413 dest: &mut ErasedRawStridedMut<'_>,
2414 operand: &ErasedRawStridedRef<'_>,
2415 starts: &ErasedRawStridedRef<'_>,
2416) -> Result<()>
2417where
2418 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2419 I: GatherIndex + KernelStorageElement,
2420{
2421 let operand_data = operand.data_as::<T>()?;
2422 let start_data = starts.data_as::<I>()?;
2423 let dest_dims = dest.dims();
2424 let dest_strides = dest.strides();
2425 let dest_offset = dest.offset();
2426 let dest_data = dest.data_as_mut::<T>()?;
2427 let operand_ref = unsafe {
2428 RawStridedRef::new_unchecked(
2429 operand_data,
2430 operand.dims(),
2431 operand.strides(),
2432 operand.offset(),
2433 )
2434 };
2435 let start_ref = unsafe {
2436 RawStridedRef::new_unchecked(start_data, starts.dims(), starts.strides(), starts.offset())
2437 };
2438 let mut dest_ref =
2439 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2440 plan.execute(&mut dest_ref, &operand_ref, &start_ref)
2441}
2442
2443fn dispatch_dynamic_update_slice_index<T>(
2444 plan: &DynamicUpdateSlicePlan,
2445 index_dtype: KernelDType,
2446 dest: &mut ErasedRawStridedMut<'_>,
2447 operand: &ErasedRawStridedRef<'_>,
2448 update: &ErasedRawStridedRef<'_>,
2449 starts: &ErasedRawStridedRef<'_>,
2450) -> Result<()>
2451where
2452 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2453{
2454 match index_dtype {
2455 KernelDType::I32 => {
2456 execute_dynamic_update_slice::<T, i32>(plan, dest, operand, update, starts)
2457 }
2458 KernelDType::I64 => {
2459 execute_dynamic_update_slice::<T, i64>(plan, dest, operand, update, starts)
2460 }
2461 _ => Err(StridedError::UnsupportedDType {
2462 dtype: index_dtype.label(),
2463 }),
2464 }
2465}
2466
2467fn execute_dynamic_update_slice<T, I>(
2468 plan: &DynamicUpdateSlicePlan,
2469 dest: &mut ErasedRawStridedMut<'_>,
2470 operand: &ErasedRawStridedRef<'_>,
2471 update: &ErasedRawStridedRef<'_>,
2472 starts: &ErasedRawStridedRef<'_>,
2473) -> Result<()>
2474where
2475 T: Copy + crate::MaybeSendSync + KernelStorageElement,
2476 I: GatherIndex + KernelStorageElement,
2477{
2478 let operand_data = operand.data_as::<T>()?;
2479 let update_data = update.data_as::<T>()?;
2480 let start_data = starts.data_as::<I>()?;
2481 let dest_dims = dest.dims();
2482 let dest_strides = dest.strides();
2483 let dest_offset = dest.offset();
2484 let dest_data = dest.data_as_mut::<T>()?;
2485 let operand_ref = unsafe {
2486 RawStridedRef::new_unchecked(
2487 operand_data,
2488 operand.dims(),
2489 operand.strides(),
2490 operand.offset(),
2491 )
2492 };
2493 let update_ref = unsafe {
2494 RawStridedRef::new_unchecked(
2495 update_data,
2496 update.dims(),
2497 update.strides(),
2498 update.offset(),
2499 )
2500 };
2501 let start_ref = unsafe {
2502 RawStridedRef::new_unchecked(start_data, starts.dims(), starts.strides(), starts.offset())
2503 };
2504 let mut dest_ref =
2505 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2506 plan.execute(&mut dest_ref, &operand_ref, &update_ref, &start_ref)
2507}
2508
2509fn dispatch_scatter_index<T>(
2510 plan: &ScatterPlan,
2511 index_dtype: KernelDType,
2512 dest: &mut ErasedRawStridedMut<'_>,
2513 operand: &ErasedRawStridedRef<'_>,
2514 scatter_indices: &ErasedRawStridedRef<'_>,
2515 updates: &ErasedRawStridedRef<'_>,
2516) -> Result<()>
2517where
2518 T: Copy + Add<Output = T> + crate::MaybeSendSync + KernelStorageElement,
2519{
2520 match index_dtype {
2521 KernelDType::I32 => {
2522 execute_scatter::<T, i32>(plan, dest, operand, scatter_indices, updates)
2523 }
2524 KernelDType::I64 => {
2525 execute_scatter::<T, i64>(plan, dest, operand, scatter_indices, updates)
2526 }
2527 _ => Err(StridedError::UnsupportedDType {
2528 dtype: index_dtype.label(),
2529 }),
2530 }
2531}
2532
2533fn execute_scatter<T, I>(
2534 plan: &ScatterPlan,
2535 dest: &mut ErasedRawStridedMut<'_>,
2536 operand: &ErasedRawStridedRef<'_>,
2537 scatter_indices: &ErasedRawStridedRef<'_>,
2538 updates: &ErasedRawStridedRef<'_>,
2539) -> Result<()>
2540where
2541 T: Copy + Add<Output = T> + crate::MaybeSendSync + KernelStorageElement,
2542 I: GatherIndex + KernelStorageElement,
2543{
2544 let operand_data = operand.data_as::<T>()?;
2545 let index_data = scatter_indices.data_as::<I>()?;
2546 let update_data = updates.data_as::<T>()?;
2547 let dest_dims = dest.dims();
2548 let dest_strides = dest.strides();
2549 let dest_offset = dest.offset();
2550 let dest_data = dest.data_as_mut::<T>()?;
2551 let operand_ref = unsafe {
2552 RawStridedRef::new_unchecked(
2553 operand_data,
2554 operand.dims(),
2555 operand.strides(),
2556 operand.offset(),
2557 )
2558 };
2559 let index_ref = unsafe {
2560 RawStridedRef::new_unchecked(
2561 index_data,
2562 scatter_indices.dims(),
2563 scatter_indices.strides(),
2564 scatter_indices.offset(),
2565 )
2566 };
2567 let update_ref = unsafe {
2568 RawStridedRef::new_unchecked(
2569 update_data,
2570 updates.dims(),
2571 updates.strides(),
2572 updates.offset(),
2573 )
2574 };
2575 let mut dest_ref =
2576 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
2577 plan.execute(&mut dest_ref, &operand_ref, &index_ref, &update_ref)
2578}
2579
2580fn read_unaligned_scalar<T>(bytes: &[u8]) -> T
2581where
2582 T: Copy,
2583{
2584 unsafe { core::ptr::read_unaligned(bytes.as_ptr().cast::<T>()) }
2585}