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