Skip to main content

strided_fused/
fused.rs

1//! Runtime-DAG fused elementwise kernels.
2
3use core::mem::MaybeUninit;
4use strided_basic::execution::is_injective_layout;
5
6use crate::{MaybeSendSync, Result, StridedError, StridedView, StridedViewMut};
7use strided_basic::execution::{
8    build_plan_fused, build_plan_fused_small, ensure_same_shape, for_each_inner_block_preordered,
9    SMALL_TENSOR_THRESHOLD,
10};
11use strided_basic::execution::{
12    map_into_validated, validate_destination_layout_without_alloc, zip_map2_into_validated,
13    zip_map3_into_validated, zip_map4_into_validated, ValidatedDestinationLayout,
14};
15
16#[cfg(feature = "parallel")]
17use strided_basic::execution::compute_costs;
18#[cfg(feature = "parallel")]
19use strided_basic::execution::{mapreduce_threaded, SendPtr, MINTHREADLENGTH};
20
21/// Runtime scalar operation for a fused elementwise plan.
22#[derive(Clone, Copy, Debug, Eq, PartialEq)]
23pub enum FusedOp {
24    Add,
25    Multiply,
26    Negate,
27    Conj,
28    Divide,
29    Abs,
30    Maximum,
31    Minimum,
32    Clamp,
33    Exp,
34    Log,
35    Sin,
36    Cos,
37    Tanh,
38    Sqrt,
39    Rsqrt,
40    Pow,
41    Expm1,
42    Log1p,
43}
44
45impl FusedOp {
46    #[inline]
47    pub const fn label(self) -> &'static str {
48        match self {
49            Self::Add => "add",
50            Self::Multiply => "multiply",
51            Self::Negate => "negate",
52            Self::Conj => "conj",
53            Self::Divide => "divide",
54            Self::Abs => "abs",
55            Self::Maximum => "maximum",
56            Self::Minimum => "minimum",
57            Self::Clamp => "clamp",
58            Self::Exp => "exp",
59            Self::Log => "log",
60            Self::Sin => "sin",
61            Self::Cos => "cos",
62            Self::Tanh => "tanh",
63            Self::Sqrt => "sqrt",
64            Self::Rsqrt => "rsqrt",
65            Self::Pow => "pow",
66            Self::Expm1 => "expm1",
67            Self::Log1p => "log1p",
68        }
69    }
70}
71
72/// One SSA instruction in a [`FusedPlan`].
73#[derive(Clone, Debug, Eq, PartialEq)]
74pub struct FusedInst {
75    pub op: FusedOp,
76    pub inputs: Vec<usize>,
77}
78
79/// Topologically ordered fused elementwise SSA DAG.
80///
81/// Values are numbered in evaluation order. Input values occupy
82/// `0..input_count`; each instruction appends one value after the previous
83/// inputs/instructions. For example, with `input_count == 2`, the first
84/// instruction writes value `2`, the second writes value `3`, and so on.
85/// `outputs` contains the value ids to write to `dests` in order.
86///
87/// All inputs and destinations passed to [`fused_elementwise_into`] must have
88/// the same shape and scalar type. Broadcast inputs should be represented with
89/// `StridedView::broadcast` before building the plan; the fused API does not
90/// perform implicit broadcasting.
91#[derive(Clone, Debug, Eq, PartialEq)]
92pub struct FusedPlan {
93    pub input_count: usize,
94    pub outputs: Vec<usize>,
95    pub ops: Vec<FusedInst>,
96}
97
98/// Scalar types supported by [`fused_elementwise_into`].
99pub trait FusedScalar: Copy + MaybeSendSync + 'static {
100    fn fused_dtype_label() -> &'static str {
101        core::any::type_name::<Self>()
102    }
103
104    fn supports_fused_op(_op: FusedOp) -> bool {
105        true
106    }
107
108    fn fused_add(self, rhs: Self) -> Self;
109    fn fused_multiply(self, rhs: Self) -> Self;
110    fn fused_negate(self) -> Self;
111    fn fused_conj(self) -> Self;
112    fn fused_divide(self, rhs: Self) -> Self;
113    fn fused_abs(self) -> Self;
114    fn fused_maximum(self, rhs: Self) -> Self;
115    fn fused_minimum(self, rhs: Self) -> Self;
116    fn fused_clamp(self, min: Self, max: Self) -> Self;
117    fn fused_exp(self) -> Self;
118    fn fused_log(self) -> Self;
119    fn fused_sin(self) -> Self;
120    fn fused_cos(self) -> Self;
121    fn fused_tanh(self) -> Self;
122    fn fused_sqrt(self) -> Self;
123    fn fused_rsqrt(self) -> Self;
124    fn fused_pow(self, rhs: Self) -> Self;
125    fn fused_expm1(self) -> Self;
126    fn fused_log1p(self) -> Self;
127}
128
129macro_rules! unsupported_fused_op {
130    ($op:literal, $ty:literal) => {
131        unreachable!("unsupported fused op {} for dtype {}", $op, $ty)
132    };
133}
134
135macro_rules! impl_real_fused_scalar {
136    ($ty:ty) => {
137        impl FusedScalar for $ty {
138            #[inline(always)]
139            fn fused_add(self, rhs: Self) -> Self {
140                self + rhs
141            }
142
143            #[inline(always)]
144            fn fused_multiply(self, rhs: Self) -> Self {
145                self * rhs
146            }
147
148            #[inline(always)]
149            fn fused_negate(self) -> Self {
150                -self
151            }
152
153            #[inline(always)]
154            fn fused_conj(self) -> Self {
155                self
156            }
157
158            #[inline(always)]
159            fn fused_divide(self, rhs: Self) -> Self {
160                self / rhs
161            }
162
163            #[inline(always)]
164            fn fused_abs(self) -> Self {
165                self.abs()
166            }
167
168            #[inline(always)]
169            fn fused_maximum(self, rhs: Self) -> Self {
170                self.max(rhs)
171            }
172
173            #[inline(always)]
174            fn fused_minimum(self, rhs: Self) -> Self {
175                self.min(rhs)
176            }
177
178            #[inline(always)]
179            fn fused_clamp(self, min: Self, max: Self) -> Self {
180                self.fused_maximum(min).fused_minimum(max)
181            }
182
183            #[inline(always)]
184            fn fused_exp(self) -> Self {
185                self.exp()
186            }
187
188            #[inline(always)]
189            fn fused_log(self) -> Self {
190                self.ln()
191            }
192
193            #[inline(always)]
194            fn fused_sin(self) -> Self {
195                self.sin()
196            }
197
198            #[inline(always)]
199            fn fused_cos(self) -> Self {
200                self.cos()
201            }
202
203            #[inline(always)]
204            fn fused_tanh(self) -> Self {
205                self.tanh()
206            }
207
208            #[inline(always)]
209            fn fused_sqrt(self) -> Self {
210                self.sqrt()
211            }
212
213            #[inline(always)]
214            fn fused_rsqrt(self) -> Self {
215                1.0 / self.sqrt()
216            }
217
218            #[inline(always)]
219            fn fused_pow(self, rhs: Self) -> Self {
220                self.powf(rhs)
221            }
222
223            #[inline(always)]
224            fn fused_expm1(self) -> Self {
225                self.exp_m1()
226            }
227
228            #[inline(always)]
229            fn fused_log1p(self) -> Self {
230                self.ln_1p()
231            }
232        }
233    };
234}
235
236macro_rules! impl_complex_fused_scalar {
237    ($ty:ty, $div:path) => {
238        impl FusedScalar for $ty {
239            #[inline(always)]
240            fn fused_add(self, rhs: Self) -> Self {
241                self + rhs
242            }
243
244            #[inline(always)]
245            fn fused_multiply(self, rhs: Self) -> Self {
246                self * rhs
247            }
248
249            #[inline(always)]
250            fn fused_negate(self) -> Self {
251                -self
252            }
253
254            #[inline(always)]
255            fn fused_conj(self) -> Self {
256                num_complex::Complex::conj(&self)
257            }
258
259            #[inline(always)]
260            fn fused_divide(self, rhs: Self) -> Self {
261                $div(self, rhs)
262            }
263
264            #[inline(always)]
265            fn fused_abs(self) -> Self {
266                Self::new(self.norm(), 0.0)
267            }
268
269            #[inline(always)]
270            fn fused_maximum(self, rhs: Self) -> Self {
271                if self.norm_sqr() >= rhs.norm_sqr() {
272                    self
273                } else {
274                    rhs
275                }
276            }
277
278            #[inline(always)]
279            fn fused_minimum(self, rhs: Self) -> Self {
280                if self.norm_sqr() <= rhs.norm_sqr() {
281                    self
282                } else {
283                    rhs
284                }
285            }
286
287            #[inline(always)]
288            fn fused_clamp(self, min: Self, max: Self) -> Self {
289                self.fused_maximum(min).fused_minimum(max)
290            }
291
292            #[inline(always)]
293            fn fused_exp(self) -> Self {
294                self.exp()
295            }
296
297            #[inline(always)]
298            fn fused_log(self) -> Self {
299                self.ln()
300            }
301
302            #[inline(always)]
303            fn fused_sin(self) -> Self {
304                self.sin()
305            }
306
307            #[inline(always)]
308            fn fused_cos(self) -> Self {
309                self.cos()
310            }
311
312            #[inline(always)]
313            fn fused_tanh(self) -> Self {
314                self.tanh()
315            }
316
317            #[inline(always)]
318            fn fused_sqrt(self) -> Self {
319                self.sqrt()
320            }
321
322            #[inline(always)]
323            fn fused_rsqrt(self) -> Self {
324                Self::new(1.0, 0.0) / self.sqrt()
325            }
326
327            #[inline(always)]
328            fn fused_pow(self, rhs: Self) -> Self {
329                self.powc(rhs)
330            }
331
332            #[inline(always)]
333            fn fused_expm1(self) -> Self {
334                self.exp() - Self::new(1.0, 0.0)
335            }
336
337            #[inline(always)]
338            fn fused_log1p(self) -> Self {
339                (self + Self::new(1.0, 0.0)).ln()
340            }
341        }
342    };
343}
344
345impl_real_fused_scalar!(f32);
346impl_real_fused_scalar!(f64);
347impl_complex_fused_scalar!(
348    num_complex::Complex32,
349    strided_basic::robust_complex_divide_f32
350);
351impl_complex_fused_scalar!(
352    num_complex::Complex64,
353    strided_basic::robust_complex_divide_f64
354);
355
356macro_rules! impl_signed_integer_fused_scalar {
357    ($ty:ty, $label:literal) => {
358        impl FusedScalar for $ty {
359            #[inline]
360            fn fused_dtype_label() -> &'static str {
361                $label
362            }
363
364            #[inline]
365            fn supports_fused_op(op: FusedOp) -> bool {
366                matches!(
367                    op,
368                    FusedOp::Add
369                        | FusedOp::Multiply
370                        | FusedOp::Negate
371                        | FusedOp::Conj
372                        | FusedOp::Abs
373                        | FusedOp::Maximum
374                        | FusedOp::Minimum
375                        | FusedOp::Clamp
376                )
377            }
378
379            #[inline(always)]
380            fn fused_add(self, rhs: Self) -> Self {
381                self.wrapping_add(rhs)
382            }
383
384            #[inline(always)]
385            fn fused_multiply(self, rhs: Self) -> Self {
386                self.wrapping_mul(rhs)
387            }
388
389            #[inline(always)]
390            fn fused_negate(self) -> Self {
391                self.wrapping_neg()
392            }
393
394            #[inline(always)]
395            fn fused_conj(self) -> Self {
396                self
397            }
398
399            #[inline(always)]
400            fn fused_divide(self, _rhs: Self) -> Self {
401                unsupported_fused_op!("divide", $label)
402            }
403
404            #[inline(always)]
405            fn fused_abs(self) -> Self {
406                self.wrapping_abs()
407            }
408
409            #[inline(always)]
410            fn fused_maximum(self, rhs: Self) -> Self {
411                self.max(rhs)
412            }
413
414            #[inline(always)]
415            fn fused_minimum(self, rhs: Self) -> Self {
416                self.min(rhs)
417            }
418
419            #[inline(always)]
420            fn fused_clamp(self, min: Self, max: Self) -> Self {
421                self.fused_maximum(min).fused_minimum(max)
422            }
423
424            #[inline(always)]
425            fn fused_exp(self) -> Self {
426                unsupported_fused_op!("exp", $label)
427            }
428
429            #[inline(always)]
430            fn fused_log(self) -> Self {
431                unsupported_fused_op!("log", $label)
432            }
433
434            #[inline(always)]
435            fn fused_sin(self) -> Self {
436                unsupported_fused_op!("sin", $label)
437            }
438
439            #[inline(always)]
440            fn fused_cos(self) -> Self {
441                unsupported_fused_op!("cos", $label)
442            }
443
444            #[inline(always)]
445            fn fused_tanh(self) -> Self {
446                unsupported_fused_op!("tanh", $label)
447            }
448
449            #[inline(always)]
450            fn fused_sqrt(self) -> Self {
451                unsupported_fused_op!("sqrt", $label)
452            }
453
454            #[inline(always)]
455            fn fused_rsqrt(self) -> Self {
456                unsupported_fused_op!("rsqrt", $label)
457            }
458
459            #[inline(always)]
460            fn fused_pow(self, _rhs: Self) -> Self {
461                unsupported_fused_op!("pow", $label)
462            }
463
464            #[inline(always)]
465            fn fused_expm1(self) -> Self {
466                unsupported_fused_op!("expm1", $label)
467            }
468
469            #[inline(always)]
470            fn fused_log1p(self) -> Self {
471                unsupported_fused_op!("log1p", $label)
472            }
473        }
474    };
475}
476
477impl_signed_integer_fused_scalar!(i32, "i32");
478impl_signed_integer_fused_scalar!(i64, "i64");
479
480impl FusedScalar for bool {
481    #[inline]
482    fn fused_dtype_label() -> &'static str {
483        "bool"
484    }
485
486    #[inline]
487    fn supports_fused_op(op: FusedOp) -> bool {
488        matches!(op, FusedOp::Conj)
489    }
490
491    #[inline(always)]
492    fn fused_add(self, _rhs: Self) -> Self {
493        unsupported_fused_op!("add", "bool")
494    }
495
496    #[inline(always)]
497    fn fused_multiply(self, _rhs: Self) -> Self {
498        unsupported_fused_op!("multiply", "bool")
499    }
500
501    #[inline(always)]
502    fn fused_negate(self) -> Self {
503        unsupported_fused_op!("negate", "bool")
504    }
505
506    #[inline(always)]
507    fn fused_conj(self) -> Self {
508        self
509    }
510
511    #[inline(always)]
512    fn fused_divide(self, _rhs: Self) -> Self {
513        unsupported_fused_op!("divide", "bool")
514    }
515
516    #[inline(always)]
517    fn fused_abs(self) -> Self {
518        unsupported_fused_op!("abs", "bool")
519    }
520
521    #[inline(always)]
522    fn fused_maximum(self, _rhs: Self) -> Self {
523        unsupported_fused_op!("maximum", "bool")
524    }
525
526    #[inline(always)]
527    fn fused_minimum(self, _rhs: Self) -> Self {
528        unsupported_fused_op!("minimum", "bool")
529    }
530
531    #[inline(always)]
532    fn fused_clamp(self, _min: Self, _max: Self) -> Self {
533        unsupported_fused_op!("clamp", "bool")
534    }
535
536    #[inline(always)]
537    fn fused_exp(self) -> Self {
538        unsupported_fused_op!("exp", "bool")
539    }
540
541    #[inline(always)]
542    fn fused_log(self) -> Self {
543        unsupported_fused_op!("log", "bool")
544    }
545
546    #[inline(always)]
547    fn fused_sin(self) -> Self {
548        unsupported_fused_op!("sin", "bool")
549    }
550
551    #[inline(always)]
552    fn fused_cos(self) -> Self {
553        unsupported_fused_op!("cos", "bool")
554    }
555
556    #[inline(always)]
557    fn fused_tanh(self) -> Self {
558        unsupported_fused_op!("tanh", "bool")
559    }
560
561    #[inline(always)]
562    fn fused_sqrt(self) -> Self {
563        unsupported_fused_op!("sqrt", "bool")
564    }
565
566    #[inline(always)]
567    fn fused_rsqrt(self) -> Self {
568        unsupported_fused_op!("rsqrt", "bool")
569    }
570
571    #[inline(always)]
572    fn fused_pow(self, _rhs: Self) -> Self {
573        unsupported_fused_op!("pow", "bool")
574    }
575
576    #[inline(always)]
577    fn fused_expm1(self) -> Self {
578        unsupported_fused_op!("expm1", "bool")
579    }
580
581    #[inline(always)]
582    fn fused_log1p(self) -> Self {
583        unsupported_fused_op!("log1p", "bool")
584    }
585}
586
587#[inline]
588fn op_arity(op: FusedOp) -> usize {
589    match op {
590        FusedOp::Negate
591        | FusedOp::Conj
592        | FusedOp::Abs
593        | FusedOp::Exp
594        | FusedOp::Log
595        | FusedOp::Sin
596        | FusedOp::Cos
597        | FusedOp::Tanh
598        | FusedOp::Sqrt
599        | FusedOp::Rsqrt
600        | FusedOp::Expm1
601        | FusedOp::Log1p => 1,
602        FusedOp::Add
603        | FusedOp::Multiply
604        | FusedOp::Divide
605        | FusedOp::Maximum
606        | FusedOp::Minimum
607        | FusedOp::Pow => 2,
608        FusedOp::Clamp => 3,
609    }
610}
611
612pub(crate) fn validate_plan(
613    plan: &FusedPlan,
614    input_count: usize,
615    output_count: usize,
616) -> Result<()> {
617    if input_count != plan.input_count {
618        return Err(StridedError::RankMismatch(input_count, plan.input_count));
619    }
620    if output_count != plan.outputs.len() {
621        return Err(StridedError::RankMismatch(output_count, plan.outputs.len()));
622    }
623    if output_count == 0 {
624        return Err(StridedError::RankMismatch(0, 1));
625    }
626
627    let mut value_count = plan.input_count;
628    for inst in &plan.ops {
629        let expected_arity = op_arity(inst.op);
630        if inst.inputs.len() != expected_arity {
631            return Err(StridedError::RankMismatch(
632                inst.inputs.len(),
633                expected_arity,
634            ));
635        }
636        for &input in &inst.inputs {
637            if input >= value_count {
638                return Err(StridedError::InvalidAxis {
639                    axis: input,
640                    rank: value_count,
641                });
642            }
643        }
644        value_count += 1;
645    }
646
647    for &output in &plan.outputs {
648        if output >= value_count {
649            return Err(StridedError::InvalidAxis {
650                axis: output,
651                rank: value_count,
652            });
653        }
654    }
655
656    Ok(())
657}
658
659pub(crate) fn validate_plan_for_scalar<T: FusedScalar>(
660    plan: &FusedPlan,
661    input_count: usize,
662    output_count: usize,
663) -> Result<()> {
664    validate_plan(plan, input_count, output_count)?;
665    for inst in &plan.ops {
666        if !T::supports_fused_op(inst.op) {
667            return Err(StridedError::UnsupportedOp {
668                op: inst.op.label(),
669                dtype: T::fused_dtype_label(),
670            });
671        }
672    }
673    Ok(())
674}
675
676fn validate_shapes<T: FusedScalar>(
677    dests: &[StridedViewMut<'_, T>],
678    inputs: &[StridedView<'_, T>],
679) -> Result<()> {
680    let dims = dests[0].dims();
681    for dest in dests {
682        validate_destination_layout(dest)?;
683    }
684    for dest in &dests[1..] {
685        ensure_same_shape(dims, dest.dims())?;
686    }
687    for input in inputs {
688        ensure_same_shape(dims, input.dims())?;
689    }
690    Ok(())
691}
692
693fn validate_destination_layout<T>(dest: &StridedViewMut<'_, T>) -> Result<()> {
694    if is_injective_layout(dest.dims(), dest.strides()) {
695        Ok(())
696    } else {
697        Err(StridedError::NonInjectiveOutputLayout)
698    }
699}
700
701#[inline(always)]
702fn eval_op<T: FusedScalar>(op: FusedOp, regs: &[T], inputs: &[usize]) -> T {
703    match op {
704        FusedOp::Negate
705        | FusedOp::Conj
706        | FusedOp::Abs
707        | FusedOp::Exp
708        | FusedOp::Log
709        | FusedOp::Sin
710        | FusedOp::Cos
711        | FusedOp::Tanh
712        | FusedOp::Sqrt
713        | FusedOp::Rsqrt
714        | FusedOp::Expm1
715        | FusedOp::Log1p => eval_unary(op, regs[inputs[0]]),
716        FusedOp::Add
717        | FusedOp::Multiply
718        | FusedOp::Divide
719        | FusedOp::Maximum
720        | FusedOp::Minimum
721        | FusedOp::Pow => eval_binary(op, regs[inputs[0]], regs[inputs[1]]),
722        FusedOp::Clamp => eval_ternary(op, regs[inputs[0]], regs[inputs[1]], regs[inputs[2]]),
723    }
724}
725
726#[inline(always)]
727fn eval_unary<T: FusedScalar>(op: FusedOp, x: T) -> T {
728    match op {
729        FusedOp::Negate => x.fused_negate(),
730        FusedOp::Conj => x.fused_conj(),
731        FusedOp::Abs => x.fused_abs(),
732        FusedOp::Exp => x.fused_exp(),
733        FusedOp::Log => x.fused_log(),
734        FusedOp::Sin => x.fused_sin(),
735        FusedOp::Cos => x.fused_cos(),
736        FusedOp::Tanh => x.fused_tanh(),
737        FusedOp::Sqrt => x.fused_sqrt(),
738        FusedOp::Rsqrt => x.fused_rsqrt(),
739        FusedOp::Expm1 => x.fused_expm1(),
740        FusedOp::Log1p => x.fused_log1p(),
741        _ => unreachable!("not a unary fused op: {op:?}"),
742    }
743}
744
745#[inline(always)]
746fn eval_binary<T: FusedScalar>(op: FusedOp, a: T, b: T) -> T {
747    match op {
748        FusedOp::Add => a.fused_add(b),
749        FusedOp::Multiply => a.fused_multiply(b),
750        FusedOp::Divide => a.fused_divide(b),
751        FusedOp::Maximum => a.fused_maximum(b),
752        FusedOp::Minimum => a.fused_minimum(b),
753        FusedOp::Pow => a.fused_pow(b),
754        _ => unreachable!("not a binary fused op: {op:?}"),
755    }
756}
757
758#[inline(always)]
759fn eval_ternary<T: FusedScalar>(op: FusedOp, a: T, b: T, c: T) -> T {
760    match op {
761        FusedOp::Clamp => a.fused_clamp(b, c),
762        _ => unreachable!("not a ternary fused op: {op:?}"),
763    }
764}
765
766#[derive(Clone, Copy)]
767enum StaticFusedKind {
768    Unary(FusedOp, usize),
769    Binary(FusedOp, usize, usize),
770    Ternary(FusedOp, usize, usize, usize),
771    AddMulLeft,
772    AddMulRight,
773    MulAddExp,
774    DivClampSqrtRsqrt,
775}
776
777#[cfg(test)]
778std::thread_local! {
779    static UNINITIALIZED_STATIC_FAMILY_HITS: core::cell::Cell<[usize; 7]> =
780        const { core::cell::Cell::new([0; 7]) };
781}
782
783#[cfg(test)]
784impl StaticFusedKind {
785    fn test_index(self) -> usize {
786        match self {
787            Self::Unary(..) => 0,
788            Self::Binary(..) => 1,
789            Self::Ternary(..) => 2,
790            Self::AddMulLeft => 3,
791            Self::AddMulRight => 4,
792            Self::MulAddExp => 5,
793            Self::DivClampSqrtRsqrt => 6,
794        }
795    }
796}
797
798#[cfg(all(test, feature = "parallel"))]
799fn reset_uninitialized_static_family_hits() {
800    UNINITIALIZED_STATIC_FAMILY_HITS.set([0; 7]);
801}
802
803#[cfg(all(test, feature = "parallel"))]
804fn uninitialized_static_family_hits() -> [usize; 7] {
805    UNINITIALIZED_STATIC_FAMILY_HITS.get()
806}
807
808#[cfg(test)]
809fn record_uninitialized_static_family_hit(kind: StaticFusedKind) {
810    UNINITIALIZED_STATIC_FAMILY_HITS.set({
811        let mut hits = UNINITIALIZED_STATIC_FAMILY_HITS.get();
812        hits[kind.test_index()] += 1;
813        hits
814    });
815}
816
817fn classify_static_specialization(plan: &FusedPlan) -> Option<StaticFusedKind> {
818    if plan.outputs.len() != 1 {
819        return None;
820    }
821    if let [inst] = plan.ops.as_slice() {
822        if plan.outputs[0] != plan.input_count {
823            return None;
824        }
825        return match (op_arity(inst.op), inst.inputs.as_slice()) {
826            (1, [a]) => Some(StaticFusedKind::Unary(inst.op, *a)),
827            (2, [a, b]) => Some(StaticFusedKind::Binary(inst.op, *a, *b)),
828            (3, [a, b, c]) => Some(StaticFusedKind::Ternary(inst.op, *a, *b, *c)),
829            _ => None,
830        };
831    }
832    if plan.input_count == 2
833        && plan.outputs.as_slice() == [3]
834        && plan.ops.len() == 2
835        && plan.ops[0].op == FusedOp::Add
836        && plan.ops[0].inputs.as_slice() == [0, 1]
837        && plan.ops[1].op == FusedOp::Multiply
838    {
839        return match plan.ops[1].inputs.as_slice() {
840            [2, 0] => Some(StaticFusedKind::AddMulLeft),
841            [0, 2] => Some(StaticFusedKind::AddMulRight),
842            _ => None,
843        };
844    }
845    if plan.input_count == 3
846        && plan.outputs.as_slice() == [5]
847        && plan.ops.len() == 3
848        && plan.ops[0].op == FusedOp::Multiply
849        && plan.ops[0].inputs.as_slice() == [0, 1]
850        && plan.ops[1].op == FusedOp::Add
851        && plan.ops[1].inputs.as_slice() == [3, 2]
852        && plan.ops[2].op == FusedOp::Exp
853        && plan.ops[2].inputs.as_slice() == [4]
854    {
855        return Some(StaticFusedKind::MulAddExp);
856    }
857    if plan.input_count == 4
858        && plan.outputs.as_slice() == [8]
859        && plan.ops.len() == 5
860        && plan.ops[0].op == FusedOp::Divide
861        && plan.ops[0].inputs.as_slice() == [0, 1]
862        && plan.ops[1].op == FusedOp::Maximum
863        && plan.ops[1].inputs.as_slice() == [4, 2]
864        && plan.ops[2].op == FusedOp::Minimum
865        && plan.ops[2].inputs.as_slice() == [5, 3]
866        && plan.ops[3].op == FusedOp::Sqrt
867        && plan.ops[3].inputs.as_slice() == [6]
868        && plan.ops[4].op == FusedOp::Rsqrt
869        && plan.ops[4].inputs.as_slice() == [7]
870    {
871        return Some(StaticFusedKind::DivClampSqrtRsqrt);
872    }
873    None
874}
875
876trait StaticOutput<T: FusedScalar> {
877    type Value: Copy + MaybeSendSync;
878
879    #[cfg(test)]
880    const IS_UNINITIALIZED: bool;
881
882    fn write(value: T) -> Self::Value;
883}
884
885struct InitializedStaticOutput;
886
887impl<T: FusedScalar> StaticOutput<T> for InitializedStaticOutput {
888    type Value = T;
889
890    #[cfg(test)]
891    const IS_UNINITIALIZED: bool = false;
892
893    #[inline(always)]
894    fn write(value: T) -> T {
895        value
896    }
897}
898
899struct UninitializedStaticOutput;
900
901impl<T: FusedScalar> StaticOutput<T> for UninitializedStaticOutput {
902    type Value = MaybeUninit<T>;
903
904    #[cfg(test)]
905    const IS_UNINITIALIZED: bool = true;
906
907    #[inline(always)]
908    fn write(value: T) -> MaybeUninit<T> {
909        MaybeUninit::new(value)
910    }
911}
912
913fn try_static_specialization_validated<T, O>(
914    dest: &mut StridedViewMut<'_, O::Value>,
915    inputs: &[StridedView<'_, T>],
916    plan: &FusedPlan,
917    validated: ValidatedDestinationLayout,
918) -> Result<bool>
919where
920    T: FusedScalar,
921    O: StaticOutput<T>,
922{
923    let Some(kind) = classify_static_specialization(plan) else {
924        return Ok(false);
925    };
926    #[cfg(test)]
927    if O::IS_UNINITIALIZED {
928        record_uninitialized_static_family_hit(kind);
929    }
930
931    match kind {
932        StaticFusedKind::Unary(op, a) => {
933            // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
934            unsafe {
935                map_into_validated(dest, &inputs[a], |x| O::write(eval_unary(op, x)), validated)
936            }?
937        }
938        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
939        StaticFusedKind::Binary(op, a, b) => unsafe {
940            zip_map2_into_validated(
941                dest,
942                &inputs[a],
943                &inputs[b],
944                |x, y| O::write(eval_binary(op, x, y)),
945                validated,
946            )
947        }?,
948        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
949        StaticFusedKind::Ternary(op, a, b, c) => unsafe {
950            zip_map3_into_validated(
951                dest,
952                &inputs[a],
953                &inputs[b],
954                &inputs[c],
955                |x, y, z| O::write(eval_ternary(op, x, y, z)),
956                validated,
957            )
958        }?,
959        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
960        StaticFusedKind::AddMulLeft => unsafe {
961            zip_map2_into_validated(
962                dest,
963                &inputs[0],
964                &inputs[1],
965                |a, b| O::write(a.fused_add(b).fused_multiply(a)),
966                validated,
967            )
968        }?,
969        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
970        StaticFusedKind::AddMulRight => unsafe {
971            zip_map2_into_validated(
972                dest,
973                &inputs[0],
974                &inputs[1],
975                |a, b| O::write(a.fused_multiply(a.fused_add(b))),
976                validated,
977            )
978        }?,
979        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
980        StaticFusedKind::MulAddExp => unsafe {
981            zip_map3_into_validated(
982                dest,
983                &inputs[0],
984                &inputs[1],
985                &inputs[2],
986                |a, b, c| O::write(a.fused_multiply(b).fused_add(c).fused_exp()),
987                validated,
988            )
989        }?,
990        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
991        StaticFusedKind::DivClampSqrtRsqrt => unsafe {
992            zip_map4_into_validated(
993                dest,
994                &inputs[0],
995                &inputs[1],
996                &inputs[2],
997                &inputs[3],
998                |a, b, lo, hi| {
999                    O::write(
1000                        a.fused_divide(b)
1001                            .fused_maximum(lo)
1002                            .fused_minimum(hi)
1003                            .fused_sqrt()
1004                            .fused_rsqrt(),
1005                    )
1006                },
1007                validated,
1008            )
1009        }?,
1010    }
1011    Ok(true)
1012}
1013
1014fn try_static_specialization<T: FusedScalar>(
1015    dests: &mut [StridedViewMut<'_, T>],
1016    inputs: &[StridedView<'_, T>],
1017    plan: &FusedPlan,
1018) -> Result<bool> {
1019    if dests.len() != 1 {
1020        return Ok(false);
1021    }
1022    let validated = validate_destination_layout_without_alloc(dests[0].dims(), dests[0].strides())?;
1023    try_static_specialization_validated::<T, InitializedStaticOutput>(
1024        &mut dests[0],
1025        inputs,
1026        plan,
1027        validated,
1028    )
1029}
1030
1031unsafe fn interpret_inner_loop<T: FusedScalar>(
1032    dst_ptrs: &[*mut T],
1033    input_ptrs: &[*const T],
1034    plan: &FusedPlan,
1035    offsets: &[isize],
1036    len: usize,
1037    strides: &[isize],
1038) {
1039    let output_count = dst_ptrs.len();
1040    let mut regs = Vec::with_capacity(plan.input_count + plan.ops.len());
1041
1042    for i in 0..len {
1043        let i = i as isize;
1044        regs.clear();
1045
1046        for (input_index, &input_ptr) in input_ptrs.iter().enumerate() {
1047            let stride_index = output_count + input_index;
1048            regs.push(*input_ptr.offset(offsets[stride_index] + i * strides[stride_index]));
1049        }
1050
1051        for inst in &plan.ops {
1052            regs.push(eval_op(inst.op, &regs, &inst.inputs));
1053        }
1054
1055        for (output_index, &dst_ptr) in dst_ptrs.iter().enumerate() {
1056            *dst_ptr.offset(offsets[output_index] + i * strides[output_index]) =
1057                regs[plan.outputs[output_index]];
1058        }
1059    }
1060}
1061
1062fn interpret_fused_elementwise_into<T: FusedScalar>(
1063    dests: &mut [StridedViewMut<'_, T>],
1064    inputs: &[StridedView<'_, T>],
1065    plan: &FusedPlan,
1066) -> Result<()> {
1067    #[cfg(feature = "parallel")]
1068    {
1069        let dims = dests[0].dims().to_vec();
1070        if dests[0].len() == 0 {
1071            return Ok(());
1072        }
1073
1074        let dst_ptrs: Vec<*mut T> = dests.iter_mut().map(|dest| dest.as_mut_ptr()).collect();
1075        let input_ptrs: Vec<*const T> = inputs.iter().map(StridedView::ptr).collect();
1076
1077        let mut strides_list: Vec<&[isize]> = Vec::with_capacity(dests.len() + inputs.len());
1078        for dest in dests.iter() {
1079            strides_list.push(dest.strides());
1080        }
1081        for input in inputs {
1082            strides_list.push(input.strides());
1083        }
1084
1085        let elem_size = std::mem::size_of::<T>();
1086        let total = dests[0].len();
1087        let (fused_dims, ordered_strides, kernel_plan) = if total <= SMALL_TENSOR_THRESHOLD {
1088            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1089            unsafe { build_plan_fused_small(&dims, &strides_list) }
1090        } else {
1091            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1092            unsafe { build_plan_fused(&dims, &strides_list, Some(0), elem_size) }
1093        };
1094
1095        let total: usize = fused_dims.iter().product();
1096        let nthreads = strided_basic::execution::rayon_threads();
1097        if total > MINTHREADLENGTH && nthreads > 1 {
1098            let dst_send: Vec<SendPtr<T>> = dst_ptrs.iter().map(|&ptr| SendPtr(ptr)).collect();
1099            let input_send: Vec<SendPtr<T>> = input_ptrs
1100                .iter()
1101                .map(|&ptr| SendPtr(ptr as *mut T))
1102                .collect();
1103
1104            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1105
1106            let costs = unsafe { compute_costs(&ordered_strides) };
1107            let initial_offsets = vec![0isize; ordered_strides.len()];
1108            let run_partition = |dims: &[usize],
1109                                 blocks: &[usize],
1110                                 strides_list: &[Vec<isize>],
1111                                 offsets: &[isize]|
1112             -> Result<()> {
1113                let dst_ptrs: Vec<*mut T> = dst_send.iter().map(|ptr| ptr.as_ptr()).collect();
1114                let input_ptrs: Vec<*const T> =
1115                    input_send.iter().map(|ptr| ptr.as_const()).collect();
1116                let run_block = |offsets: &[isize], len: usize, strides: &[isize]| -> Result<()> {
1117                    // SAFETY: validated shapes/layouts and the derived block bound every access.
1118                    unsafe {
1119                        interpret_inner_loop(&dst_ptrs, &input_ptrs, plan, offsets, len, strides);
1120                    }
1121                    Ok(())
1122                };
1123                // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1124                unsafe {
1125                    for_each_inner_block_preordered(dims, blocks, strides_list, offsets, run_block)
1126                }
1127            };
1128            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1129            return unsafe {
1130                mapreduce_threaded(
1131                    &fused_dims,
1132                    &kernel_plan.block,
1133                    &ordered_strides,
1134                    &initial_offsets,
1135                    &costs,
1136                    nthreads,
1137                    0,
1138                    1,
1139                    &run_partition,
1140                )
1141            };
1142        }
1143    }
1144
1145    interpret_fused_elementwise_into_serial(dests, inputs, plan)
1146}
1147
1148fn interpret_fused_elementwise_into_serial<T: FusedScalar>(
1149    dests: &mut [StridedViewMut<'_, T>],
1150    inputs: &[StridedView<'_, T>],
1151    plan: &FusedPlan,
1152) -> Result<()> {
1153    let dims = dests[0].dims().to_vec();
1154    if dests[0].len() == 0 {
1155        return Ok(());
1156    }
1157
1158    let dst_ptrs: Vec<*mut T> = dests.iter_mut().map(|dest| dest.as_mut_ptr()).collect();
1159    let input_ptrs: Vec<*const T> = inputs.iter().map(StridedView::ptr).collect();
1160
1161    let mut strides_list: Vec<&[isize]> = Vec::with_capacity(dests.len() + inputs.len());
1162    for dest in dests.iter() {
1163        strides_list.push(dest.strides());
1164    }
1165    for input in inputs {
1166        strides_list.push(input.strides());
1167    }
1168
1169    let elem_size = std::mem::size_of::<T>();
1170    let total = dests[0].len();
1171    let (fused_dims, ordered_strides, kernel_plan) = if total <= SMALL_TENSOR_THRESHOLD {
1172        // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1173        unsafe { build_plan_fused_small(&dims, &strides_list) }
1174    } else {
1175        // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1176        unsafe { build_plan_fused(&dims, &strides_list, Some(0), elem_size) }
1177    };
1178
1179    let initial_offsets = vec![0isize; ordered_strides.len()];
1180    let run_block = |offsets: &[isize], len: usize, strides: &[isize]| -> Result<()> {
1181        // SAFETY: validated shapes/layouts and the derived block bound every access.
1182        unsafe {
1183            interpret_inner_loop(&dst_ptrs, &input_ptrs, plan, offsets, len, strides);
1184        }
1185        Ok(())
1186    };
1187    // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1188    unsafe {
1189        for_each_inner_block_preordered(
1190            &fused_dims,
1191            &kernel_plan.block,
1192            &ordered_strides,
1193            &initial_offsets,
1194            run_block,
1195        )
1196    }
1197}
1198
1199pub(crate) fn fused_elementwise_into_serial<T: FusedScalar>(
1200    dests: &mut [StridedViewMut<'_, T>],
1201    inputs: &[StridedView<'_, T>],
1202    plan: &FusedPlan,
1203) -> Result<()> {
1204    validate_plan_for_scalar::<T>(plan, inputs.len(), dests.len())?;
1205    validate_shapes(dests, inputs)?;
1206    interpret_fused_elementwise_into_serial(dests, inputs, plan)
1207}
1208
1209unsafe fn interpret_inner_loop_uninit<T: FusedScalar>(
1210    dst_ptr: *mut MaybeUninit<T>,
1211    input_ptrs: &[*const T],
1212    plan: &FusedPlan,
1213    offsets: &[isize],
1214    len: usize,
1215    strides: &[isize],
1216) {
1217    let mut regs = Vec::with_capacity(plan.input_count + plan.ops.len());
1218    for i in 0..len {
1219        let i = i as isize;
1220        regs.clear();
1221        for (input_index, &input_ptr) in input_ptrs.iter().enumerate() {
1222            let stride_index = 1 + input_index;
1223            regs.push(*input_ptr.offset(offsets[stride_index] + i * strides[stride_index]));
1224        }
1225        for inst in &plan.ops {
1226            regs.push(eval_op(inst.op, &regs, &inst.inputs));
1227        }
1228        *dst_ptr.offset(offsets[0] + i * strides[0]) = MaybeUninit::new(regs[plan.outputs[0]]);
1229    }
1230}
1231
1232pub(crate) fn fused_elementwise_into_uninit<T: FusedScalar>(
1233    dest: &mut StridedViewMut<'_, MaybeUninit<T>>,
1234    inputs: &[StridedView<'_, T>],
1235    plan: &FusedPlan,
1236    serial: bool,
1237    validated: ValidatedDestinationLayout,
1238) -> Result<()> {
1239    #[cfg(not(feature = "parallel"))]
1240    let _ = serial;
1241    validate_plan_for_scalar::<T>(plan, inputs.len(), 1)?;
1242    for input in inputs {
1243        ensure_same_shape(dest.dims(), input.dims())?;
1244    }
1245
1246    if !serial
1247        && try_static_specialization_validated::<T, UninitializedStaticOutput>(
1248            dest, inputs, plan, validated,
1249        )?
1250    {
1251        return Ok(());
1252    }
1253
1254    let dims = dest.dims();
1255    if dest.len() == 0 {
1256        return Ok(());
1257    }
1258    let dst_ptr = dest.as_mut_ptr();
1259    let input_ptrs: Vec<*const T> = inputs.iter().map(StridedView::ptr).collect();
1260    let mut strides_list: Vec<&[isize]> = Vec::with_capacity(1 + inputs.len());
1261    strides_list.push(dest.strides());
1262    for input in inputs {
1263        strides_list.push(input.strides());
1264    }
1265    let total = dest.len();
1266    let (fused_dims, ordered_strides, kernel_plan) = if total <= SMALL_TENSOR_THRESHOLD {
1267        // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1268        unsafe { build_plan_fused_small(dims, &strides_list) }
1269    } else {
1270        // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1271        unsafe { build_plan_fused(dims, &strides_list, Some(0), core::mem::size_of::<T>()) }
1272    };
1273
1274    #[cfg(feature = "parallel")]
1275    {
1276        let total: usize = fused_dims.iter().product();
1277        let nthreads = strided_basic::execution::rayon_threads();
1278        if !serial && total > MINTHREADLENGTH && nthreads > 1 {
1279            let dst_send = SendPtr(dst_ptr);
1280            let input_send: Vec<SendPtr<T>> = input_ptrs
1281                .iter()
1282                .map(|&ptr| SendPtr(ptr as *mut T))
1283                .collect();
1284            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1285            let costs = unsafe { compute_costs(&ordered_strides) };
1286            let initial_offsets = vec![0isize; ordered_strides.len()];
1287            let run_partition = |dims: &[usize],
1288                                 blocks: &[usize],
1289                                 strides_list: &[Vec<isize>],
1290                                 offsets: &[isize]|
1291             -> Result<()> {
1292                let input_ptrs: Vec<*const T> =
1293                    input_send.iter().map(|ptr| ptr.as_const()).collect();
1294                let run_block = |offsets: &[isize], len: usize, strides: &[isize]| -> Result<()> {
1295                    // SAFETY: validated shapes/layouts and the derived block bound every access.
1296                    unsafe {
1297                        interpret_inner_loop_uninit(
1298                            dst_send.as_ptr(),
1299                            &input_ptrs,
1300                            plan,
1301                            offsets,
1302                            len,
1303                            strides,
1304                        );
1305                    }
1306                    Ok(())
1307                };
1308                // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1309                unsafe {
1310                    for_each_inner_block_preordered(dims, blocks, strides_list, offsets, run_block)
1311                }
1312            };
1313            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1314            return unsafe {
1315                mapreduce_threaded(
1316                    &fused_dims,
1317                    &kernel_plan.block,
1318                    &ordered_strides,
1319                    &initial_offsets,
1320                    &costs,
1321                    nthreads,
1322                    0,
1323                    1,
1324                    &run_partition,
1325                )
1326            };
1327        }
1328    }
1329
1330    let initial_offsets = vec![0isize; ordered_strides.len()];
1331    let run_block = |offsets: &[isize], len: usize, strides: &[isize]| -> Result<()> {
1332        // SAFETY: validated shapes/layouts and the derived block bound every access.
1333        unsafe {
1334            interpret_inner_loop_uninit(dst_ptr, &input_ptrs, plan, offsets, len, strides);
1335        }
1336        Ok(())
1337    };
1338    // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1339    unsafe {
1340        for_each_inner_block_preordered(
1341            &fused_dims,
1342            &kernel_plan.block,
1343            &ordered_strides,
1344            &initial_offsets,
1345            run_block,
1346        )
1347    }
1348}
1349
1350/// Evaluate a runtime-DAG elementwise plan into one or more destinations.
1351///
1352/// The plan is validated before any destination is written:
1353///
1354/// - `inputs.len()` must equal `plan.input_count`;
1355/// - `dests.len()` must equal `plan.outputs.len()`;
1356/// - instruction operands must reference earlier SSA values with the right
1357///   arity for their [`FusedOp`];
1358/// - every input and destination must have exactly the destination shape;
1359/// - each mutable destination layout must be injective, so two logical output
1360///   elements never map to the same memory address.
1361///
1362/// The implementation dispatches known single-output plans to existing static
1363/// `map_into`/`zip_map*_into` kernels and uses a generic interpreter fallback
1364/// for arbitrary validated DAGs. Overlapping source/destination memory is not
1365/// supported by the strided kernels generally.
1366///
1367/// Real `Maximum`, `Minimum`, and `Clamp` use Rust `f32`/`f64` `max`/`min`
1368/// semantics. Complex `Abs` returns the norm in the real component; complex
1369/// `Maximum`, `Minimum`, and `Clamp` compare by squared norm. Signed integer
1370/// `Add`, `Multiply`, `Negate`, and `Abs` use wrapping arithmetic. `bool`
1371/// supports only copy-like identity plans and `Conj`; ambiguous arithmetic and
1372/// transcendental op/dtype pairs are rejected before any destination is written.
1373pub fn fused_elementwise_into<T: FusedScalar>(
1374    dests: &mut [StridedViewMut<'_, T>],
1375    inputs: &[StridedView<'_, T>],
1376    plan: &FusedPlan,
1377) -> Result<()> {
1378    validate_plan_for_scalar::<T>(plan, inputs.len(), dests.len())?;
1379    validate_shapes(dests, inputs)?;
1380    if try_static_specialization(dests, inputs, plan)? {
1381        return Ok(());
1382    }
1383    interpret_fused_elementwise_into(dests, inputs, plan)
1384}
1385
1386#[cfg(test)]
1387#[path = "fused/tests/tests.rs"]
1388mod tests;