Skip to main content

strided_kernel/
erased.rs

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/// Runtime unary operation for [`erased_map_into`].
7#[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/// Runtime binary operation for [`erased_zip_into`].
28#[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
54/// Apply one runtime-selected unary operation without compiling a plan.
55///
56/// The destination must not overlap the input. Real and complex dtypes support
57/// every [`ErasedMapOp`], signed integers use wrapping negate/abs semantics,
58/// and `bool` supports only [`ErasedMapOp::Conj`]. Complex absolute value has
59/// the real output contract `c32 -> f32` and `c64 -> f64`; all other supported
60/// unary operations preserve dtype.
61///
62/// # Errors
63///
64/// Returns a typed [`StridedError`] for dtype, shape, output-layout, overlap,
65/// or unsupported dtype/op contracts. Validation completes before any write.
66pub 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    // SAFETY: input/output overlap was rejected before forming references.
77    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
100/// Apply one runtime-selected binary operation without compiling a plan.
101///
102/// The destination must not overlap either input. Real dtypes support every
103/// [`ErasedZipOp`]. Signed integers support every operation with wrapping
104/// arithmetic and a pre-write zero-divisor check. Complex dtypes support
105/// add/subtract/multiply/divide. `bool` has no binary one-shot operations.
106///
107/// # Errors
108///
109/// Returns a typed [`StridedError`] for dtype, shape, output-layout, overlap,
110/// or unsupported dtype/op contracts. Validation completes before any write.
111pub 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    // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
125    let lhs = unsafe { lhs.try_as_ref_after_no_overlap() }?;
126    // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
127    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/// Dtype-erased static-slice wrapper.
145#[derive(Clone, Debug)]
146pub struct ErasedSlicePlan {
147    dtype: KernelDType,
148    plan: SlicePlan,
149}
150
151/// Dtype-erased reverse wrapper.
152#[derive(Clone, Debug)]
153pub struct ErasedReversePlan {
154    dtype: KernelDType,
155    plan: ReversePlan,
156}
157
158/// Dtype-erased pad wrapper.
159#[derive(Clone, Debug)]
160pub struct ErasedPadPlan {
161    dtype: KernelDType,
162    plan: PadPlan,
163}
164
165impl ErasedSlicePlan {
166    /// Validate and store a static slice plan for one dtype and fixed layout set.
167    #[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    /// Execute a static slice into an erased output descriptor.
204    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    /// Execute a static slice as a full overwrite of uninitialized output storage.
229    /// On success, every reachable destination slot is fully overwritten;
230    /// unreachable holes are neither read nor initialized. Validation errors
231    /// are returned before any destination write. A panic during execution
232    /// may leave a partially initialized `MaybeUninit` destination, which is
233    /// still safely droppable; no readable value is promised for unwritten
234    /// reachable slots.
235    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        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
245        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    /// Validate and store a reverse plan for one dtype and fixed layout set.
264    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    /// Execute a reverse into an erased output descriptor.
289    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    /// Execute reverse as a full overwrite of uninitialized output storage.
314    /// On success, every reachable destination slot is fully overwritten;
315    /// unreachable holes are neither read nor initialized. Validation errors
316    /// are returned before any destination write. A panic during execution
317    /// may leave a partially initialized `MaybeUninit` destination, which is
318    /// still safely droppable; no readable value is promised for unwritten
319    /// reachable slots.
320    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        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
330        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    /// Validate and store a pad plan for one dtype and fixed layout set.
349    #[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    /// Execute pad into an erased output descriptor using one dtype scalar as fill.
386    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    /// Execute pad as a full overwrite of uninitialized output storage.
413    /// On success, every reachable destination slot is fully overwritten;
414    /// unreachable holes are neither read nor initialized. Validation errors
415    /// are returned before any destination write. A panic during execution
416    /// may leave a partially initialized `MaybeUninit` destination, which is
417    /// still safely droppable; no readable value is promised for unwritten
418    /// reachable slots.
419    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        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
431        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/// Dtype-erased gather wrapper.
449///
450/// This is the erased replay boundary for indexed reads. Value buffers use the
451/// configured value dtype, while the index descriptor must use `i32` or `i64`.
452#[derive(Clone, Debug)]
453pub struct ErasedGatherPlan {
454    dtype: KernelDType,
455    index_dtype: KernelDType,
456    plan: GatherPlan,
457}
458
459/// Dtype-erased fixed-window dynamic-slice wrapper.
460#[derive(Clone, Debug)]
461pub struct ErasedDynamicSlicePlan {
462    dtype: KernelDType,
463    index_dtype: KernelDType,
464    plan: DynamicSlicePlan,
465}
466
467/// Dtype-erased dynamic-update-slice wrapper.
468#[derive(Clone, Debug)]
469pub struct ErasedDynamicUpdateSlicePlan {
470    dtype: KernelDType,
471    index_dtype: KernelDType,
472    plan: DynamicUpdateSlicePlan,
473}
474
475/// Dtype-erased additive scatter wrapper.
476#[derive(Clone, Debug)]
477pub struct ErasedScatterPlan {
478    dtype: KernelDType,
479    index_dtype: KernelDType,
480    plan: ScatterPlan,
481}
482
483impl ErasedGatherPlan {
484    /// Validate and store a gather plan for one value dtype, index dtype, and layout set.
485    #[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    /// Execute an indexed read into an erased output descriptor.
530    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    /// Execute gather into a destination whose reachable slots may be
599    /// uninitialized. All validation precedes the first destination write.
600    /// On success, every reachable destination slot is fully overwritten;
601    /// unreachable holes are neither read nor initialized. Validation errors
602    /// are returned before any destination write. A panic during execution
603    /// may leave a partially initialized `MaybeUninit` destination, which is
604    /// still safely droppable; no readable value is promised for unwritten
605    /// reachable slots.
606    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        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
619        let operand = &unsafe { operand.try_as_ref_after_no_overlap() }?;
620        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
621        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    /// Validate and store a dynamic-slice plan for one value dtype, index dtype, and layout set.
686    #[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    /// Execute a fixed-window dynamic slice into an erased output descriptor.
731    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    /// Execute dynamic slice into a destination whose reachable slots may be
800    /// uninitialized.
801    /// On success, every reachable destination slot is fully overwritten;
802    /// unreachable holes are neither read nor initialized. Validation errors
803    /// are returned before any destination write. A panic during execution may
804    /// leave reachable slots partially initialized, but the `MaybeUninit`
805    /// destination remains safely droppable.
806    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        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
819        let operand = &unsafe { operand.try_as_ref_after_no_overlap() }?;
820        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
821        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    /// Validate and store a dynamic-update-slice plan for one value dtype, index dtype, and layout set.
886    #[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    /// Execute a dynamic update slice into an erased output descriptor.
933    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    /// On success, the copy phase initializes every reachable destination
1011    /// slot before the read-modify-write phase. Unreachable holes are neither
1012    /// read nor initialized. Validation errors before the copy leave the
1013    /// destination untouched; an error or panic after the copy may leave a
1014    /// mixture of old and new reachable values, all initialized and safely
1015    /// droppable.
1016    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        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
1032        let operand = &unsafe { operand.try_as_ref_after_no_overlap() }?;
1033        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
1034        let update = &unsafe { update.try_as_ref_after_no_overlap() }?;
1035        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
1036        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    /// Validate and store an additive scatter plan for one value dtype, index dtype, and layout set.
1108    #[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    /// Execute additive scatter into an erased output descriptor.
1157    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    /// On success, the copy phase initializes every reachable destination
1227    /// slot before the read-modify-write phase. Unreachable holes are neither
1228    /// read nor initialized. Validation errors before the copy leave the
1229    /// destination untouched; an error or panic after the copy may leave a
1230    /// mixture of old and new reachable values, all initialized and safely
1231    /// droppable.
1232    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        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
1248        let operand = &unsafe { operand.try_as_ref_after_no_overlap() }?;
1249        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
1250        let scatter_indices = &unsafe { scatter_indices.try_as_ref_after_no_overlap() }?;
1251        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
1252        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    // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
1353
1354    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    // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
1389    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    // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
1433    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            // SAFETY: RawStridedRef construction validated every reachable offset.
1714            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        // SAFETY: RawStridedRef construction validated every reachable offset.
1756        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    // SAFETY: the erased entry rejected overlap before constructing these descriptors.
1903    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    // SAFETY: the erased entry rejected overlap before constructing these descriptors.
1962    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    // SAFETY: the erased entry rejected overlap before constructing these descriptors.
2041    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    // SAFETY: the erased entry rejected overlap before constructing these descriptors.
2124    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}