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