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