1#![doc(hidden)]
2
3pub type Result<T> = tenferro_tensor::Result<T>;
15pub use tenferro_tensor::{DType, Error, Tensor, TypedTensor};
16
17use num_complex::{Complex32, Complex64};
18use std::mem::size_of_val;
19use strided_basic::{ErasedRawStridedPtr, ErasedRawStridedRef, ExecContext, KernelDType};
20use strided_fused::{ErasedFusedPlan, FusedInst, FusedOp, FusedPlan};
21use tenferro_cpu_basic::{
22 erased_raw_strided_ref, erased_raw_strided_uninit_mut, typed_host_data, BufferPool, PoolScalar,
23 PooledUninitOutput,
24};
25use tenferro_tensor::backend::{
26 ElementwiseFusionInputView, ElementwiseFusionOp, ElementwiseFusionPlan,
27};
28use tenferro_tensor::col_major_strides;
29
30const ELEMENTWISE_FUSION_OP: &str = "execute_elementwise_fusion";
31
32const ELEMENTWISE_FUSION_MIN_ELEMENTS: usize = 16 * 1024;
33
34pub const ERASED_FUSION_MAX_INPUTS: usize = 4;
46
47fn validate_elementwise_fusion_inputs(
48 inputs: &[&Tensor],
49 plan: &ElementwiseFusionPlan,
50) -> crate::Result<bool> {
51 if inputs.len() != plan.input_count() {
52 return Err(crate::Error::invalid_argument(
53 ELEMENTWISE_FUSION_OP,
54 "inputs",
55 format!(
56 "plan expects {} inputs but backend received {}",
57 plan.input_count(),
58 inputs.len()
59 ),
60 ));
61 }
62 if plan.input_views().len() != plan.input_count() {
63 return Err(crate::Error::invalid_argument(
64 ELEMENTWISE_FUSION_OP,
65 "input_views",
66 format!(
67 "plan has {} input views for {} inputs",
68 plan.input_views().len(),
69 plan.input_count()
70 ),
71 ));
72 }
73 if plan.outputs().is_empty() {
74 return Ok(false);
75 }
76 for input in inputs {
77 if input.dtype() != plan.dtype() {
78 return Err(crate::Error::dtype_mismatch(
79 ELEMENTWISE_FUSION_OP,
80 input.dtype(),
81 plan.dtype(),
82 ));
83 }
84 }
85 Ok(true)
86}
87
88fn strided_fused_op(op: ElementwiseFusionOp) -> FusedOp {
89 match op {
90 ElementwiseFusionOp::Add => FusedOp::Add,
91 ElementwiseFusionOp::Multiply => FusedOp::Multiply,
92 ElementwiseFusionOp::Negate => FusedOp::Negate,
93 ElementwiseFusionOp::Conj => FusedOp::Conj,
94 ElementwiseFusionOp::Divide => FusedOp::Divide,
95 ElementwiseFusionOp::Abs => FusedOp::Abs,
96 ElementwiseFusionOp::Maximum => FusedOp::Maximum,
97 ElementwiseFusionOp::Minimum => FusedOp::Minimum,
98 ElementwiseFusionOp::Clamp => FusedOp::Clamp,
99 ElementwiseFusionOp::Exp => FusedOp::Exp,
100 ElementwiseFusionOp::Log => FusedOp::Log,
101 ElementwiseFusionOp::Sin => FusedOp::Sin,
102 ElementwiseFusionOp::Cos => FusedOp::Cos,
103 ElementwiseFusionOp::Tanh => FusedOp::Tanh,
104 ElementwiseFusionOp::Sqrt => FusedOp::Sqrt,
105 ElementwiseFusionOp::Rsqrt => FusedOp::Rsqrt,
106 ElementwiseFusionOp::Pow => FusedOp::Pow,
107 ElementwiseFusionOp::Expm1 => FusedOp::Expm1,
108 ElementwiseFusionOp::Log1p => FusedOp::Log1p,
109 ElementwiseFusionOp::Remainder | ElementwiseFusionOp::Erf => {
113 unreachable!("ops without a strided_fused instruction are declined before CPU fusion")
114 }
115 }
116}
117
118fn plan_uses_unfused_op(plan: &ElementwiseFusionPlan) -> bool {
123 plan.ops().iter().any(|inst| {
124 matches!(
125 inst.op(),
126 ElementwiseFusionOp::Remainder | ElementwiseFusionOp::Erf
127 )
128 })
129}
130
131fn plan_uses_ordered_op(plan: &ElementwiseFusionPlan) -> bool {
132 plan.ops().iter().any(|inst| {
133 matches!(
134 inst.op(),
135 ElementwiseFusionOp::Maximum
136 | ElementwiseFusionOp::Minimum
137 | ElementwiseFusionOp::Clamp
138 )
139 })
140}
141
142fn should_defer_to_broadcast_multiply_special_case(plan: &ElementwiseFusionPlan) -> bool {
143 !plan.input_views().iter().all(|view| view.is_identity())
144 && plan.ops().len() == 1
145 && plan.outputs() == [plan.input_count()]
146 && plan.ops()[0].op() == ElementwiseFusionOp::Multiply
147}
148
149fn single_output_strided_fused_plan(plan: &ElementwiseFusionPlan, output: usize) -> FusedPlan {
150 FusedPlan {
151 input_count: plan.input_count(),
152 outputs: vec![output],
153 ops: plan
154 .ops()
155 .iter()
156 .map(|inst| FusedInst {
157 op: strided_fused_op(inst.op()),
158 inputs: inst.inputs().to_vec(),
159 })
160 .collect(),
161 }
162}
163
164fn kernel_dtype(dtype: DType) -> KernelDType {
165 match dtype {
166 DType::F32 => KernelDType::F32,
167 DType::F64 => KernelDType::F64,
168 DType::I32 => KernelDType::I32,
169 DType::I64 => KernelDType::I64,
170 DType::Bool => KernelDType::Bool,
171 DType::C32 => KernelDType::C32,
172 DType::C64 => KernelDType::C64,
173 DType::External(_) => unreachable!("KernelDType covers the preset scalars"),
177 }
178}
179
180fn typed_bytes<T>(data: &[T]) -> &[u8] {
181 unsafe { std::slice::from_raw_parts(data.as_ptr().cast::<u8>(), size_of_val(data)) }
184}
185
186struct ErasedFusionInput<'a> {
187 data: &'a [u8],
188 dims: Vec<usize>,
189 strides: Vec<isize>,
190}
191
192fn tensor_host_bytes<'a>(op: &'static str, input: &'a Tensor) -> crate::Result<&'a [u8]> {
193 macro_rules! bytes {
194 ($tensor:expr) => {
195 typed_host_data(op, $tensor).map(typed_bytes)
196 };
197 }
198
199 match input.dtype() {
200 DType::F32 => bytes!(fused_host::<f32>(op, input)?),
201 DType::F64 => bytes!(fused_host::<f64>(op, input)?),
202 DType::I32 => bytes!(fused_host::<i32>(op, input)?),
203 DType::I64 => bytes!(fused_host::<i64>(op, input)?),
204 DType::Bool => bytes!(fused_host::<bool>(op, input)?),
205 DType::C32 => bytes!(fused_host::<Complex32>(op, input)?),
206 DType::C64 => bytes!(fused_host::<Complex64>(op, input)?),
207 DType::External(_) => Err(crate::Error::unsupported_dtype(
210 op,
211 input.dtype(),
212 "an externally defined payload is not a fused input",
213 )),
214 }
215}
216fn fused_host<'a, T: tenferro_tensor::TensorScalar>(
221 op: &'static str,
222 input: &'a Tensor,
223) -> crate::Result<&'a TypedTensor<T>> {
224 input.as_typed::<T>().ok_or_else(|| {
225 crate::Error::unsupported_dtype(
226 op,
227 input.dtype(),
228 "an externally defined payload is not a fused input",
229 )
230 })
231}
232
233fn erased_fusion_input<'a>(
234 input: &'a Tensor,
235 view: &ElementwiseFusionInputView,
236) -> crate::Result<ErasedFusionInput<'a>> {
237 let data = tensor_host_bytes(ELEMENTWISE_FUSION_OP, input)?;
238 let base_shape = input.shape();
239 let base_strides = col_major_strides(base_shape)?;
240 let ElementwiseFusionInputView::BroadcastInDim { shape, dims } = view else {
241 return Ok(ErasedFusionInput {
242 data,
243 dims: base_shape.to_vec(),
244 strides: base_strides,
245 });
246 };
247
248 if dims.len() != base_shape.len() {
249 return Err(crate::Error::invalid_argument(
250 ELEMENTWISE_FUSION_OP,
251 "configuration",
252 format!(
253 "broadcast dims length {} does not match input rank {}",
254 dims.len(),
255 base_shape.len()
256 ),
257 ));
258 }
259
260 let mut strides = vec![0; shape.len()];
261 let mut seen = vec![false; shape.len()];
262 for (source_axis, &target_axis) in dims.iter().enumerate() {
263 if target_axis >= shape.len() {
264 return Err(crate::Error::axis_out_of_bounds(
265 ELEMENTWISE_FUSION_OP,
266 target_axis,
267 shape.len(),
268 ));
269 }
270 if seen[target_axis] {
271 return Err(crate::Error::duplicate_axis(
272 ELEMENTWISE_FUSION_OP,
273 target_axis,
274 "broadcast dims",
275 ));
276 }
277 seen[target_axis] = true;
278 let source_dim = base_shape[source_axis];
279 let target_dim = shape[target_axis];
280 if source_dim != target_dim && source_dim != 1 {
281 return Err(crate::Error::shape_mismatch(
282 ELEMENTWISE_FUSION_OP,
283 shape.to_vec(),
284 base_shape.to_vec(),
285 ));
286 }
287 if source_dim == target_dim {
288 strides[target_axis] = base_strides[source_axis];
289 }
290 }
291
292 Ok(ErasedFusionInput {
293 data,
294 dims: shape.to_vec(),
295 strides,
296 })
297}
298
299#[doc(hidden)]
300pub fn elementwise_fusion_with_pool(
301 buffers: &mut BufferPool,
302 exec_context: &ExecContext,
303 inputs: &[&Tensor],
304 plan: &ElementwiseFusionPlan,
305) -> crate::Result<Option<Vec<Tensor>>> {
306 if !validate_elementwise_fusion_inputs(inputs, plan)? {
307 return Ok(None);
308 }
309 if inputs.is_empty() || inputs.len() > ERASED_FUSION_MAX_INPUTS {
310 return Ok(None);
311 }
312 if plan_uses_unfused_op(plan) {
313 return Ok(None);
314 }
315 if should_defer_to_broadcast_multiply_special_case(plan) {
316 return Ok(None);
317 }
318 if !dtype_supports_erased_fusion(plan.dtype(), plan) {
319 return Ok(None);
320 }
321
322 let input_layouts = inputs
323 .iter()
324 .zip(plan.input_views())
325 .map(|(input, view)| erased_fusion_input(input, view))
326 .collect::<crate::Result<Vec<_>>>()?;
327 let shape = input_layouts[0].dims.clone();
328 if input_layouts
329 .iter()
330 .skip(1)
331 .any(|input| input.dims != shape)
332 {
333 return Ok(None);
334 }
335 let element_count =
336 tenferro_tensor::validate::checked_shape_product(ELEMENTWISE_FUSION_OP, "shape", &shape)?;
337 if element_count < ELEMENTWISE_FUSION_MIN_ELEMENTS {
338 return Ok(None);
339 }
340
341 let dtype = kernel_dtype(plan.dtype());
342 let input_refs = input_layouts
343 .iter()
344 .map(|input| {
345 unsafe { erased_raw_strided_ref(dtype, input.data, &input.dims, &input.strides, 0) }
349 .map_err(|err| crate::Error::backend_source(ELEMENTWISE_FUSION_OP, err))
350 })
351 .collect::<crate::Result<Vec<_>>>()?;
352
353 execute_erased_fused_outputs(buffers, exec_context, dtype, &input_refs, &shape, plan).map(Some)
354}
355
356fn dtype_supports_erased_fusion(dtype: DType, plan: &ElementwiseFusionPlan) -> bool {
357 match dtype {
358 DType::F32 | DType::F64 => true,
359 DType::C32 | DType::C64 => !plan_uses_ordered_op(plan),
360 DType::I32 | DType::I64 => plan.ops().iter().all(|inst| {
361 matches!(
362 inst.op(),
363 ElementwiseFusionOp::Add
364 | ElementwiseFusionOp::Multiply
365 | ElementwiseFusionOp::Negate
366 | ElementwiseFusionOp::Conj
367 | ElementwiseFusionOp::Abs
368 | ElementwiseFusionOp::Maximum
369 | ElementwiseFusionOp::Minimum
370 | ElementwiseFusionOp::Clamp
371 )
372 }),
373 DType::Bool => plan
374 .ops()
375 .iter()
376 .all(|inst| inst.op() == ElementwiseFusionOp::Conj),
377 DType::External(_) => false,
380 }
381}
382
383fn execute_erased_fused_outputs(
384 buffers: &mut BufferPool,
385 exec_context: &ExecContext,
386 dtype: KernelDType,
387 input_refs: &[ErasedRawStridedRef<'_>],
388 shape: &[usize],
389 plan: &ElementwiseFusionPlan,
390) -> crate::Result<Vec<Tensor>> {
391 let input_ptrs: Vec<_> = input_refs
392 .iter()
393 .map(ErasedRawStridedPtr::from_ref)
394 .collect();
395 match dtype {
396 KernelDType::F32 => plan
397 .outputs()
398 .iter()
399 .map(|&output| {
400 execute_erased_fused_output::<f32>(
401 buffers,
402 exec_context,
403 dtype,
404 &input_ptrs,
405 shape,
406 plan,
407 output,
408 Tensor::from_typed::<f32>,
409 )
410 })
411 .collect(),
412 KernelDType::F64 => plan
413 .outputs()
414 .iter()
415 .map(|&output| {
416 execute_erased_fused_output::<f64>(
417 buffers,
418 exec_context,
419 dtype,
420 &input_ptrs,
421 shape,
422 plan,
423 output,
424 Tensor::from_typed::<f64>,
425 )
426 })
427 .collect(),
428 KernelDType::I32 => plan
429 .outputs()
430 .iter()
431 .map(|&output| {
432 execute_erased_fused_output::<i32>(
433 buffers,
434 exec_context,
435 dtype,
436 &input_ptrs,
437 shape,
438 plan,
439 output,
440 Tensor::from_typed::<i32>,
441 )
442 })
443 .collect(),
444 KernelDType::I64 => plan
445 .outputs()
446 .iter()
447 .map(|&output| {
448 execute_erased_fused_output::<i64>(
449 buffers,
450 exec_context,
451 dtype,
452 &input_ptrs,
453 shape,
454 plan,
455 output,
456 Tensor::from_typed::<i64>,
457 )
458 })
459 .collect(),
460 KernelDType::Bool => plan
461 .outputs()
462 .iter()
463 .map(|&output| {
464 execute_erased_fused_output::<bool>(
465 buffers,
466 exec_context,
467 dtype,
468 &input_ptrs,
469 shape,
470 plan,
471 output,
472 Tensor::from_typed::<bool>,
473 )
474 })
475 .collect(),
476 KernelDType::C32 => plan
477 .outputs()
478 .iter()
479 .map(|&output| {
480 execute_erased_fused_output::<num_complex::Complex32>(
481 buffers,
482 exec_context,
483 dtype,
484 &input_ptrs,
485 shape,
486 plan,
487 output,
488 Tensor::from_typed::<tenferro_tensor::Complex32>,
489 )
490 })
491 .collect(),
492 KernelDType::C64 => plan
493 .outputs()
494 .iter()
495 .map(|&output| {
496 execute_erased_fused_output::<num_complex::Complex64>(
497 buffers,
498 exec_context,
499 dtype,
500 &input_ptrs,
501 shape,
502 plan,
503 output,
504 Tensor::from_typed::<tenferro_tensor::Complex64>,
505 )
506 })
507 .collect(),
508 _ => Err(crate::Error::unsupported(
509 ELEMENTWISE_FUSION_OP,
510 format!(
511 "unsupported dtype {}; supported dtypes: F32/F64/I32/I64/Bool/C32/C64",
512 dtype.label()
513 ),
514 )),
515 }
516}
517
518#[allow(clippy::too_many_arguments)]
519fn execute_erased_fused_output<T>(
520 buffers: &mut BufferPool,
521 exec_context: &ExecContext,
522 dtype: KernelDType,
523 input_ptrs: &[ErasedRawStridedPtr<'_>],
524 shape: &[usize],
525 plan: &ElementwiseFusionPlan,
526 output: usize,
527 wrap: fn(TypedTensor<T>) -> Tensor,
528) -> crate::Result<Tensor>
529where
530 T: Clone + PoolScalar,
531{
532 let fused_plan = single_output_strided_fused_plan(plan, output);
533 let erased_plan = ErasedFusedPlan::compile(dtype, fused_plan)
534 .map_err(|err| crate::Error::backend_source(ELEMENTWISE_FUSION_OP, err))?;
535 let mut out = PooledUninitOutput::<T>::new(buffers, shape.to_vec())?;
536 let output_strides = col_major_strides(shape)?;
537 let mut dest = unsafe {
541 erased_raw_strided_uninit_mut(dtype, out.as_uninit_bytes_mut(), shape, &output_strides, 0)
542 }
543 .map_err(|err| crate::Error::backend_source(ELEMENTWISE_FUSION_OP, err))?;
544 erased_plan
545 .execute_uninit(exec_context, &mut dest, input_ptrs)
546 .map_err(|err| crate::Error::backend_source(ELEMENTWISE_FUSION_OP, err))?;
547 Ok(wrap(unsafe { out.assume_init()? }))
549}
550
551#[cfg(test)]
552mod tests;