Skip to main content

strided_kernel/erased/
uninit.rs

1//! Dtype-erased elementwise entry points that fully overwrite uninitialized
2//! output storage.
3//!
4//! Every entry point validates dtype, shape, destination injectivity, and
5//! input/output overlap before the first write. `Ok(())` means every reachable
6//! destination element was initialized. A panic during replay may leave the
7//! destination partially initialized, which remains safe to drop because the
8//! storage is `MaybeUninit`.
9
10use super::{erased_raw_ref, map_output_dtype, raw_any, OneShotScalar};
11use crate::*;
12use core::mem::MaybeUninit;
13use core::num::Wrapping;
14use num_complex::{Complex32, Complex64};
15use strided_basic::execution::{
16    check_dtype, ensure_same_shape, map_raw_into_validated,
17    validate_destination_layout_without_alloc, validate_uninit_no_overlap,
18    zip_map2_raw_into_validated, ValidatedDestinationLayout,
19};
20
21/// Apply one runtime-selected unary operation into uninitialized storage.
22///
23/// The dtype contract matches [`erased_map_into`](crate::erased_map_into):
24/// complex absolute value writes `c32 -> f32` and `c64 -> f64`, `bool`
25/// supports only [`ErasedMapOp::Conj`], and signed integers use wrapping
26/// negate/abs.
27///
28/// # Examples
29///
30/// ```
31/// use core::mem::MaybeUninit;
32/// use strided_kernel::{
33///     erased_map_into_uninit, ErasedMapOp, ErasedRawStridedPtr, ErasedRawStridedRef,
34///     ErasedRawStridedUninitMut, ExecContext, KernelDType,
35/// };
36///
37/// let input = [-1.5_f64, 2.0];
38/// let input = ErasedRawStridedRef::from_slice(&input, &[2], &[1], 0).unwrap();
39/// let mut out = [MaybeUninit::<f64>::uninit(); 2];
40/// let mut dest = ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[2], &[1], 0).unwrap();
41/// erased_map_into_uninit(
42///     KernelDType::F64,
43///     ErasedMapOp::Abs,
44///     &ExecContext::serial(),
45///     &mut dest,
46///     &ErasedRawStridedPtr::from_ref(&input),
47/// )
48/// .unwrap();
49/// // SAFETY: a successful call initializes every reachable element.
50/// assert_eq!(unsafe { [out[0].assume_init(), out[1].assume_init()] }, [1.5, 2.0]);
51/// ```
52///
53/// # Errors
54///
55/// Returns a typed [`StridedError`] for dtype, shape, output-layout, overlap,
56/// or unsupported dtype/op contracts. Validation completes before any write.
57pub fn erased_map_into_uninit(
58    input_dtype: KernelDType,
59    op: ErasedMapOp,
60    ctx: &ExecContext,
61    dest: &mut ErasedRawStridedUninitMut<'_>,
62    input: &ErasedRawStridedPtr<'_>,
63) -> Result<()> {
64    check_dtype(input_dtype, input.dtype())?;
65    check_dtype(map_output_dtype(input_dtype, op)?, dest.dtype())?;
66    validate_uninit_no_overlap(dest, input, 0)?;
67    // SAFETY: input/output overlap was rejected before forming references.
68    let input = unsafe { input.try_as_ref_after_no_overlap() }?;
69
70    ctx.run(|| match (input_dtype, op) {
71        (KernelDType::C32, ErasedMapOp::Abs) => {
72            uninit_map_with::<f32, Complex32>(dest, &input, |value| value.norm())
73        }
74        (KernelDType::C64, ErasedMapOp::Abs) => {
75            uninit_map_with::<f64, Complex64>(dest, &input, |value| value.norm())
76        }
77        (KernelDType::F32, _) => uninit_map::<f32>(op, dest, &input),
78        (KernelDType::F64, _) => uninit_map::<f64>(op, dest, &input),
79        (KernelDType::I32, _) => uninit_map::<i32>(op, dest, &input),
80        (KernelDType::I64, _) => uninit_map::<i64>(op, dest, &input),
81        (KernelDType::Bool, _) => uninit_map::<bool>(op, dest, &input),
82        (KernelDType::C32, _) => uninit_map::<Complex32>(op, dest, &input),
83        (KernelDType::C64, _) => uninit_map::<Complex64>(op, dest, &input),
84        _ => Err(StridedError::UnsupportedDType {
85            dtype: input_dtype.label(),
86        }),
87    })
88}
89
90/// Apply one runtime-selected binary operation into uninitialized storage.
91///
92/// The dtype contract matches [`erased_zip_into`](crate::erased_zip_into).
93/// Signed-integer divide and remainder scan the divisor before any write and
94/// return [`StridedError::IntegerDivisionByZero`] when it contains zero; all
95/// other integer arithmetic wraps. Real maximum/minimum propagate NaN.
96///
97/// # Examples
98///
99/// ```
100/// use core::mem::MaybeUninit;
101/// use strided_kernel::{
102///     erased_zip_into_uninit, ErasedRawStridedPtr, ErasedRawStridedRef,
103///     ErasedRawStridedUninitMut, ErasedZipOp, ExecContext, KernelDType,
104/// };
105///
106/// let lhs = [i32::MAX, 7];
107/// let rhs = [1_i32, 2];
108/// let lhs = ErasedRawStridedRef::from_slice(&lhs, &[2], &[1], 0).unwrap();
109/// let rhs = ErasedRawStridedRef::from_slice(&rhs, &[2], &[1], 0).unwrap();
110/// let mut out = [MaybeUninit::<i32>::uninit(); 2];
111/// let mut dest = ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[2], &[1], 0).unwrap();
112/// erased_zip_into_uninit(
113///     KernelDType::I32,
114///     ErasedZipOp::Add,
115///     &ExecContext::serial(),
116///     &mut dest,
117///     &ErasedRawStridedPtr::from_ref(&lhs),
118///     &ErasedRawStridedPtr::from_ref(&rhs),
119/// )
120/// .unwrap();
121/// // SAFETY: a successful call initializes every reachable element.
122/// assert_eq!(unsafe { [out[0].assume_init(), out[1].assume_init()] }, [i32::MIN, 9]);
123/// ```
124///
125/// # Errors
126///
127/// Returns a typed [`StridedError`] for dtype, shape, output-layout, overlap,
128/// unsupported dtype/op contracts, or an integer zero divisor. Validation
129/// completes before any write.
130pub fn erased_zip_into_uninit(
131    dtype: KernelDType,
132    op: ErasedZipOp,
133    ctx: &ExecContext,
134    dest: &mut ErasedRawStridedUninitMut<'_>,
135    lhs: &ErasedRawStridedPtr<'_>,
136    rhs: &ErasedRawStridedPtr<'_>,
137) -> Result<()> {
138    check_dtype(dtype, dest.dtype())?;
139    check_dtype(dtype, lhs.dtype())?;
140    check_dtype(dtype, rhs.dtype())?;
141    validate_uninit_no_overlap(dest, lhs, 0)?;
142    validate_uninit_no_overlap(dest, rhs, 1)?;
143    // SAFETY: input/output overlap was rejected before forming references.
144    let lhs = unsafe { lhs.try_as_ref_after_no_overlap() }?;
145    // SAFETY: input/output overlap was rejected before forming references.
146    let rhs = unsafe { rhs.try_as_ref_after_no_overlap() }?;
147
148    ctx.run(|| match dtype {
149        KernelDType::F32 => uninit_zip::<f32>(op, dest, &lhs, &rhs),
150        KernelDType::F64 => uninit_zip::<f64>(op, dest, &lhs, &rhs),
151        KernelDType::I32 => uninit_zip::<i32>(op, dest, &lhs, &rhs),
152        KernelDType::I64 => uninit_zip::<i64>(op, dest, &lhs, &rhs),
153        KernelDType::Bool => uninit_zip::<bool>(op, dest, &lhs, &rhs),
154        KernelDType::C32 => uninit_zip::<Complex32>(op, dest, &lhs, &rhs),
155        KernelDType::C64 => uninit_zip::<Complex64>(op, dest, &lhs, &rhs),
156        _ => Err(StridedError::UnsupportedDType {
157            dtype: dtype.label(),
158        }),
159    })
160}
161
162/// Compare two same-dtype operands elementwise into uninitialized `bool`
163/// storage.
164///
165/// Real, signed-integer, and `bool` operands support every [`CompareOp`] with
166/// `PartialOrd` semantics, so any comparison involving NaN is `false`.
167/// Complex operands support only [`CompareOp::Eq`]; ordered comparisons return
168/// [`StridedError::UnsupportedOp`]. The destination dtype must be `bool`.
169///
170/// # Examples
171///
172/// ```
173/// use core::mem::MaybeUninit;
174/// use strided_kernel::{
175///     erased_compare_into_uninit, CompareOp, ErasedRawStridedPtr, ErasedRawStridedRef,
176///     ErasedRawStridedUninitMut, ExecContext, KernelDType,
177/// };
178///
179/// let lhs = [1.0_f64, f64::NAN];
180/// let rhs = [2.0_f64, 0.0];
181/// let lhs = ErasedRawStridedRef::from_slice(&lhs, &[2], &[1], 0).unwrap();
182/// let rhs = ErasedRawStridedRef::from_slice(&rhs, &[2], &[1], 0).unwrap();
183/// let mut out = [MaybeUninit::<bool>::uninit(); 2];
184/// let mut dest = ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[2], &[1], 0).unwrap();
185/// erased_compare_into_uninit(
186///     KernelDType::F64,
187///     CompareOp::Lt,
188///     &ExecContext::serial(),
189///     &mut dest,
190///     &ErasedRawStridedPtr::from_ref(&lhs),
191///     &ErasedRawStridedPtr::from_ref(&rhs),
192/// )
193/// .unwrap();
194/// // SAFETY: a successful call initializes every reachable element.
195/// assert_eq!(unsafe { [out[0].assume_init(), out[1].assume_init()] }, [true, false]);
196/// ```
197///
198/// # Errors
199///
200/// Returns a typed [`StridedError`] for dtype, shape, output-layout, overlap,
201/// or unsupported dtype/op contracts. Validation completes before any write.
202pub fn erased_compare_into_uninit(
203    dtype: KernelDType,
204    op: CompareOp,
205    ctx: &ExecContext,
206    dest: &mut ErasedRawStridedUninitMut<'_>,
207    lhs: &ErasedRawStridedPtr<'_>,
208    rhs: &ErasedRawStridedPtr<'_>,
209) -> Result<()> {
210    check_dtype(KernelDType::Bool, dest.dtype())?;
211    check_dtype(dtype, lhs.dtype())?;
212    check_dtype(dtype, rhs.dtype())?;
213    validate_uninit_no_overlap(dest, lhs, 0)?;
214    validate_uninit_no_overlap(dest, rhs, 1)?;
215    // SAFETY: input/output overlap was rejected before forming references.
216    let lhs = unsafe { lhs.try_as_ref_after_no_overlap() }?;
217    // SAFETY: input/output overlap was rejected before forming references.
218    let rhs = unsafe { rhs.try_as_ref_after_no_overlap() }?;
219
220    ctx.run(|| match dtype {
221        KernelDType::F32 => uninit_compare_ordered::<f32>(op, dest, &lhs, &rhs),
222        KernelDType::F64 => uninit_compare_ordered::<f64>(op, dest, &lhs, &rhs),
223        KernelDType::I32 => uninit_compare_ordered::<i32>(op, dest, &lhs, &rhs),
224        KernelDType::I64 => uninit_compare_ordered::<i64>(op, dest, &lhs, &rhs),
225        KernelDType::Bool => uninit_compare_ordered::<bool>(op, dest, &lhs, &rhs),
226        KernelDType::C32 => uninit_compare_eq::<Complex32>(op, dest, &lhs, &rhs),
227        KernelDType::C64 => uninit_compare_eq::<Complex64>(op, dest, &lhs, &rhs),
228        _ => Err(StridedError::UnsupportedDType {
229            dtype: dtype.label(),
230        }),
231    })
232}
233
234/// Select elementwise between two same-dtype operands by a `bool` predicate,
235/// writing uninitialized storage: `dest[i] = if pred[i] { on_true[i] } else
236/// { on_false[i] }`.
237///
238/// Every dtype is supported. The predicate dtype must be `bool`.
239///
240/// # Examples
241///
242/// ```
243/// use core::mem::MaybeUninit;
244/// use strided_kernel::{
245///     erased_select_into_uninit, ErasedRawStridedPtr, ErasedRawStridedRef,
246///     ErasedRawStridedUninitMut, ExecContext, KernelDType,
247/// };
248///
249/// let pred = [true, false];
250/// let on_true = [1_i64, 2];
251/// let on_false = [10_i64, 20];
252/// let pred = ErasedRawStridedRef::from_slice(&pred, &[2], &[1], 0).unwrap();
253/// let on_true = ErasedRawStridedRef::from_slice(&on_true, &[2], &[1], 0).unwrap();
254/// let on_false = ErasedRawStridedRef::from_slice(&on_false, &[2], &[1], 0).unwrap();
255/// let mut out = [MaybeUninit::<i64>::uninit(); 2];
256/// let mut dest = ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[2], &[1], 0).unwrap();
257/// erased_select_into_uninit(
258///     KernelDType::I64,
259///     &ExecContext::serial(),
260///     &mut dest,
261///     &ErasedRawStridedPtr::from_ref(&pred),
262///     &ErasedRawStridedPtr::from_ref(&on_true),
263///     &ErasedRawStridedPtr::from_ref(&on_false),
264/// )
265/// .unwrap();
266/// // SAFETY: a successful call initializes every reachable element.
267/// assert_eq!(unsafe { [out[0].assume_init(), out[1].assume_init()] }, [1, 20]);
268/// ```
269///
270/// # Errors
271///
272/// Returns a typed [`StridedError`] for dtype, shape, output-layout, or
273/// overlap contracts. Validation completes before any write.
274pub fn erased_select_into_uninit(
275    dtype: KernelDType,
276    ctx: &ExecContext,
277    dest: &mut ErasedRawStridedUninitMut<'_>,
278    pred: &ErasedRawStridedPtr<'_>,
279    on_true: &ErasedRawStridedPtr<'_>,
280    on_false: &ErasedRawStridedPtr<'_>,
281) -> Result<()> {
282    check_dtype(dtype, dest.dtype())?;
283    check_dtype(KernelDType::Bool, pred.dtype())?;
284    check_dtype(dtype, on_true.dtype())?;
285    check_dtype(dtype, on_false.dtype())?;
286    validate_uninit_no_overlap(dest, pred, 0)?;
287    validate_uninit_no_overlap(dest, on_true, 1)?;
288    validate_uninit_no_overlap(dest, on_false, 2)?;
289    // SAFETY: input/output overlap was rejected before forming references.
290    let pred = unsafe { pred.try_as_ref_after_no_overlap() }?;
291    // SAFETY: input/output overlap was rejected before forming references.
292    let on_true = unsafe { on_true.try_as_ref_after_no_overlap() }?;
293    // SAFETY: input/output overlap was rejected before forming references.
294    let on_false = unsafe { on_false.try_as_ref_after_no_overlap() }?;
295
296    ctx.run(|| match dtype {
297        KernelDType::F32 => uninit_select::<f32>(dest, &pred, &on_true, &on_false),
298        KernelDType::F64 => uninit_select::<f64>(dest, &pred, &on_true, &on_false),
299        KernelDType::I32 => uninit_select::<i32>(dest, &pred, &on_true, &on_false),
300        KernelDType::I64 => uninit_select::<i64>(dest, &pred, &on_true, &on_false),
301        KernelDType::Bool => uninit_select::<bool>(dest, &pred, &on_true, &on_false),
302        KernelDType::C32 => uninit_select::<Complex32>(dest, &pred, &on_true, &on_false),
303        KernelDType::C64 => uninit_select::<Complex64>(dest, &pred, &on_true, &on_false),
304        _ => Err(StridedError::UnsupportedDType {
305            dtype: dtype.label(),
306        }),
307    })
308}
309
310/// Clamp elementwise into uninitialized storage:
311/// `dest[i] = minimum(hi[i], maximum(lo[i], x[i]))`.
312///
313/// `maximum` and `minimum` are [`ErasedZipOp::Maximum`] and
314/// [`ErasedZipOp::Minimum`], so a NaN in any real operand yields NaN and
315/// `lo > hi` yields `hi`. Supported dtypes are `f32`, `f64`, `i32`, and `i64`.
316///
317/// # Examples
318///
319/// ```
320/// use core::mem::MaybeUninit;
321/// use strided_kernel::{
322///     erased_clamp_into_uninit, ErasedRawStridedPtr, ErasedRawStridedRef,
323///     ErasedRawStridedUninitMut, ExecContext, KernelDType,
324/// };
325///
326/// let x = [-3.0_f32, 0.5, 9.0];
327/// let lo = [0.0_f32; 1];
328/// let hi = [1.0_f32; 1];
329/// let x = ErasedRawStridedRef::from_slice(&x, &[3], &[1], 0).unwrap();
330/// // Stride-0 descriptors broadcast one bound over the output.
331/// let lo = ErasedRawStridedRef::from_slice(&lo, &[3], &[0], 0).unwrap();
332/// let hi = ErasedRawStridedRef::from_slice(&hi, &[3], &[0], 0).unwrap();
333/// let mut out = [MaybeUninit::<f32>::uninit(); 3];
334/// let mut dest = ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[3], &[1], 0).unwrap();
335/// erased_clamp_into_uninit(
336///     KernelDType::F32,
337///     &ExecContext::serial(),
338///     &mut dest,
339///     &ErasedRawStridedPtr::from_ref(&x),
340///     &ErasedRawStridedPtr::from_ref(&lo),
341///     &ErasedRawStridedPtr::from_ref(&hi),
342/// )
343/// .unwrap();
344/// // SAFETY: a successful call initializes every reachable element.
345/// let out = unsafe { out.map(|value| value.assume_init()) };
346/// assert_eq!(out, [0.0, 0.5, 1.0]);
347/// ```
348///
349/// # Errors
350///
351/// Returns a typed [`StridedError`] for dtype, shape, output-layout, or
352/// overlap contracts. Validation completes before any write.
353pub fn erased_clamp_into_uninit(
354    dtype: KernelDType,
355    ctx: &ExecContext,
356    dest: &mut ErasedRawStridedUninitMut<'_>,
357    x: &ErasedRawStridedPtr<'_>,
358    lo: &ErasedRawStridedPtr<'_>,
359    hi: &ErasedRawStridedPtr<'_>,
360) -> Result<()> {
361    check_dtype(dtype, dest.dtype())?;
362    check_dtype(dtype, x.dtype())?;
363    check_dtype(dtype, lo.dtype())?;
364    check_dtype(dtype, hi.dtype())?;
365    validate_uninit_no_overlap(dest, x, 0)?;
366    validate_uninit_no_overlap(dest, lo, 1)?;
367    validate_uninit_no_overlap(dest, hi, 2)?;
368    // SAFETY: input/output overlap was rejected before forming references.
369    let x = unsafe { x.try_as_ref_after_no_overlap() }?;
370    // SAFETY: input/output overlap was rejected before forming references.
371    let lo = unsafe { lo.try_as_ref_after_no_overlap() }?;
372    // SAFETY: input/output overlap was rejected before forming references.
373    let hi = unsafe { hi.try_as_ref_after_no_overlap() }?;
374
375    ctx.run(|| match dtype {
376        KernelDType::F32 => uninit_clamp::<f32>(dest, &x, &lo, &hi),
377        KernelDType::F64 => uninit_clamp::<f64>(dest, &x, &lo, &hi),
378        KernelDType::I32 => uninit_clamp::<i32>(dest, &x, &lo, &hi),
379        KernelDType::I64 => uninit_clamp::<i64>(dest, &x, &lo, &hi),
380        _ => Err(StridedError::UnsupportedDType {
381            dtype: dtype.label(),
382        }),
383    })
384}
385
386/// Broadcast two same-dtype operands onto the destination axes and multiply
387/// them into uninitialized storage.
388///
389/// `lhs_axes[k]` names the destination axis that `lhs` axis `k` maps to; each
390/// mapped extent must equal the destination extent or be one, and unmapped
391/// destination axes broadcast. Supported dtypes are `f32`, `f64`, `c32`,
392/// `c64`, and the signed integers, whose products wrap on overflow.
393///
394/// # Examples
395///
396/// ```
397/// use core::mem::MaybeUninit;
398/// use strided_kernel::{
399///     erased_broadcast_mul_into_uninit, ErasedRawStridedPtr, ErasedRawStridedRef,
400///     ErasedRawStridedUninitMut, ExecContext, KernelDType,
401/// };
402///
403/// // Outer product: out[i, j] = a[i] * b[j], column-major 2 x 3 output.
404/// let a = [1.0_f64, 2.0];
405/// let b = [10.0_f64, 20.0, 30.0];
406/// let a = ErasedRawStridedRef::from_slice(&a, &[2], &[1], 0).unwrap();
407/// let b = ErasedRawStridedRef::from_slice(&b, &[3], &[1], 0).unwrap();
408/// let mut out = [MaybeUninit::<f64>::uninit(); 6];
409/// let mut dest =
410///     ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[2, 3], &[1, 2], 0).unwrap();
411/// erased_broadcast_mul_into_uninit(
412///     KernelDType::F64,
413///     &ExecContext::serial(),
414///     &mut dest,
415///     &ErasedRawStridedPtr::from_ref(&a),
416///     &[0],
417///     &ErasedRawStridedPtr::from_ref(&b),
418///     &[1],
419/// )
420/// .unwrap();
421/// // SAFETY: a successful call initializes every reachable element.
422/// let out = unsafe { out.map(|value| value.assume_init()) };
423/// assert_eq!(out, [10.0, 20.0, 20.0, 40.0, 30.0, 60.0]);
424/// ```
425///
426/// # Errors
427///
428/// Returns a typed [`StridedError`] for dtype, axis mapping, shape,
429/// output-layout, or overlap contracts. Validation completes before any write.
430pub fn erased_broadcast_mul_into_uninit(
431    dtype: KernelDType,
432    ctx: &ExecContext,
433    dest: &mut ErasedRawStridedUninitMut<'_>,
434    lhs: &ErasedRawStridedPtr<'_>,
435    lhs_axes: &[usize],
436    rhs: &ErasedRawStridedPtr<'_>,
437    rhs_axes: &[usize],
438) -> Result<()> {
439    check_dtype(dtype, dest.dtype())?;
440    check_dtype(dtype, lhs.dtype())?;
441    check_dtype(dtype, rhs.dtype())?;
442    validate_uninit_no_overlap(dest, lhs, 0)?;
443    validate_uninit_no_overlap(dest, rhs, 1)?;
444    // SAFETY: input/output overlap was rejected before forming references.
445    let lhs = unsafe { lhs.try_as_ref_after_no_overlap() }?;
446    // SAFETY: input/output overlap was rejected before forming references.
447    let rhs = unsafe { rhs.try_as_ref_after_no_overlap() }?;
448
449    ctx.run(|| match dtype {
450        KernelDType::F32 => uninit_broadcast_mul::<f32>(dest, &lhs, lhs_axes, &rhs, rhs_axes),
451        KernelDType::F64 => uninit_broadcast_mul::<f64>(dest, &lhs, lhs_axes, &rhs, rhs_axes),
452        KernelDType::C32 => uninit_broadcast_mul::<Complex32>(dest, &lhs, lhs_axes, &rhs, rhs_axes),
453        KernelDType::C64 => uninit_broadcast_mul::<Complex64>(dest, &lhs, lhs_axes, &rhs, rhs_axes),
454        KernelDType::I32 => {
455            uninit_broadcast_mul_wrapping::<i32>(dest, &lhs, lhs_axes, &rhs, rhs_axes)
456        }
457        KernelDType::I64 => {
458            uninit_broadcast_mul_wrapping::<i64>(dest, &lhs, lhs_axes, &rhs, rhs_axes)
459        }
460        _ => Err(StridedError::UnsupportedDType {
461            dtype: dtype.label(),
462        }),
463    })
464}
465
466/// Validate the destination layout and the input shapes.
467///
468/// Returns `None` for an empty destination, which needs no writes.
469fn prepare_destination(
470    dest: &ErasedRawStridedUninitMut<'_>,
471    inputs: &[&[usize]],
472) -> Result<Option<ValidatedDestinationLayout>> {
473    let validated = validate_destination_layout_without_alloc(dest.dims(), dest.strides())?;
474    for input_dims in inputs {
475        ensure_same_shape(dest.dims(), input_dims)?;
476    }
477    if dest.dims().contains(&0) {
478        Ok(None)
479    } else {
480        Ok(Some(validated))
481    }
482}
483
484fn uninit_raw_mut<'a, T: KernelStorageElement>(
485    dest: &'a mut ErasedRawStridedUninitMut<'_>,
486) -> Result<RawStridedMut<'a, MaybeUninit<T>>> {
487    let dims = dest.dims();
488    let strides = dest.strides();
489    let offset = dest.offset();
490    let data = dest.data_as_uninit_mut::<T>()?;
491    // SAFETY: the uninit descriptor validated every reachable offset against
492    // its storage at construction.
493    Ok(unsafe { RawStridedMut::new_unchecked(data, dims, strides, offset) })
494}
495
496fn uninit_view_mut<'a, T: KernelStorageElement>(
497    dest: &'a mut ErasedRawStridedUninitMut<'_>,
498) -> Result<StridedViewMut<'a, MaybeUninit<T>>> {
499    let dims = dest.dims();
500    let strides = dest.strides();
501    let offset = dest.offset();
502    let data = dest.data_as_uninit_mut::<T>()?;
503    // SAFETY: the uninit descriptor validated every reachable offset against
504    // its storage at construction.
505    Ok(unsafe { StridedViewMut::new_unchecked(data, dims, strides, offset) })
506}
507
508fn erased_view<'a, T: KernelStorageElement>(
509    src: &'a ErasedRawStridedRef<'a>,
510) -> Result<StridedView<'a, T>> {
511    let data = src.data_as::<T>()?;
512    // SAFETY: the descriptor validated every reachable offset at construction.
513    Ok(unsafe { StridedView::new_unchecked(data, src.dims(), src.strides(), src.offset()) })
514}
515
516fn uninit_map<T: OneShotScalar>(
517    op: ErasedMapOp,
518    dest: &mut ErasedRawStridedUninitMut<'_>,
519    input: &ErasedRawStridedRef<'_>,
520) -> Result<()> {
521    if !T::supports_map(op) {
522        return Err(StridedError::UnsupportedOp {
523            op: op.label(),
524            dtype: T::one_shot_dtype_label(),
525        });
526    }
527    // Select the operation once, outside the element loop, so each replay
528    // closure is monomorphic and the inner loop can vectorize.
529    match op {
530        ErasedMapOp::Negate => {
531            uninit_map_with::<T, T>(dest, input, |value| T::map(ErasedMapOp::Negate, value))
532        }
533        ErasedMapOp::Conj => {
534            uninit_map_with::<T, T>(dest, input, |value| T::map(ErasedMapOp::Conj, value))
535        }
536        ErasedMapOp::Abs => {
537            uninit_map_with::<T, T>(dest, input, |value| T::map(ErasedMapOp::Abs, value))
538        }
539        ErasedMapOp::Sign => {
540            uninit_map_with::<T, T>(dest, input, |value| T::map(ErasedMapOp::Sign, value))
541        }
542    }
543}
544
545fn uninit_map_with<D, A>(
546    dest: &mut ErasedRawStridedUninitMut<'_>,
547    input: &ErasedRawStridedRef<'_>,
548    map: impl Fn(A) -> D + crate::MaybeSync,
549) -> Result<()>
550where
551    D: Copy + crate::MaybeSendSync + KernelStorageElement,
552    A: Copy + crate::MaybeSendSync + KernelStorageElement,
553{
554    let Some(validated) = prepare_destination(dest, &[input.dims()])? else {
555        return Ok(());
556    };
557    let input = erased_raw_ref::<A>(input)?;
558    let mut dest = uninit_raw_mut::<D>(dest)?;
559    // SAFETY: matching shapes and this destination layout were validated above,
560    // and the caller rejected input/output overlap.
561    unsafe {
562        map_raw_into_validated::<MaybeUninit<D>, A, Identity>(
563            &mut dest,
564            &input,
565            |value| MaybeUninit::new(map(value)),
566            validated,
567        )
568    }
569}
570
571fn uninit_zip<T: OneShotScalar>(
572    op: ErasedZipOp,
573    dest: &mut ErasedRawStridedUninitMut<'_>,
574    lhs: &ErasedRawStridedRef<'_>,
575    rhs: &ErasedRawStridedRef<'_>,
576) -> Result<()> {
577    if !T::supports_zip(op) {
578        return Err(StridedError::UnsupportedOp {
579            op: op.label(),
580            dtype: T::one_shot_dtype_label(),
581        });
582    }
583    let Some(validated) = prepare_destination(dest, &[lhs.dims(), rhs.dims()])? else {
584        return Ok(());
585    };
586    let lhs = erased_raw_ref::<T>(lhs)?;
587    let rhs = erased_raw_ref::<T>(rhs)?;
588    if matches!(op, ErasedZipOp::Divide | ErasedZipOp::Remainder)
589        && T::INTEGER
590        && raw_any(&rhs, T::is_zero)?
591    {
592        return Err(StridedError::IntegerDivisionByZero { op: op.label() });
593    }
594    let mut dest = uninit_raw_mut::<T>(dest)?;
595    // Select the operation once, outside the element loop, so each replay
596    // closure is monomorphic and the inner loop can vectorize.
597    macro_rules! replay {
598        ($op:ident) => {
599            // SAFETY: matching shapes and this destination layout were
600            // validated above, and the caller rejected input/output overlap.
601            unsafe {
602                zip_map2_raw_into_validated::<MaybeUninit<T>, T, T, Identity, Identity>(
603                    &mut dest,
604                    &lhs,
605                    &rhs,
606                    |lhs, rhs| MaybeUninit::new(T::zip(ErasedZipOp::$op, lhs, rhs)),
607                    validated,
608                )
609            }
610        };
611    }
612    match op {
613        ErasedZipOp::Add => replay!(Add),
614        ErasedZipOp::Subtract => replay!(Subtract),
615        ErasedZipOp::Multiply => replay!(Multiply),
616        ErasedZipOp::Divide => replay!(Divide),
617        ErasedZipOp::Remainder => replay!(Remainder),
618        ErasedZipOp::Maximum => replay!(Maximum),
619        ErasedZipOp::Minimum => replay!(Minimum),
620    }
621}
622
623fn uninit_compare_with<T>(
624    dest: &mut ErasedRawStridedUninitMut<'_>,
625    lhs: &ErasedRawStridedRef<'_>,
626    rhs: &ErasedRawStridedRef<'_>,
627    compare: impl Fn(T, T) -> bool + crate::MaybeSync,
628) -> Result<()>
629where
630    T: Copy + crate::MaybeSendSync + KernelStorageElement,
631{
632    let Some(validated) = prepare_destination(dest, &[lhs.dims(), rhs.dims()])? else {
633        return Ok(());
634    };
635    let lhs = erased_raw_ref::<T>(lhs)?;
636    let rhs = erased_raw_ref::<T>(rhs)?;
637    let mut dest = uninit_raw_mut::<bool>(dest)?;
638    // SAFETY: matching shapes and this destination layout were validated above,
639    // and the caller rejected input/output overlap.
640    unsafe {
641        zip_map2_raw_into_validated::<MaybeUninit<bool>, T, T, Identity, Identity>(
642            &mut dest,
643            &lhs,
644            &rhs,
645            |lhs, rhs| MaybeUninit::new(compare(lhs, rhs)),
646            validated,
647        )
648    }
649}
650
651fn uninit_compare_ordered<T>(
652    op: CompareOp,
653    dest: &mut ErasedRawStridedUninitMut<'_>,
654    lhs: &ErasedRawStridedRef<'_>,
655    rhs: &ErasedRawStridedRef<'_>,
656) -> Result<()>
657where
658    T: Copy + crate::MaybeSendSync + KernelStorageElement + PartialOrd,
659{
660    // The comparison is selected once, outside the element loop.
661    match op {
662        CompareOp::Eq => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a == b),
663        CompareOp::Lt => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a < b),
664        CompareOp::Le => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a <= b),
665        CompareOp::Gt => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a > b),
666        CompareOp::Ge => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a >= b),
667        _ => Err(StridedError::UnsupportedOp {
668            op: compare_label(op),
669            dtype: T::DTYPE.label(),
670        }),
671    }
672}
673
674fn uninit_compare_eq<T>(
675    op: CompareOp,
676    dest: &mut ErasedRawStridedUninitMut<'_>,
677    lhs: &ErasedRawStridedRef<'_>,
678    rhs: &ErasedRawStridedRef<'_>,
679) -> Result<()>
680where
681    T: Copy + crate::MaybeSendSync + KernelStorageElement + PartialEq,
682{
683    match op {
684        CompareOp::Eq => uninit_compare_with::<T>(dest, lhs, rhs, |a, b| a == b),
685        _ => Err(StridedError::UnsupportedOp {
686            op: compare_label(op),
687            dtype: T::DTYPE.label(),
688        }),
689    }
690}
691
692fn compare_label(op: CompareOp) -> &'static str {
693    match op {
694        CompareOp::Eq => "eq",
695        CompareOp::Lt => "lt",
696        CompareOp::Le => "le",
697        CompareOp::Gt => "gt",
698        CompareOp::Ge => "ge",
699        _ => "compare",
700    }
701}
702
703fn uninit_select<T>(
704    dest: &mut ErasedRawStridedUninitMut<'_>,
705    pred: &ErasedRawStridedRef<'_>,
706    on_true: &ErasedRawStridedRef<'_>,
707    on_false: &ErasedRawStridedRef<'_>,
708) -> Result<()>
709where
710    T: Copy + crate::MaybeSendSync + KernelStorageElement,
711{
712    let dims: [&[usize]; 3] = [pred.dims(), on_true.dims(), on_false.dims()];
713    if prepare_destination(dest, &dims)?.is_none() {
714        return Ok(());
715    }
716    let pred = erased_view::<bool>(pred)?;
717    let on_true = erased_view::<T>(on_true)?;
718    let on_false = erased_view::<T>(on_false)?;
719    let mut dest = uninit_view_mut::<T>(dest)?;
720    zip_map3_into(&mut dest, &pred, &on_true, &on_false, |pred, a, b| {
721        MaybeUninit::new(if pred { a } else { b })
722    })
723}
724
725fn uninit_clamp<T: OneShotScalar>(
726    dest: &mut ErasedRawStridedUninitMut<'_>,
727    x: &ErasedRawStridedRef<'_>,
728    lo: &ErasedRawStridedRef<'_>,
729    hi: &ErasedRawStridedRef<'_>,
730) -> Result<()> {
731    let dims: [&[usize]; 3] = [x.dims(), lo.dims(), hi.dims()];
732    if prepare_destination(dest, &dims)?.is_none() {
733        return Ok(());
734    }
735    let x = erased_view::<T>(x)?;
736    let lo = erased_view::<T>(lo)?;
737    let hi = erased_view::<T>(hi)?;
738    let mut dest = uninit_view_mut::<T>(dest)?;
739    zip_map3_into(&mut dest, &x, &lo, &hi, |x, lo, hi| {
740        MaybeUninit::new(T::clamp(x, lo, hi))
741    })
742}
743
744fn uninit_broadcast_mul<T>(
745    dest: &mut ErasedRawStridedUninitMut<'_>,
746    lhs: &ErasedRawStridedRef<'_>,
747    lhs_axes: &[usize],
748    rhs: &ErasedRawStridedRef<'_>,
749    rhs_axes: &[usize],
750) -> Result<()>
751where
752    T: Copy + crate::MaybeSendSync + KernelStorageElement + core::ops::Mul<Output = T>,
753{
754    let lhs = erased_view::<T>(lhs)?;
755    let rhs = erased_view::<T>(rhs)?;
756    let mut dest = uninit_view_mut::<T>(dest)?;
757    broadcast_mul_into_uninit(&mut dest, &lhs, lhs_axes, &rhs, rhs_axes)
758}
759
760fn uninit_broadcast_mul_wrapping<T>(
761    dest: &mut ErasedRawStridedUninitMut<'_>,
762    lhs: &ErasedRawStridedRef<'_>,
763    lhs_axes: &[usize],
764    rhs: &ErasedRawStridedRef<'_>,
765    rhs_axes: &[usize],
766) -> Result<()>
767where
768    T: Copy + crate::MaybeSendSync + KernelStorageElement,
769    Wrapping<T>: core::ops::Mul<Output = Wrapping<T>> + crate::MaybeSendSync,
770{
771    let lhs_data = wrapping_slice(lhs.data_as::<T>()?);
772    let rhs_data = wrapping_slice(rhs.data_as::<T>()?);
773    // SAFETY: the descriptors validated every reachable offset at construction,
774    // and the reinterpreted slices keep their element count.
775    let lhs = unsafe {
776        StridedView::<Wrapping<T>, Identity>::new_unchecked(
777            lhs_data,
778            lhs.dims(),
779            lhs.strides(),
780            lhs.offset(),
781        )
782    };
783    // SAFETY: as above.
784    let rhs = unsafe {
785        StridedView::<Wrapping<T>, Identity>::new_unchecked(
786            rhs_data,
787            rhs.dims(),
788            rhs.strides(),
789            rhs.offset(),
790        )
791    };
792    let dims = dest.dims();
793    let strides = dest.strides();
794    let offset = dest.offset();
795    let data = dest.data_as_uninit_mut::<T>()?;
796    // SAFETY: `Wrapping<T>` is `repr(transparent)` over `T`, so
797    // `MaybeUninit<Wrapping<T>>` has the layout of `MaybeUninit<T>`, and the
798    // slice keeps its element count and exclusive borrow.
799    let data = unsafe {
800        core::slice::from_raw_parts_mut(
801            data.as_mut_ptr().cast::<MaybeUninit<Wrapping<T>>>(),
802            data.len(),
803        )
804    };
805    // SAFETY: the uninit descriptor validated every reachable offset.
806    let mut dest = unsafe { StridedViewMut::new_unchecked(data, dims, strides, offset) };
807    broadcast_mul_into_uninit(&mut dest, &lhs, lhs_axes, &rhs, rhs_axes)
808}
809
810fn wrapping_slice<T>(data: &[T]) -> &[Wrapping<T>] {
811    // SAFETY: `Wrapping<T>` is `repr(transparent)` over `T`; the slice keeps
812    // its element count and shared borrow.
813    unsafe { core::slice::from_raw_parts(data.as_ptr().cast::<Wrapping<T>>(), data.len()) }
814}