Skip to main content

tenferro_cpu_fused/
lib.rs

1#![doc(hidden)]
2
3//! Fused CPU elementwise execution.
4
5/// Internal result alias for the fused CPU adapter.
6///
7/// # Examples
8///
9/// ```rust
10/// use tenferro_cpu_fused::Result;
11/// let result: Result<()> = Ok(());
12/// assert!(result.is_ok());
13/// ```
14pub 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
34/// Largest input count `strided_fused::ErasedFusedPlan` accepts.
35///
36/// [`elementwise_fusion_with_pool`] declines larger plans, and the CPU runtime
37/// engine reports this limit so the prepared planner does not plan such
38/// regions at all.
39///
40/// # Examples
41///
42/// ```rust
43/// assert_eq!(tenferro_cpu_fused::ERASED_FUSION_MAX_INPUTS, 4);
44/// ```
45pub 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        // INVARIANT: `plan_uses_unfused_op` declines every plan containing
110        // these ops before a strided plan is built; strided_fused has no
111        // remainder or `erf` instruction.
112        ElementwiseFusionOp::Remainder | ElementwiseFusionOp::Erf => {
113            unreachable!("ops without a strided_fused instruction are declined before CPU fusion")
114        }
115    }
116}
117
118/// Whether the plan contains an op that `strided_fused` cannot replay.
119///
120/// `erf` has no `FusedOp` instruction yet, so a region containing it is
121/// declined and its ops run unfused.
122fn 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        // INVARIANT: `KernelDType` is the fixed compiled-kernel vocabulary used by
174        // the fused path; an externally defined scalar has no entry in it and is
175        // rejected by `dtype_supports_erased_fusion` before this point.
176        DType::External(_) => unreachable!("KernelDType covers the preset scalars"),
177    }
178}
179
180fn typed_bytes<T>(data: &[T]) -> &[u8] {
181    // SAFETY: `data` is an aligned typed slice. The returned byte slice has
182    // the same lifetime and exact byte length, and is read-only.
183    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        // A caller-owned payload is opaque here, so the fused path rejects it
208        // instead of reading bytes it cannot interpret.
209        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}
216/// The typed tensor behind `input`, or the refusal this path produces for one.
217///
218/// Callers reach this from a match on `input.dtype()`, so `None` means the tag table
219/// and the runtime dtype disagree rather than a caller mistake.
220fn 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            // SAFETY: fusion inputs are initialized typed storage with matching
346            // dtype and alignment; validated layouts bound every reachable read
347            // for the retained input borrow.
348            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        // An externally defined scalar has no fused kernel, so the fixed-dtype
378        // fused path rejects it rather than guessing a representation.
379        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    // SAFETY: `out` exclusively owns the output allocation with matching
538    // dtype/alignment and the fused plan overwrites every reachable element
539    // before `assume_init` exposes typed storage.
540    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    // SAFETY: the fused replay writes every logical destination element and retains no destination view.
548    Ok(wrap(unsafe { out.assume_init()? }))
549}
550
551#[cfg(test)]
552mod tests;