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#[derive(Clone, Debug)]
12pub struct ErasedFusedPlan {
13 dtype: KernelDType,
14 plan: FusedPlan,
15}
16
17impl ErasedFusedPlan {
18 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 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 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 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 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 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 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}