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) => {
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                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!(num_complex::Complex32);
348impl_complex_fused_scalar!(num_complex::Complex64);
349
350macro_rules! impl_signed_integer_fused_scalar {
351    ($ty:ty, $label:literal) => {
352        impl FusedScalar for $ty {
353            #[inline]
354            fn fused_dtype_label() -> &'static str {
355                $label
356            }
357
358            #[inline]
359            fn supports_fused_op(op: FusedOp) -> bool {
360                matches!(
361                    op,
362                    FusedOp::Add
363                        | FusedOp::Multiply
364                        | FusedOp::Negate
365                        | FusedOp::Conj
366                        | FusedOp::Abs
367                        | FusedOp::Maximum
368                        | FusedOp::Minimum
369                        | FusedOp::Clamp
370                )
371            }
372
373            #[inline(always)]
374            fn fused_add(self, rhs: Self) -> Self {
375                self.wrapping_add(rhs)
376            }
377
378            #[inline(always)]
379            fn fused_multiply(self, rhs: Self) -> Self {
380                self.wrapping_mul(rhs)
381            }
382
383            #[inline(always)]
384            fn fused_negate(self) -> Self {
385                self.wrapping_neg()
386            }
387
388            #[inline(always)]
389            fn fused_conj(self) -> Self {
390                self
391            }
392
393            #[inline(always)]
394            fn fused_divide(self, _rhs: Self) -> Self {
395                unsupported_fused_op!("divide", $label)
396            }
397
398            #[inline(always)]
399            fn fused_abs(self) -> Self {
400                self.wrapping_abs()
401            }
402
403            #[inline(always)]
404            fn fused_maximum(self, rhs: Self) -> Self {
405                self.max(rhs)
406            }
407
408            #[inline(always)]
409            fn fused_minimum(self, rhs: Self) -> Self {
410                self.min(rhs)
411            }
412
413            #[inline(always)]
414            fn fused_clamp(self, min: Self, max: Self) -> Self {
415                self.fused_maximum(min).fused_minimum(max)
416            }
417
418            #[inline(always)]
419            fn fused_exp(self) -> Self {
420                unsupported_fused_op!("exp", $label)
421            }
422
423            #[inline(always)]
424            fn fused_log(self) -> Self {
425                unsupported_fused_op!("log", $label)
426            }
427
428            #[inline(always)]
429            fn fused_sin(self) -> Self {
430                unsupported_fused_op!("sin", $label)
431            }
432
433            #[inline(always)]
434            fn fused_cos(self) -> Self {
435                unsupported_fused_op!("cos", $label)
436            }
437
438            #[inline(always)]
439            fn fused_tanh(self) -> Self {
440                unsupported_fused_op!("tanh", $label)
441            }
442
443            #[inline(always)]
444            fn fused_sqrt(self) -> Self {
445                unsupported_fused_op!("sqrt", $label)
446            }
447
448            #[inline(always)]
449            fn fused_rsqrt(self) -> Self {
450                unsupported_fused_op!("rsqrt", $label)
451            }
452
453            #[inline(always)]
454            fn fused_pow(self, _rhs: Self) -> Self {
455                unsupported_fused_op!("pow", $label)
456            }
457
458            #[inline(always)]
459            fn fused_expm1(self) -> Self {
460                unsupported_fused_op!("expm1", $label)
461            }
462
463            #[inline(always)]
464            fn fused_log1p(self) -> Self {
465                unsupported_fused_op!("log1p", $label)
466            }
467        }
468    };
469}
470
471impl_signed_integer_fused_scalar!(i32, "i32");
472impl_signed_integer_fused_scalar!(i64, "i64");
473
474impl FusedScalar for bool {
475    #[inline]
476    fn fused_dtype_label() -> &'static str {
477        "bool"
478    }
479
480    #[inline]
481    fn supports_fused_op(op: FusedOp) -> bool {
482        matches!(op, FusedOp::Conj)
483    }
484
485    #[inline(always)]
486    fn fused_add(self, _rhs: Self) -> Self {
487        unsupported_fused_op!("add", "bool")
488    }
489
490    #[inline(always)]
491    fn fused_multiply(self, _rhs: Self) -> Self {
492        unsupported_fused_op!("multiply", "bool")
493    }
494
495    #[inline(always)]
496    fn fused_negate(self) -> Self {
497        unsupported_fused_op!("negate", "bool")
498    }
499
500    #[inline(always)]
501    fn fused_conj(self) -> Self {
502        self
503    }
504
505    #[inline(always)]
506    fn fused_divide(self, _rhs: Self) -> Self {
507        unsupported_fused_op!("divide", "bool")
508    }
509
510    #[inline(always)]
511    fn fused_abs(self) -> Self {
512        unsupported_fused_op!("abs", "bool")
513    }
514
515    #[inline(always)]
516    fn fused_maximum(self, _rhs: Self) -> Self {
517        unsupported_fused_op!("maximum", "bool")
518    }
519
520    #[inline(always)]
521    fn fused_minimum(self, _rhs: Self) -> Self {
522        unsupported_fused_op!("minimum", "bool")
523    }
524
525    #[inline(always)]
526    fn fused_clamp(self, _min: Self, _max: Self) -> Self {
527        unsupported_fused_op!("clamp", "bool")
528    }
529
530    #[inline(always)]
531    fn fused_exp(self) -> Self {
532        unsupported_fused_op!("exp", "bool")
533    }
534
535    #[inline(always)]
536    fn fused_log(self) -> Self {
537        unsupported_fused_op!("log", "bool")
538    }
539
540    #[inline(always)]
541    fn fused_sin(self) -> Self {
542        unsupported_fused_op!("sin", "bool")
543    }
544
545    #[inline(always)]
546    fn fused_cos(self) -> Self {
547        unsupported_fused_op!("cos", "bool")
548    }
549
550    #[inline(always)]
551    fn fused_tanh(self) -> Self {
552        unsupported_fused_op!("tanh", "bool")
553    }
554
555    #[inline(always)]
556    fn fused_sqrt(self) -> Self {
557        unsupported_fused_op!("sqrt", "bool")
558    }
559
560    #[inline(always)]
561    fn fused_rsqrt(self) -> Self {
562        unsupported_fused_op!("rsqrt", "bool")
563    }
564
565    #[inline(always)]
566    fn fused_pow(self, _rhs: Self) -> Self {
567        unsupported_fused_op!("pow", "bool")
568    }
569
570    #[inline(always)]
571    fn fused_expm1(self) -> Self {
572        unsupported_fused_op!("expm1", "bool")
573    }
574
575    #[inline(always)]
576    fn fused_log1p(self) -> Self {
577        unsupported_fused_op!("log1p", "bool")
578    }
579}
580
581#[inline]
582fn op_arity(op: FusedOp) -> usize {
583    match op {
584        FusedOp::Negate
585        | FusedOp::Conj
586        | FusedOp::Abs
587        | FusedOp::Exp
588        | FusedOp::Log
589        | FusedOp::Sin
590        | FusedOp::Cos
591        | FusedOp::Tanh
592        | FusedOp::Sqrt
593        | FusedOp::Rsqrt
594        | FusedOp::Expm1
595        | FusedOp::Log1p => 1,
596        FusedOp::Add
597        | FusedOp::Multiply
598        | FusedOp::Divide
599        | FusedOp::Maximum
600        | FusedOp::Minimum
601        | FusedOp::Pow => 2,
602        FusedOp::Clamp => 3,
603    }
604}
605
606pub(crate) fn validate_plan(
607    plan: &FusedPlan,
608    input_count: usize,
609    output_count: usize,
610) -> Result<()> {
611    if input_count != plan.input_count {
612        return Err(StridedError::RankMismatch(input_count, plan.input_count));
613    }
614    if output_count != plan.outputs.len() {
615        return Err(StridedError::RankMismatch(output_count, plan.outputs.len()));
616    }
617    if output_count == 0 {
618        return Err(StridedError::RankMismatch(0, 1));
619    }
620
621    let mut value_count = plan.input_count;
622    for inst in &plan.ops {
623        let expected_arity = op_arity(inst.op);
624        if inst.inputs.len() != expected_arity {
625            return Err(StridedError::RankMismatch(
626                inst.inputs.len(),
627                expected_arity,
628            ));
629        }
630        for &input in &inst.inputs {
631            if input >= value_count {
632                return Err(StridedError::InvalidAxis {
633                    axis: input,
634                    rank: value_count,
635                });
636            }
637        }
638        value_count += 1;
639    }
640
641    for &output in &plan.outputs {
642        if output >= value_count {
643            return Err(StridedError::InvalidAxis {
644                axis: output,
645                rank: value_count,
646            });
647        }
648    }
649
650    Ok(())
651}
652
653pub(crate) fn validate_plan_for_scalar<T: FusedScalar>(
654    plan: &FusedPlan,
655    input_count: usize,
656    output_count: usize,
657) -> Result<()> {
658    validate_plan(plan, input_count, output_count)?;
659    for inst in &plan.ops {
660        if !T::supports_fused_op(inst.op) {
661            return Err(StridedError::UnsupportedOp {
662                op: inst.op.label(),
663                dtype: T::fused_dtype_label(),
664            });
665        }
666    }
667    Ok(())
668}
669
670fn validate_shapes<T: FusedScalar>(
671    dests: &[StridedViewMut<'_, T>],
672    inputs: &[StridedView<'_, T>],
673) -> Result<()> {
674    let dims = dests[0].dims();
675    for dest in dests {
676        validate_destination_layout(dest)?;
677    }
678    for dest in &dests[1..] {
679        ensure_same_shape(dims, dest.dims())?;
680    }
681    for input in inputs {
682        ensure_same_shape(dims, input.dims())?;
683    }
684    Ok(())
685}
686
687fn validate_destination_layout<T>(dest: &StridedViewMut<'_, T>) -> Result<()> {
688    if is_injective_layout(dest.dims(), dest.strides()) {
689        Ok(())
690    } else {
691        Err(StridedError::NonInjectiveOutputLayout)
692    }
693}
694
695#[inline(always)]
696fn eval_op<T: FusedScalar>(op: FusedOp, regs: &[T], inputs: &[usize]) -> T {
697    match op {
698        FusedOp::Negate
699        | FusedOp::Conj
700        | FusedOp::Abs
701        | FusedOp::Exp
702        | FusedOp::Log
703        | FusedOp::Sin
704        | FusedOp::Cos
705        | FusedOp::Tanh
706        | FusedOp::Sqrt
707        | FusedOp::Rsqrt
708        | FusedOp::Expm1
709        | FusedOp::Log1p => eval_unary(op, regs[inputs[0]]),
710        FusedOp::Add
711        | FusedOp::Multiply
712        | FusedOp::Divide
713        | FusedOp::Maximum
714        | FusedOp::Minimum
715        | FusedOp::Pow => eval_binary(op, regs[inputs[0]], regs[inputs[1]]),
716        FusedOp::Clamp => eval_ternary(op, regs[inputs[0]], regs[inputs[1]], regs[inputs[2]]),
717    }
718}
719
720#[inline(always)]
721fn eval_unary<T: FusedScalar>(op: FusedOp, x: T) -> T {
722    match op {
723        FusedOp::Negate => x.fused_negate(),
724        FusedOp::Conj => x.fused_conj(),
725        FusedOp::Abs => x.fused_abs(),
726        FusedOp::Exp => x.fused_exp(),
727        FusedOp::Log => x.fused_log(),
728        FusedOp::Sin => x.fused_sin(),
729        FusedOp::Cos => x.fused_cos(),
730        FusedOp::Tanh => x.fused_tanh(),
731        FusedOp::Sqrt => x.fused_sqrt(),
732        FusedOp::Rsqrt => x.fused_rsqrt(),
733        FusedOp::Expm1 => x.fused_expm1(),
734        FusedOp::Log1p => x.fused_log1p(),
735        _ => unreachable!("not a unary fused op: {op:?}"),
736    }
737}
738
739#[inline(always)]
740fn eval_binary<T: FusedScalar>(op: FusedOp, a: T, b: T) -> T {
741    match op {
742        FusedOp::Add => a.fused_add(b),
743        FusedOp::Multiply => a.fused_multiply(b),
744        FusedOp::Divide => a.fused_divide(b),
745        FusedOp::Maximum => a.fused_maximum(b),
746        FusedOp::Minimum => a.fused_minimum(b),
747        FusedOp::Pow => a.fused_pow(b),
748        _ => unreachable!("not a binary fused op: {op:?}"),
749    }
750}
751
752#[inline(always)]
753fn eval_ternary<T: FusedScalar>(op: FusedOp, a: T, b: T, c: T) -> T {
754    match op {
755        FusedOp::Clamp => a.fused_clamp(b, c),
756        _ => unreachable!("not a ternary fused op: {op:?}"),
757    }
758}
759
760#[derive(Clone, Copy)]
761enum StaticFusedKind {
762    Unary(FusedOp, usize),
763    Binary(FusedOp, usize, usize),
764    Ternary(FusedOp, usize, usize, usize),
765    AddMulLeft,
766    AddMulRight,
767    MulAddExp,
768    DivClampSqrtRsqrt,
769}
770
771#[cfg(test)]
772std::thread_local! {
773    static UNINITIALIZED_STATIC_FAMILY_HITS: core::cell::Cell<[usize; 7]> =
774        const { core::cell::Cell::new([0; 7]) };
775}
776
777#[cfg(test)]
778impl StaticFusedKind {
779    fn test_index(self) -> usize {
780        match self {
781            Self::Unary(..) => 0,
782            Self::Binary(..) => 1,
783            Self::Ternary(..) => 2,
784            Self::AddMulLeft => 3,
785            Self::AddMulRight => 4,
786            Self::MulAddExp => 5,
787            Self::DivClampSqrtRsqrt => 6,
788        }
789    }
790}
791
792#[cfg(all(test, feature = "parallel"))]
793fn reset_uninitialized_static_family_hits() {
794    UNINITIALIZED_STATIC_FAMILY_HITS.set([0; 7]);
795}
796
797#[cfg(all(test, feature = "parallel"))]
798fn uninitialized_static_family_hits() -> [usize; 7] {
799    UNINITIALIZED_STATIC_FAMILY_HITS.get()
800}
801
802#[cfg(test)]
803fn record_uninitialized_static_family_hit(kind: StaticFusedKind) {
804    UNINITIALIZED_STATIC_FAMILY_HITS.set({
805        let mut hits = UNINITIALIZED_STATIC_FAMILY_HITS.get();
806        hits[kind.test_index()] += 1;
807        hits
808    });
809}
810
811fn classify_static_specialization(plan: &FusedPlan) -> Option<StaticFusedKind> {
812    if plan.outputs.len() != 1 {
813        return None;
814    }
815    if let [inst] = plan.ops.as_slice() {
816        if plan.outputs[0] != plan.input_count {
817            return None;
818        }
819        return match (op_arity(inst.op), inst.inputs.as_slice()) {
820            (1, [a]) => Some(StaticFusedKind::Unary(inst.op, *a)),
821            (2, [a, b]) => Some(StaticFusedKind::Binary(inst.op, *a, *b)),
822            (3, [a, b, c]) => Some(StaticFusedKind::Ternary(inst.op, *a, *b, *c)),
823            _ => None,
824        };
825    }
826    if plan.input_count == 2
827        && plan.outputs.as_slice() == [3]
828        && plan.ops.len() == 2
829        && plan.ops[0].op == FusedOp::Add
830        && plan.ops[0].inputs.as_slice() == [0, 1]
831        && plan.ops[1].op == FusedOp::Multiply
832    {
833        return match plan.ops[1].inputs.as_slice() {
834            [2, 0] => Some(StaticFusedKind::AddMulLeft),
835            [0, 2] => Some(StaticFusedKind::AddMulRight),
836            _ => None,
837        };
838    }
839    if plan.input_count == 3
840        && plan.outputs.as_slice() == [5]
841        && plan.ops.len() == 3
842        && plan.ops[0].op == FusedOp::Multiply
843        && plan.ops[0].inputs.as_slice() == [0, 1]
844        && plan.ops[1].op == FusedOp::Add
845        && plan.ops[1].inputs.as_slice() == [3, 2]
846        && plan.ops[2].op == FusedOp::Exp
847        && plan.ops[2].inputs.as_slice() == [4]
848    {
849        return Some(StaticFusedKind::MulAddExp);
850    }
851    if plan.input_count == 4
852        && plan.outputs.as_slice() == [8]
853        && plan.ops.len() == 5
854        && plan.ops[0].op == FusedOp::Divide
855        && plan.ops[0].inputs.as_slice() == [0, 1]
856        && plan.ops[1].op == FusedOp::Maximum
857        && plan.ops[1].inputs.as_slice() == [4, 2]
858        && plan.ops[2].op == FusedOp::Minimum
859        && plan.ops[2].inputs.as_slice() == [5, 3]
860        && plan.ops[3].op == FusedOp::Sqrt
861        && plan.ops[3].inputs.as_slice() == [6]
862        && plan.ops[4].op == FusedOp::Rsqrt
863        && plan.ops[4].inputs.as_slice() == [7]
864    {
865        return Some(StaticFusedKind::DivClampSqrtRsqrt);
866    }
867    None
868}
869
870trait StaticOutput<T: FusedScalar> {
871    type Value: Copy + MaybeSendSync;
872
873    #[cfg(test)]
874    const IS_UNINITIALIZED: bool;
875
876    fn write(value: T) -> Self::Value;
877}
878
879struct InitializedStaticOutput;
880
881impl<T: FusedScalar> StaticOutput<T> for InitializedStaticOutput {
882    type Value = T;
883
884    #[cfg(test)]
885    const IS_UNINITIALIZED: bool = false;
886
887    #[inline(always)]
888    fn write(value: T) -> T {
889        value
890    }
891}
892
893struct UninitializedStaticOutput;
894
895impl<T: FusedScalar> StaticOutput<T> for UninitializedStaticOutput {
896    type Value = MaybeUninit<T>;
897
898    #[cfg(test)]
899    const IS_UNINITIALIZED: bool = true;
900
901    #[inline(always)]
902    fn write(value: T) -> MaybeUninit<T> {
903        MaybeUninit::new(value)
904    }
905}
906
907fn try_static_specialization_validated<T, O>(
908    dest: &mut StridedViewMut<'_, O::Value>,
909    inputs: &[StridedView<'_, T>],
910    plan: &FusedPlan,
911    validated: ValidatedDestinationLayout,
912) -> Result<bool>
913where
914    T: FusedScalar,
915    O: StaticOutput<T>,
916{
917    let Some(kind) = classify_static_specialization(plan) else {
918        return Ok(false);
919    };
920    #[cfg(test)]
921    if O::IS_UNINITIALIZED {
922        record_uninitialized_static_family_hit(kind);
923    }
924
925    match kind {
926        StaticFusedKind::Unary(op, a) => {
927            // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
928            unsafe {
929                map_into_validated(dest, &inputs[a], |x| O::write(eval_unary(op, x)), validated)
930            }?
931        }
932        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
933        StaticFusedKind::Binary(op, a, b) => unsafe {
934            zip_map2_into_validated(
935                dest,
936                &inputs[a],
937                &inputs[b],
938                |x, y| O::write(eval_binary(op, x, y)),
939                validated,
940            )
941        }?,
942        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
943        StaticFusedKind::Ternary(op, a, b, c) => unsafe {
944            zip_map3_into_validated(
945                dest,
946                &inputs[a],
947                &inputs[b],
948                &inputs[c],
949                |x, y, z| O::write(eval_ternary(op, x, y, z)),
950                validated,
951            )
952        }?,
953        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
954        StaticFusedKind::AddMulLeft => unsafe {
955            zip_map2_into_validated(
956                dest,
957                &inputs[0],
958                &inputs[1],
959                |a, b| O::write(a.fused_add(b).fused_multiply(a)),
960                validated,
961            )
962        }?,
963        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
964        StaticFusedKind::AddMulRight => unsafe {
965            zip_map2_into_validated(
966                dest,
967                &inputs[0],
968                &inputs[1],
969                |a, b| O::write(a.fused_multiply(a.fused_add(b))),
970                validated,
971            )
972        }?,
973        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
974        StaticFusedKind::MulAddExp => unsafe {
975            zip_map3_into_validated(
976                dest,
977                &inputs[0],
978                &inputs[1],
979                &inputs[2],
980                |a, b, c| O::write(a.fused_multiply(b).fused_add(c).fused_exp()),
981                validated,
982            )
983        }?,
984        // SAFETY: matching shapes and this destination layout were validated before specialization/replay.
985        StaticFusedKind::DivClampSqrtRsqrt => unsafe {
986            zip_map4_into_validated(
987                dest,
988                &inputs[0],
989                &inputs[1],
990                &inputs[2],
991                &inputs[3],
992                |a, b, lo, hi| {
993                    O::write(
994                        a.fused_divide(b)
995                            .fused_maximum(lo)
996                            .fused_minimum(hi)
997                            .fused_sqrt()
998                            .fused_rsqrt(),
999                    )
1000                },
1001                validated,
1002            )
1003        }?,
1004    }
1005    Ok(true)
1006}
1007
1008fn try_static_specialization<T: FusedScalar>(
1009    dests: &mut [StridedViewMut<'_, T>],
1010    inputs: &[StridedView<'_, T>],
1011    plan: &FusedPlan,
1012) -> Result<bool> {
1013    if dests.len() != 1 {
1014        return Ok(false);
1015    }
1016    let validated = validate_destination_layout_without_alloc(dests[0].dims(), dests[0].strides())?;
1017    try_static_specialization_validated::<T, InitializedStaticOutput>(
1018        &mut dests[0],
1019        inputs,
1020        plan,
1021        validated,
1022    )
1023}
1024
1025unsafe fn interpret_inner_loop<T: FusedScalar>(
1026    dst_ptrs: &[*mut T],
1027    input_ptrs: &[*const T],
1028    plan: &FusedPlan,
1029    offsets: &[isize],
1030    len: usize,
1031    strides: &[isize],
1032) {
1033    let output_count = dst_ptrs.len();
1034    let mut regs = Vec::with_capacity(plan.input_count + plan.ops.len());
1035
1036    for i in 0..len {
1037        let i = i as isize;
1038        regs.clear();
1039
1040        for (input_index, &input_ptr) in input_ptrs.iter().enumerate() {
1041            let stride_index = output_count + input_index;
1042            regs.push(*input_ptr.offset(offsets[stride_index] + i * strides[stride_index]));
1043        }
1044
1045        for inst in &plan.ops {
1046            regs.push(eval_op(inst.op, &regs, &inst.inputs));
1047        }
1048
1049        for (output_index, &dst_ptr) in dst_ptrs.iter().enumerate() {
1050            *dst_ptr.offset(offsets[output_index] + i * strides[output_index]) =
1051                regs[plan.outputs[output_index]];
1052        }
1053    }
1054}
1055
1056fn interpret_fused_elementwise_into<T: FusedScalar>(
1057    dests: &mut [StridedViewMut<'_, T>],
1058    inputs: &[StridedView<'_, T>],
1059    plan: &FusedPlan,
1060) -> Result<()> {
1061    #[cfg(feature = "parallel")]
1062    {
1063        let dims = dests[0].dims().to_vec();
1064        if dests[0].len() == 0 {
1065            return Ok(());
1066        }
1067
1068        let dst_ptrs: Vec<*mut T> = dests.iter_mut().map(|dest| dest.as_mut_ptr()).collect();
1069        let input_ptrs: Vec<*const T> = inputs.iter().map(StridedView::ptr).collect();
1070
1071        let mut strides_list: Vec<&[isize]> = Vec::with_capacity(dests.len() + inputs.len());
1072        for dest in dests.iter() {
1073            strides_list.push(dest.strides());
1074        }
1075        for input in inputs {
1076            strides_list.push(input.strides());
1077        }
1078
1079        let elem_size = std::mem::size_of::<T>();
1080        let total = dests[0].len();
1081        let (fused_dims, ordered_strides, kernel_plan) = if total <= SMALL_TENSOR_THRESHOLD {
1082            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1083            unsafe { build_plan_fused_small(&dims, &strides_list) }
1084        } else {
1085            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1086            unsafe { build_plan_fused(&dims, &strides_list, Some(0), elem_size) }
1087        };
1088
1089        let total: usize = fused_dims.iter().product();
1090        let nthreads = strided_basic::execution::rayon_threads();
1091        if total > MINTHREADLENGTH && nthreads > 1 {
1092            let dst_send: Vec<SendPtr<T>> = dst_ptrs.iter().map(|&ptr| SendPtr(ptr)).collect();
1093            let input_send: Vec<SendPtr<T>> = input_ptrs
1094                .iter()
1095                .map(|&ptr| SendPtr(ptr as *mut T))
1096                .collect();
1097
1098            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1099
1100            let costs = unsafe { compute_costs(&ordered_strides) };
1101            let initial_offsets = vec![0isize; ordered_strides.len()];
1102            let run_partition = |dims: &[usize],
1103                                 blocks: &[usize],
1104                                 strides_list: &[Vec<isize>],
1105                                 offsets: &[isize]|
1106             -> Result<()> {
1107                let dst_ptrs: Vec<*mut T> = dst_send.iter().map(|ptr| ptr.as_ptr()).collect();
1108                let input_ptrs: Vec<*const T> =
1109                    input_send.iter().map(|ptr| ptr.as_const()).collect();
1110                let run_block = |offsets: &[isize], len: usize, strides: &[isize]| -> Result<()> {
1111                    // SAFETY: validated shapes/layouts and the derived block bound every access.
1112                    unsafe {
1113                        interpret_inner_loop(&dst_ptrs, &input_ptrs, plan, offsets, len, strides);
1114                    }
1115                    Ok(())
1116                };
1117                // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1118                unsafe {
1119                    for_each_inner_block_preordered(dims, blocks, strides_list, offsets, run_block)
1120                }
1121            };
1122            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1123            return unsafe {
1124                mapreduce_threaded(
1125                    &fused_dims,
1126                    &kernel_plan.block,
1127                    &ordered_strides,
1128                    &initial_offsets,
1129                    &costs,
1130                    nthreads,
1131                    0,
1132                    1,
1133                    &run_partition,
1134                )
1135            };
1136        }
1137    }
1138
1139    interpret_fused_elementwise_into_serial(dests, inputs, plan)
1140}
1141
1142fn interpret_fused_elementwise_into_serial<T: FusedScalar>(
1143    dests: &mut [StridedViewMut<'_, T>],
1144    inputs: &[StridedView<'_, T>],
1145    plan: &FusedPlan,
1146) -> Result<()> {
1147    let dims = dests[0].dims().to_vec();
1148    if dests[0].len() == 0 {
1149        return Ok(());
1150    }
1151
1152    let dst_ptrs: Vec<*mut T> = dests.iter_mut().map(|dest| dest.as_mut_ptr()).collect();
1153    let input_ptrs: Vec<*const T> = inputs.iter().map(StridedView::ptr).collect();
1154
1155    let mut strides_list: Vec<&[isize]> = Vec::with_capacity(dests.len() + inputs.len());
1156    for dest in dests.iter() {
1157        strides_list.push(dest.strides());
1158    }
1159    for input in inputs {
1160        strides_list.push(input.strides());
1161    }
1162
1163    let elem_size = std::mem::size_of::<T>();
1164    let total = dests[0].len();
1165    let (fused_dims, ordered_strides, kernel_plan) = if total <= SMALL_TENSOR_THRESHOLD {
1166        // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1167        unsafe { build_plan_fused_small(&dims, &strides_list) }
1168    } else {
1169        // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1170        unsafe { build_plan_fused(&dims, &strides_list, Some(0), elem_size) }
1171    };
1172
1173    let initial_offsets = vec![0isize; ordered_strides.len()];
1174    let run_block = |offsets: &[isize], len: usize, strides: &[isize]| -> Result<()> {
1175        // SAFETY: validated shapes/layouts and the derived block bound every access.
1176        unsafe {
1177            interpret_inner_loop(&dst_ptrs, &input_ptrs, plan, offsets, len, strides);
1178        }
1179        Ok(())
1180    };
1181    // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1182    unsafe {
1183        for_each_inner_block_preordered(
1184            &fused_dims,
1185            &kernel_plan.block,
1186            &ordered_strides,
1187            &initial_offsets,
1188            run_block,
1189        )
1190    }
1191}
1192
1193pub(crate) fn fused_elementwise_into_serial<T: FusedScalar>(
1194    dests: &mut [StridedViewMut<'_, T>],
1195    inputs: &[StridedView<'_, T>],
1196    plan: &FusedPlan,
1197) -> Result<()> {
1198    validate_plan_for_scalar::<T>(plan, inputs.len(), dests.len())?;
1199    validate_shapes(dests, inputs)?;
1200    interpret_fused_elementwise_into_serial(dests, inputs, plan)
1201}
1202
1203unsafe fn interpret_inner_loop_uninit<T: FusedScalar>(
1204    dst_ptr: *mut MaybeUninit<T>,
1205    input_ptrs: &[*const T],
1206    plan: &FusedPlan,
1207    offsets: &[isize],
1208    len: usize,
1209    strides: &[isize],
1210) {
1211    let mut regs = Vec::with_capacity(plan.input_count + plan.ops.len());
1212    for i in 0..len {
1213        let i = i as isize;
1214        regs.clear();
1215        for (input_index, &input_ptr) in input_ptrs.iter().enumerate() {
1216            let stride_index = 1 + input_index;
1217            regs.push(*input_ptr.offset(offsets[stride_index] + i * strides[stride_index]));
1218        }
1219        for inst in &plan.ops {
1220            regs.push(eval_op(inst.op, &regs, &inst.inputs));
1221        }
1222        *dst_ptr.offset(offsets[0] + i * strides[0]) = MaybeUninit::new(regs[plan.outputs[0]]);
1223    }
1224}
1225
1226pub(crate) fn fused_elementwise_into_uninit<T: FusedScalar>(
1227    dest: &mut StridedViewMut<'_, MaybeUninit<T>>,
1228    inputs: &[StridedView<'_, T>],
1229    plan: &FusedPlan,
1230    serial: bool,
1231    validated: ValidatedDestinationLayout,
1232) -> Result<()> {
1233    #[cfg(not(feature = "parallel"))]
1234    let _ = serial;
1235    validate_plan_for_scalar::<T>(plan, inputs.len(), 1)?;
1236    for input in inputs {
1237        ensure_same_shape(dest.dims(), input.dims())?;
1238    }
1239
1240    if !serial
1241        && try_static_specialization_validated::<T, UninitializedStaticOutput>(
1242            dest, inputs, plan, validated,
1243        )?
1244    {
1245        return Ok(());
1246    }
1247
1248    let dims = dest.dims();
1249    if dest.len() == 0 {
1250        return Ok(());
1251    }
1252    let dst_ptr = dest.as_mut_ptr();
1253    let input_ptrs: Vec<*const T> = inputs.iter().map(StridedView::ptr).collect();
1254    let mut strides_list: Vec<&[isize]> = Vec::with_capacity(1 + inputs.len());
1255    strides_list.push(dest.strides());
1256    for input in inputs {
1257        strides_list.push(input.strides());
1258    }
1259    let total = dest.len();
1260    let (fused_dims, ordered_strides, kernel_plan) = if total <= SMALL_TENSOR_THRESHOLD {
1261        // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1262        unsafe { build_plan_fused_small(dims, &strides_list) }
1263    } else {
1264        // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1265        unsafe { build_plan_fused(dims, &strides_list, Some(0), core::mem::size_of::<T>()) }
1266    };
1267
1268    #[cfg(feature = "parallel")]
1269    {
1270        let total: usize = fused_dims.iter().product();
1271        let nthreads = strided_basic::execution::rayon_threads();
1272        if !serial && total > MINTHREADLENGTH && nthreads > 1 {
1273            let dst_send = SendPtr(dst_ptr);
1274            let input_send: Vec<SendPtr<T>> = input_ptrs
1275                .iter()
1276                .map(|&ptr| SendPtr(ptr as *mut T))
1277                .collect();
1278            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1279            let costs = unsafe { compute_costs(&ordered_strides) };
1280            let initial_offsets = vec![0isize; ordered_strides.len()];
1281            let run_partition = |dims: &[usize],
1282                                 blocks: &[usize],
1283                                 strides_list: &[Vec<isize>],
1284                                 offsets: &[isize]|
1285             -> Result<()> {
1286                let input_ptrs: Vec<*const T> =
1287                    input_send.iter().map(|ptr| ptr.as_const()).collect();
1288                let run_block = |offsets: &[isize], len: usize, strides: &[isize]| -> Result<()> {
1289                    // SAFETY: validated shapes/layouts and the derived block bound every access.
1290                    unsafe {
1291                        interpret_inner_loop_uninit(
1292                            dst_send.as_ptr(),
1293                            &input_ptrs,
1294                            plan,
1295                            offsets,
1296                            len,
1297                            strides,
1298                        );
1299                    }
1300                    Ok(())
1301                };
1302                // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1303                unsafe {
1304                    for_each_inner_block_preordered(dims, blocks, strides_list, offsets, run_block)
1305                }
1306            };
1307            // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1308            return unsafe {
1309                mapreduce_threaded(
1310                    &fused_dims,
1311                    &kernel_plan.block,
1312                    &ordered_strides,
1313                    &initial_offsets,
1314                    &costs,
1315                    nthreads,
1316                    0,
1317                    1,
1318                    &run_partition,
1319                )
1320            };
1321        }
1322    }
1323
1324    let initial_offsets = vec![0isize; ordered_strides.len()];
1325    let run_block = |offsets: &[isize], len: usize, strides: &[isize]| -> Result<()> {
1326        // SAFETY: validated shapes/layouts and the derived block bound every access.
1327        unsafe {
1328            interpret_inner_loop_uninit(dst_ptr, &input_ptrs, plan, offsets, len, strides);
1329        }
1330        Ok(())
1331    };
1332    // SAFETY: the enclosing checked kernel supplies validated layout metadata and its derived partitions.
1333    unsafe {
1334        for_each_inner_block_preordered(
1335            &fused_dims,
1336            &kernel_plan.block,
1337            &ordered_strides,
1338            &initial_offsets,
1339            run_block,
1340        )
1341    }
1342}
1343
1344/// Evaluate a runtime-DAG elementwise plan into one or more destinations.
1345///
1346/// The plan is validated before any destination is written:
1347///
1348/// - `inputs.len()` must equal `plan.input_count`;
1349/// - `dests.len()` must equal `plan.outputs.len()`;
1350/// - instruction operands must reference earlier SSA values with the right
1351///   arity for their [`FusedOp`];
1352/// - every input and destination must have exactly the destination shape;
1353/// - each mutable destination layout must be injective, so two logical output
1354///   elements never map to the same memory address.
1355///
1356/// The implementation dispatches known single-output plans to existing static
1357/// `map_into`/`zip_map*_into` kernels and uses a generic interpreter fallback
1358/// for arbitrary validated DAGs. Overlapping source/destination memory is not
1359/// supported by the strided kernels generally.
1360///
1361/// Real `Maximum`, `Minimum`, and `Clamp` use Rust `f32`/`f64` `max`/`min`
1362/// semantics. Complex `Abs` returns the norm in the real component; complex
1363/// `Maximum`, `Minimum`, and `Clamp` compare by squared norm. Signed integer
1364/// `Add`, `Multiply`, `Negate`, and `Abs` use wrapping arithmetic. `bool`
1365/// supports only copy-like identity plans and `Conj`; ambiguous arithmetic and
1366/// transcendental op/dtype pairs are rejected before any destination is written.
1367pub fn fused_elementwise_into<T: FusedScalar>(
1368    dests: &mut [StridedViewMut<'_, T>],
1369    inputs: &[StridedView<'_, T>],
1370    plan: &FusedPlan,
1371) -> Result<()> {
1372    validate_plan_for_scalar::<T>(plan, inputs.len(), dests.len())?;
1373    validate_shapes(dests, inputs)?;
1374    if try_static_specialization(dests, inputs, plan)? {
1375        return Ok(());
1376    }
1377    interpret_fused_elementwise_into(dests, inputs, plan)
1378}
1379
1380#[cfg(test)]
1381#[path = "fused/tests/tests.rs"]
1382mod tests;