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) => {
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 unsafe {
929 map_into_validated(dest, &inputs[a], |x| O::write(eval_unary(op, x)), validated)
930 }?
931 }
932 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 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 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 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 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 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, ®s, &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 unsafe { build_plan_fused_small(&dims, &strides_list) }
1084 } else {
1085 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 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 unsafe {
1113 interpret_inner_loop(&dst_ptrs, &input_ptrs, plan, offsets, len, strides);
1114 }
1115 Ok(())
1116 };
1117 unsafe {
1119 for_each_inner_block_preordered(dims, blocks, strides_list, offsets, run_block)
1120 }
1121 };
1122 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 unsafe { build_plan_fused_small(&dims, &strides_list) }
1168 } else {
1169 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 unsafe {
1177 interpret_inner_loop(&dst_ptrs, &input_ptrs, plan, offsets, len, strides);
1178 }
1179 Ok(())
1180 };
1181 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, ®s, &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 unsafe { build_plan_fused_small(dims, &strides_list) }
1263 } else {
1264 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 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 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 unsafe {
1304 for_each_inner_block_preordered(dims, blocks, strides_list, offsets, run_block)
1305 }
1306 };
1307 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 unsafe {
1328 interpret_inner_loop_uninit(dst_ptr, &input_ptrs, plan, offsets, len, strides);
1329 }
1330 Ok(())
1331 };
1332 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
1344pub 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;