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