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