Skip to main content

strided_fused/
erased.rs

1use crate::*;
2use num_complex::{Complex32, Complex64};
3use strided_basic::execution::{check_dtype, erased_view, validate_uninit_no_overlap};
4const ERASED_FUSED_INPUT_LIMIT: usize = 4;
5
6/// Dtype-erased single-output wrapper around [`FusedPlan`].
7///
8/// This is the erased replay boundary for unary map and zip-map elementwise
9/// families. It supports the same runtime op-code vocabulary as [`FusedPlan`],
10/// but only for the scalar dtypes currently implementing [`FusedScalar`].
11#[derive(Clone, Debug)]
12pub struct ErasedFusedPlan {
13    dtype: KernelDType,
14    plan: FusedPlan,
15}
16
17impl ErasedFusedPlan {
18    /// Validate and store a single-output fused elementwise plan for one dtype.
19    pub fn compile(dtype: KernelDType, plan: FusedPlan) -> Result<Self> {
20        check_fused_dtype(dtype)?;
21        if plan.input_count == 0 || plan.input_count > ERASED_FUSED_INPUT_LIMIT {
22            return Err(StridedError::UnsupportedArity {
23                arity: plan.input_count,
24                max: ERASED_FUSED_INPUT_LIMIT,
25            });
26        }
27        if plan.outputs.len() != 1 {
28            return Err(StridedError::RankMismatch(plan.outputs.len(), 1));
29        }
30        validate_fused_plan_for_dtype(dtype, &plan)?;
31        Ok(Self { dtype, plan })
32    }
33
34    #[inline]
35    pub fn dtype(&self) -> KernelDType {
36        self.dtype
37    }
38
39    #[inline]
40    pub fn plan(&self) -> &FusedPlan {
41        &self.plan
42    }
43
44    /// Execute a single-output fused elementwise plan through erased descriptors.
45    pub fn execute(
46        &self,
47        ctx: &ExecContext,
48        dest: &mut ErasedRawStridedMut<'_>,
49        inputs: &[ErasedRawStridedRef<'_>],
50    ) -> Result<()> {
51        if inputs.len() != self.plan.input_count {
52            return Err(StridedError::RankMismatch(
53                inputs.len(),
54                self.plan.input_count,
55            ));
56        }
57        check_dtype(self.dtype, dest.dtype())?;
58        for input in inputs {
59            check_dtype(self.dtype, input.dtype())?;
60        }
61
62        let result = match self.dtype {
63            KernelDType::F32 => execute_fused::<f32>(&self.plan, ctx, dest, inputs),
64            KernelDType::F64 => execute_fused::<f64>(&self.plan, ctx, dest, inputs),
65            KernelDType::I32 => execute_fused::<i32>(&self.plan, ctx, dest, inputs),
66            KernelDType::I64 => execute_fused::<i64>(&self.plan, ctx, dest, inputs),
67            KernelDType::Bool => execute_fused::<bool>(&self.plan, ctx, dest, inputs),
68            KernelDType::C32 => execute_fused::<Complex32>(&self.plan, ctx, dest, inputs),
69            KernelDType::C64 => execute_fused::<Complex64>(&self.plan, ctx, dest, inputs),
70            _ => Err(StridedError::UnsupportedDType {
71                dtype: self.dtype.label(),
72            }),
73        };
74        result
75    }
76
77    /// Execute a single-output fused plan into fully overwritten uninitialized storage.
78    ///
79    /// Dtype, shape, destination injectivity, bounds, and input/output overlap
80    /// are validated before any shared typed input descriptor is formed or any
81    /// destination byte is written. On `Ok(())`, every logical destination
82    /// element is initialized. An error leaves the destination untouched; a
83    /// panic during execution may leave partial initialization, but the backing
84    /// `MaybeUninit` storage remains safe to drop.
85    ///
86    /// # Errors
87    ///
88    /// Returns a typed dtype, input-count, shape, bounds, destination
89    /// injectivity, unsupported-operation, or input/output-overlap error. All
90    /// error-producing validation completes before execution starts.
91    /// On success, every reachable destination slot is fully overwritten;
92    /// unreachable holes are neither read nor initialized. Validation errors
93    /// are returned before any destination write. A panic during execution
94    /// may leave a partially initialized `MaybeUninit` destination, which is
95    /// still safely droppable; no readable value is promised for unwritten
96    /// reachable slots.
97    pub fn execute_uninit(
98        &self,
99        ctx: &ExecContext,
100        dest: &mut ErasedRawStridedUninitMut<'_>,
101        inputs: &[ErasedRawStridedPtr<'_>],
102    ) -> Result<()> {
103        if inputs.len() != self.plan.input_count {
104            return Err(StridedError::RankMismatch(
105                inputs.len(),
106                self.plan.input_count,
107            ));
108        }
109        check_dtype(self.dtype, dest.dtype())?;
110        for input in inputs {
111            check_dtype(self.dtype, input.dtype())?;
112        }
113        for (index, input) in inputs.iter().enumerate() {
114            validate_uninit_no_overlap(dest, input, index)?;
115            if input.dims() != dest.dims() {
116                return Err(StridedError::ShapeMismatch(
117                    input.dims().to_vec(),
118                    dest.dims().to_vec(),
119                ));
120            }
121        }
122        let validated = strided_basic::execution::validate_destination_layout_without_alloc(
123            dest.dims(),
124            dest.strides(),
125        )?;
126
127        let run = |dest: &mut ErasedRawStridedUninitMut<'_>| match self.dtype {
128            KernelDType::F32 => execute_fused_uninit_ptrs::<f32>(
129                &self.plan,
130                dest,
131                inputs,
132                ctx.is_serial(),
133                validated,
134            ),
135            KernelDType::F64 => execute_fused_uninit_ptrs::<f64>(
136                &self.plan,
137                dest,
138                inputs,
139                ctx.is_serial(),
140                validated,
141            ),
142            KernelDType::I32 => execute_fused_uninit_ptrs::<i32>(
143                &self.plan,
144                dest,
145                inputs,
146                ctx.is_serial(),
147                validated,
148            ),
149            KernelDType::I64 => execute_fused_uninit_ptrs::<i64>(
150                &self.plan,
151                dest,
152                inputs,
153                ctx.is_serial(),
154                validated,
155            ),
156            KernelDType::Bool => execute_fused_uninit_ptrs::<bool>(
157                &self.plan,
158                dest,
159                inputs,
160                ctx.is_serial(),
161                validated,
162            ),
163            KernelDType::C32 => execute_fused_uninit_ptrs::<Complex32>(
164                &self.plan,
165                dest,
166                inputs,
167                ctx.is_serial(),
168                validated,
169            ),
170            KernelDType::C64 => execute_fused_uninit_ptrs::<Complex64>(
171                &self.plan,
172                dest,
173                inputs,
174                ctx.is_serial(),
175                validated,
176            ),
177            _ => Err(StridedError::UnsupportedDType {
178                dtype: self.dtype.label(),
179            }),
180        };
181        if ctx.is_serial() {
182            run(dest)
183        } else {
184            ctx.run(|| run(dest))
185        }
186    }
187}
188
189fn check_fused_dtype(dtype: KernelDType) -> Result<()> {
190    match dtype {
191        KernelDType::F32
192        | KernelDType::F64
193        | KernelDType::I32
194        | KernelDType::I64
195        | KernelDType::Bool
196        | KernelDType::C32
197        | KernelDType::C64 => Ok(()),
198        _ => Err(StridedError::UnsupportedDType {
199            dtype: dtype.label(),
200        }),
201    }
202}
203
204fn validate_fused_plan_for_dtype(dtype: KernelDType, plan: &FusedPlan) -> Result<()> {
205    match dtype {
206        KernelDType::F32 => {
207            crate::fused::validate_plan_for_scalar::<f32>(plan, plan.input_count, 1)
208        }
209        KernelDType::F64 => {
210            crate::fused::validate_plan_for_scalar::<f64>(plan, plan.input_count, 1)
211        }
212        KernelDType::I32 => {
213            crate::fused::validate_plan_for_scalar::<i32>(plan, plan.input_count, 1)
214        }
215        KernelDType::I64 => {
216            crate::fused::validate_plan_for_scalar::<i64>(plan, plan.input_count, 1)
217        }
218        KernelDType::Bool => {
219            crate::fused::validate_plan_for_scalar::<bool>(plan, plan.input_count, 1)
220        }
221        KernelDType::C32 => {
222            crate::fused::validate_plan_for_scalar::<Complex32>(plan, plan.input_count, 1)
223        }
224        KernelDType::C64 => {
225            crate::fused::validate_plan_for_scalar::<Complex64>(plan, plan.input_count, 1)
226        }
227        _ => Err(StridedError::UnsupportedDType {
228            dtype: dtype.label(),
229        }),
230    }
231}
232
233fn execute_fused<T>(
234    plan: &FusedPlan,
235    ctx: &ExecContext,
236    dest: &mut ErasedRawStridedMut<'_>,
237    inputs: &[ErasedRawStridedRef<'_>],
238) -> Result<()>
239where
240    T: FusedScalar + KernelStorageElement,
241{
242    let dest_dims = dest.dims();
243    let dest_strides = dest.strides();
244    let dest_offset = dest.offset();
245    let dest_data = dest.data_as_mut::<T>()?;
246    let dest_view =
247        unsafe { StridedViewMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
248
249    match inputs {
250        [a] => {
251            let input_views = [erased_view::<T>(a)?];
252            let mut dests = [dest_view];
253            execute_fused_views(ctx, &mut dests, &input_views, plan)
254        }
255        [a, b] => {
256            let input_views = [erased_view::<T>(a)?, erased_view::<T>(b)?];
257            let mut dests = [dest_view];
258            execute_fused_views(ctx, &mut dests, &input_views, plan)
259        }
260        [a, b, c] => {
261            let input_views = [
262                erased_view::<T>(a)?,
263                erased_view::<T>(b)?,
264                erased_view::<T>(c)?,
265            ];
266            let mut dests = [dest_view];
267            execute_fused_views(ctx, &mut dests, &input_views, plan)
268        }
269        [a, b, c, d] => {
270            let input_views = [
271                erased_view::<T>(a)?,
272                erased_view::<T>(b)?,
273                erased_view::<T>(c)?,
274                erased_view::<T>(d)?,
275            ];
276            let mut dests = [dest_view];
277            execute_fused_views(ctx, &mut dests, &input_views, plan)
278        }
279        _ => Err(StridedError::UnsupportedArity {
280            arity: inputs.len(),
281            max: ERASED_FUSED_INPUT_LIMIT,
282        }),
283    }
284}
285
286fn execute_fused_uninit<T>(
287    plan: &FusedPlan,
288    dest: &mut ErasedRawStridedUninitMut<'_>,
289    inputs: &[ErasedRawStridedRef<'_>],
290    serial: bool,
291    validated: strided_basic::execution::ValidatedDestinationLayout,
292) -> Result<()>
293where
294    T: FusedScalar + KernelStorageElement,
295{
296    let dims = dest.dims();
297    let strides = dest.strides();
298    let offset = dest.offset();
299    let dest_data = dest.data_as_uninit_mut::<T>()?;
300    let mut dest_view = unsafe { StridedViewMut::new_unchecked(dest_data, dims, strides, offset) };
301    match inputs {
302        [a] => {
303            let input_views = [erased_view::<T>(a)?];
304            crate::fused::fused_elementwise_into_uninit(
305                &mut dest_view,
306                &input_views,
307                plan,
308                serial,
309                validated,
310            )
311        }
312        [a, b] => {
313            let input_views = [erased_view::<T>(a)?, erased_view::<T>(b)?];
314            crate::fused::fused_elementwise_into_uninit(
315                &mut dest_view,
316                &input_views,
317                plan,
318                serial,
319                validated,
320            )
321        }
322        [a, b, c] => {
323            let input_views = [
324                erased_view::<T>(a)?,
325                erased_view::<T>(b)?,
326                erased_view::<T>(c)?,
327            ];
328            crate::fused::fused_elementwise_into_uninit(
329                &mut dest_view,
330                &input_views,
331                plan,
332                serial,
333                validated,
334            )
335        }
336        [a, b, c, d] => {
337            let input_views = [
338                erased_view::<T>(a)?,
339                erased_view::<T>(b)?,
340                erased_view::<T>(c)?,
341                erased_view::<T>(d)?,
342            ];
343            crate::fused::fused_elementwise_into_uninit(
344                &mut dest_view,
345                &input_views,
346                plan,
347                serial,
348                validated,
349            )
350        }
351        _ => Err(StridedError::UnsupportedArity {
352            arity: inputs.len(),
353            max: ERASED_FUSED_INPUT_LIMIT,
354        }),
355    }
356}
357
358fn execute_fused_uninit_ptrs<T>(
359    plan: &FusedPlan,
360    dest: &mut ErasedRawStridedUninitMut<'_>,
361    inputs: &[ErasedRawStridedPtr<'_>],
362    serial: bool,
363    validated: strided_basic::execution::ValidatedDestinationLayout,
364) -> Result<()>
365where
366    T: FusedScalar + KernelStorageElement,
367{
368    match inputs {
369        [a] => {
370            // SAFETY: ErasedFusedPlan::execute_uninit checked every input against the destination before this helper.
371            let refs = unsafe { [a.try_as_ref_after_no_overlap()?] };
372            execute_fused_uninit::<T>(plan, dest, &refs, serial, validated)
373        }
374        [a, b] => {
375            // SAFETY: ErasedFusedPlan::execute_uninit checked every input against the destination before this helper.
376            let refs = unsafe {
377                [
378                    a.try_as_ref_after_no_overlap()?,
379                    b.try_as_ref_after_no_overlap()?,
380                ]
381            };
382            execute_fused_uninit::<T>(plan, dest, &refs, serial, validated)
383        }
384        [a, b, c] => {
385            // SAFETY: ErasedFusedPlan::execute_uninit checked every input against the destination before this helper.
386            let refs = unsafe {
387                [
388                    a.try_as_ref_after_no_overlap()?,
389                    b.try_as_ref_after_no_overlap()?,
390                    c.try_as_ref_after_no_overlap()?,
391                ]
392            };
393            execute_fused_uninit::<T>(plan, dest, &refs, serial, validated)
394        }
395        [a, b, c, d] => {
396            // SAFETY: ErasedFusedPlan::execute_uninit checked every input against the destination before this helper.
397            let refs = unsafe {
398                [
399                    a.try_as_ref_after_no_overlap()?,
400                    b.try_as_ref_after_no_overlap()?,
401                    c.try_as_ref_after_no_overlap()?,
402                    d.try_as_ref_after_no_overlap()?,
403                ]
404            };
405            execute_fused_uninit::<T>(plan, dest, &refs, serial, validated)
406        }
407        _ => Err(StridedError::UnsupportedArity {
408            arity: inputs.len(),
409            max: ERASED_FUSED_INPUT_LIMIT,
410        }),
411    }
412}
413
414fn execute_fused_views<T>(
415    ctx: &ExecContext,
416    dests: &mut [StridedViewMut<'_, T>],
417    inputs: &[StridedView<'_, T>],
418    plan: &FusedPlan,
419) -> Result<()>
420where
421    T: FusedScalar + KernelStorageElement,
422{
423    if ctx.is_serial() {
424        crate::fused::fused_elementwise_into_serial(dests, inputs, plan)
425    } else {
426        ctx.run(|| fused_elementwise_into(dests, inputs, plan))
427    }
428}