1use 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#[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#[derive(Clone, Debug, Eq, PartialEq)]
74pub struct FusedInst {
75 pub op: FusedOp,
76 pub inputs: Vec<usize>,
77}
78
79#[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
98pub 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 unsafe {
935 map_into_validated(dest, &inputs[a], |x| O::write(eval_unary(op, x)), validated)
936 }?
937 }
938 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 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 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 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 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 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, ®s, &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 unsafe { build_plan_fused_small(&dims, &strides_list) }
1090 } else {
1091 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 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 unsafe {
1119 interpret_inner_loop(&dst_ptrs, &input_ptrs, plan, offsets, len, strides);
1120 }
1121 Ok(())
1122 };
1123 unsafe {
1125 for_each_inner_block_preordered(dims, blocks, strides_list, offsets, run_block)
1126 }
1127 };
1128 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 unsafe { build_plan_fused_small(&dims, &strides_list) }
1174 } else {
1175 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 unsafe {
1183 interpret_inner_loop(&dst_ptrs, &input_ptrs, plan, offsets, len, strides);
1184 }
1185 Ok(())
1186 };
1187 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, ®s, &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 unsafe { build_plan_fused_small(dims, &strides_list) }
1269 } else {
1270 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 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 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 unsafe {
1310 for_each_inner_block_preordered(dims, blocks, strides_list, offsets, run_block)
1311 }
1312 };
1313 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 unsafe {
1334 interpret_inner_loop_uninit(dst_ptr, &input_ptrs, plan, offsets, len, strides);
1335 }
1336 Ok(())
1337 };
1338 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
1350pub 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;