1#![allow(private_bounds)]
2use std::iter::zip;
3use std::ops::{Add, AddAssign, Div, DivAssign, Mul, MulAssign, Neg, Sub, SubAssign};
4
5use crate::tensor::backend::Backend;
6use crate::tensor::errors::OpError;
7use crate::tensor::graph::NodeKind;
8use crate::tensor::mem_formats::layout::Layout;
9use crate::tensor::mem_formats::slice::SliceRange;
10use crate::tensor::ops::capabilities::{CanMatMul, FloatLike, NumericOp};
11use crate::tensor::ops::compute_layout;
12use crate::tensor::ops::def_op::{OpKind, OpKindScalar};
13use crate::tensor::skeleton::{
14 BakedPromise, BinaryResult, Clean, SkeletonPromise, SkeletonSlot, Tainting, UnaryResult,
15};
16use crate::tensor::traits::{Dimension, Numeric, Operand};
17use crate::tensor::{CachedTensorPromise, Tensor, TensorPromise};
18
19struct NodeWithLayout<T: Numeric, B: Backend> {
22 node: NodeKind<T, B>,
23 layout: Layout,
24}
25
26impl<T: Numeric, B: Backend> Dimension for NodeWithLayout<T, B> {
27 fn layout(&self) -> &Layout {
28 &self.layout
29 }
30}
31
32impl<T: Numeric, B: Backend> Operand<T, B> for NodeWithLayout<T, B> {
33 fn to_node(&self) -> NodeKind<T, B> {
34 self.node.clone()
35 }
36}
37
38impl<T: Numeric, B: Backend> Tainting for NodeWithLayout<T, B> {
39 type Mark = Clean;
40}
41
42#[inline]
47fn find_broadcast_target(l1: &Layout, l2: &Layout) -> Vec<usize> {
48 let (largest, smallest) = if l1.shape().len() >= l2.shape().len() {
49 (l1, l2)
50 } else {
51 (l2, l1)
52 };
53 let largest_size = largest.shape().len();
54 debug_assert!(
55 largest.shape().len() >= smallest.shape().len(),
56 "broadcast helper precondition violated"
57 );
58 let diff = largest.shape().len() - smallest.shape().len();
59
60 let mut new_shape = vec![0_usize; largest_size];
61
62 for (i, (&dim1, &dim2)) in zip(l1.shape().iter().rev(), l2.shape().iter().rev()).enumerate() {
63 new_shape[largest_size - i - 1] = dim1.max(dim2);
64 }
65
66 new_shape[..diff].copy_from_slice(&largest.shape()[..diff]);
67
68 new_shape
69}
70
71#[inline]
72fn find_broadcast_target_until_batch(l1: &Layout, l2: &Layout) -> Option<(Vec<usize>, Vec<usize>)> {
73 let (largest, smallest) = if l1.shape().len() >= l2.shape().len() {
74 (l1, l2)
75 } else {
76 (l2, l1)
77 };
78 let largest_size = largest.shape().len();
79 debug_assert!(
80 largest.shape().len() >= smallest.shape().len(),
81 "matmul-batch broadcast helper precondition violated"
82 );
83 let diff = largest.shape().len() - smallest.shape().len();
84 let smallest_diff: usize = 2.min(smallest.shape().len());
85
86 if largest_size <= 2 {
87 return None;
88 }
89
90 let mut new_l1_shape = vec![0_usize; largest_size];
91 let mut new_l2_shape = vec![0_usize; largest_size];
92
93 for dim in 0..smallest_diff {
94 new_l1_shape[largest_size - dim - 1] = l1.shape()[l1.shape().len() - dim - 1];
95 new_l2_shape[largest_size - dim - 1] = l2.shape()[l2.shape().len() - dim - 1];
96 }
97
98 for (i, (&dim1, &dim2)) in zip(l1.shape().iter().rev(), l2.shape().iter().rev())
99 .enumerate()
100 .skip(smallest_diff)
101 {
102 let max = dim1.max(dim2);
103 new_l1_shape[largest_size - i - 1] = max;
104 new_l2_shape[largest_size - i - 1] = max;
105 }
106
107 for dim in 0..diff {
108 let n = largest.shape()[dim];
109 new_l1_shape[dim] = n;
110 new_l2_shape[dim] = n;
111 }
112
113 Some((new_l1_shape, new_l2_shape))
114}
115
116#[inline]
117fn is_blas_ready<T, B, D>(source: &D) -> bool
118where
119 B: Backend,
120 D: Operand<T, B>,
121{
122 let layout = source.layout();
123 if B::SUPPORTS_NON_CONTIGUOUS_MATMUL {
124 let last_axis = layout.stride().len() - 1;
125 return layout.stride()[last_axis] != 0 && layout.stride()[last_axis - 1] != 0;
126 }
127
128 B::SUPPORTS_2D_TRANSPOSED_MATMUL && layout.is_last_axes_transposed()
132}
133
134type NodeTransform<Output, Backend> = Result<
135 (
136 NodeWithLayout<Output, Backend>,
137 NodeWithLayout<Output, Backend>,
138 Layout,
139 ),
140 OpError,
141>;
142
143#[inline]
144fn apply_transform_to_pair<T, B, D1, D2, F, N1, N2, L>(
145 lhs: &D1,
146 rhs: &D2,
147 filter: F,
148 transform_l: N1,
149 transform_r: N2,
150 compute_output_layout: L,
151) -> NodeTransform<T, B>
152where
153 T: Numeric,
154 B: Backend,
155 D1: Operand<T, B>,
156 D2: Operand<T, B>,
157 F: FnOnce(&D1, &D2) -> (bool, bool),
158 N1: FnOnce(&D1) -> Result<TensorPromise<T, B>, OpError>,
159 N2: FnOnce(&D2) -> Result<TensorPromise<T, B>, OpError>,
160 L: FnOnce(&Layout, &Layout) -> Result<Layout, OpError>,
161{
162 let (apply_l, apply_r) = filter(lhs, rhs);
163
164 let (node1, layout1, node2, layout2) = match (apply_l, apply_r) {
165 (false, false) => (
166 lhs.to_node(),
167 lhs.layout().clone(),
168 rhs.to_node(),
169 rhs.layout().clone(),
170 ),
171 (true, false) => {
172 let temp = transform_l(lhs)?;
173 let layout = temp.layout().clone();
174 (
175 NodeKind::Node(temp.graph),
176 layout,
177 rhs.to_node(),
178 rhs.layout().clone(),
179 )
180 }
181 (false, true) => {
182 let temp = transform_r(rhs)?;
183 let layout = temp.layout().clone();
184 (
185 lhs.to_node(),
186 lhs.layout().clone(),
187 NodeKind::Node(temp.graph),
188 layout,
189 )
190 }
191 (true, true) => {
192 let temp1 = transform_l(lhs)?;
193 let layout1 = temp1.layout().clone();
194 let temp2 = transform_r(rhs)?;
195 let layout2 = temp2.layout().clone();
196 (
197 NodeKind::Node(temp1.graph),
198 layout1,
199 NodeKind::Node(temp2.graph),
200 layout2,
201 )
202 }
203 };
204
205 let layout = compute_output_layout(&layout1, &layout2)?;
206 Ok((
207 NodeWithLayout {
208 node: node1,
209 layout: layout1,
210 },
211 NodeWithLayout {
212 node: node2,
213 layout: layout2,
214 },
215 layout,
216 ))
217}
218
219fn view_impl<T, B, D>(source: &D, shape: &[usize]) -> Result<TensorPromise<T, B>, OpError>
222where
223 T: Numeric,
224 B: Backend,
225 D: Operand<T, B>,
226{
227 let input = Box::new([source.to_node()]);
228 let layout = source.layout().view(shape)?;
229
230 Ok(TensorPromise::with_layout(
231 OpKind::View(layout.clone()),
232 input,
233 layout,
234 ))
235}
236
237fn broadcast_impl<T, B, D>(source: &D, shape: &[usize]) -> Result<TensorPromise<T, B>, OpError>
238where
239 T: Numeric,
240 B: Backend,
241 D: Operand<T, B>,
242{
243 let input = Box::new([source.to_node()]);
244 let layout = source.layout().broadcast(shape)?;
245
246 Ok(TensorPromise::with_layout(
247 OpKind::Broadcast(layout.clone()),
248 input,
249 layout,
250 ))
251}
252
253fn reshape_impl<T, B, D>(source: &D, shape: &[usize]) -> Result<TensorPromise<T, B>, OpError>
254where
255 T: Numeric,
256 B: Backend,
257 D: Operand<T, B>,
258{
259 let cont: TensorPromise<T, B> = as_contiguous_impl(source);
260 let layout = cont.graph.layout.view(shape)?;
261 let input = Box::new([NodeKind::Node(cont.graph)]);
262
263 Ok(TensorPromise::with_layout(
264 OpKind::View(layout.clone()),
265 input,
266 layout,
267 ))
268}
269
270fn slice_impl<T, B, D>(source: &D, range: &[SliceRange]) -> Result<TensorPromise<T, B>, OpError>
271where
272 T: Numeric,
273 B: Backend,
274 D: Operand<T, B>,
275{
276 let input = Box::new([source.to_node()]);
277 let layout = source.layout().slice(range)?;
278
279 Ok(TensorPromise::with_layout(
280 OpKind::Slice(layout.clone()),
281 input,
282 layout,
283 ))
284}
285
286fn transpose_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
287where
288 T: Numeric,
289 B: Backend,
290 D: Operand<T, B>,
291{
292 let input = Box::new([source.to_node()]);
293
294 unsafe { TensorPromise::new(OpKind::Transpose, input).unwrap_unchecked() }
295}
296
297fn transpose_axes_impl<T, B, D>(source: &D, axes: &[usize]) -> Result<TensorPromise<T, B>, OpError>
298where
299 T: Numeric,
300 B: Backend,
301 D: Operand<T, B>,
302{
303 let input = Box::new([source.to_node()]);
304 let layout = source.layout().transpose_axes(axes)?;
305
306 Ok(TensorPromise::with_layout(
307 OpKind::TransposeAxes(layout.clone()),
308 input,
309 layout,
310 ))
311}
312
313fn as_contiguous_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
314where
315 T: Numeric,
316 B: Backend,
317 D: Operand<T, B>,
318{
319 let node = source.to_node();
320 unsafe { TensorPromise::new(OpKind::AsContiguous, Box::new([node])).unwrap_unchecked() }
321}
322
323fn add_scalar_impl<T, B, D>(lhs: &D, rhs: T) -> TensorPromise<T, B>
326where
327 T: Numeric,
328 B: Backend,
329 D: Operand<T, B>,
330{
331 unsafe {
332 TensorPromise::new(
333 OpKind::ScalarOp(OpKindScalar::AxBy(T::MUL_NEUTRAL, rhs)),
334 Box::new([lhs.to_node()]),
335 )
336 .unwrap_unchecked()
337 }
338}
339
340fn sub_scalar_impl<T, B, D>(lhs: &D, rhs: T) -> TensorPromise<T, B>
341where
342 T: Numeric,
343 B: Backend,
344 D: Operand<T, B>,
345 T: Numeric + Neg<Output = T>,
346{
347 unsafe {
348 TensorPromise::new(
349 OpKind::ScalarOp(OpKindScalar::AxBy(T::MUL_NEUTRAL, -rhs)),
350 Box::new([lhs.to_node()]),
351 )
352 .unwrap_unchecked()
353 }
354}
355
356fn mul_scalar_impl<T, B, D>(lhs: &D, rhs: T) -> TensorPromise<T, B>
357where
358 T: Numeric,
359 B: Backend,
360 D: Operand<T, B>,
361{
362 unsafe {
363 TensorPromise::new(
364 OpKind::ScalarOp(OpKindScalar::AxBy(rhs, T::SUM_NEUTRAL)),
365 Box::new([lhs.to_node()]),
366 )
367 .unwrap_unchecked()
368 }
369}
370
371fn div_scalar_impl<T, B, D>(lhs: &D, rhs: T) -> TensorPromise<T, B>
372where
373 T: Numeric,
374 B: Backend,
375 D: Operand<T, B>,
376{
377 if rhs == T::SUM_NEUTRAL {
378 panic!("cannot divide by zero. stop.")
379 }
380
381 unsafe {
382 TensorPromise::new(
383 OpKind::ScalarOp(OpKindScalar::AxBy(T::MUL_NEUTRAL / rhs, T::SUM_NEUTRAL)),
384 Box::new([lhs.to_node()]),
385 )
386 .unwrap_unchecked()
387 }
388}
389
390fn exp_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
391where
392 T: Numeric,
393 B: Backend,
394 D: Operand<T, B>,
395{
396 let input = Box::new([source.to_node()]);
397
398 unsafe { TensorPromise::new(OpKind::ScalarOp(OpKindScalar::Exp), input).unwrap_unchecked() }
399}
400
401fn ln_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
402where
403 T: Numeric,
404 B: Backend,
405 D: Operand<T, B>,
406{
407 let input = Box::new([source.to_node()]);
408
409 unsafe { TensorPromise::new(OpKind::ScalarOp(OpKindScalar::Ln), input).unwrap_unchecked() }
410}
411
412fn log2_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
413where
414 T: Numeric,
415 B: Backend,
416 D: Operand<T, B>,
417{
418 let input = Box::new([source.to_node()]);
419
420 unsafe { TensorPromise::new(OpKind::ScalarOp(OpKindScalar::Log2), input).unwrap_unchecked() }
421}
422
423fn relu_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
424where
425 T: Numeric,
426 B: Backend,
427 D: Operand<T, B>,
428{
429 let input = Box::new([source.to_node()]);
430
431 unsafe { TensorPromise::new(OpKind::ScalarOp(OpKindScalar::ReLU), input).unwrap_unchecked() }
432}
433
434fn tanh_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
435where
436 T: Numeric,
437 B: Backend,
438 D: Operand<T, B>,
439{
440 let input = Box::new([source.to_node()]);
441
442 unsafe { TensorPromise::new(OpKind::ScalarOp(OpKindScalar::Tanh), input).unwrap_unchecked() }
443}
444
445fn add_tensor_impl<T, B, D1, D2>(lhs: &D1, rhs: &D2) -> TensorPromise<T, B>
448where
449 T: Numeric,
450 B: Backend,
451 D1: Operand<T, B>,
452 D2: Operand<T, B>,
453{
454 let target = find_broadcast_target(lhs.layout(), rhs.layout());
455
456 let result = apply_transform_to_pair(
457 lhs,
458 rhs,
459 |l, r| (l.layout().shape() != target, r.layout().shape() != target),
460 |x| broadcast_impl(x, &target),
461 |x| broadcast_impl(x, &target),
462 |l1, l2| compute_layout(&OpKind::<T>::Add, &[l1, l2]),
463 );
464
465 if let Err(err) = result {
466 panic!("{}", err);
467 }
468
469 let (lhs_b, rhs_b, layout) = unsafe { result.unwrap_unchecked() };
470 TensorPromise::with_layout(
471 OpKind::Add,
472 [lhs_b.to_node(), rhs_b.to_node()].into(),
473 layout,
474 )
475}
476
477fn sub_tensor_impl<T, B, D1, D2>(lhs: &D1, rhs: &D2) -> TensorPromise<T, B>
478where
479 T: Numeric,
480 B: Backend,
481 D1: Operand<T, B>,
482 D2: Operand<T, B>,
483{
484 let target = find_broadcast_target(lhs.layout(), rhs.layout());
485
486 let result = apply_transform_to_pair(
487 lhs,
488 rhs,
489 |l, r| (l.layout().shape() != target, r.layout().shape() != target),
490 |x| broadcast_impl(x, &target),
491 |x| broadcast_impl(x, &target),
492 |l1, l2| compute_layout(&OpKind::<T>::Sub, &[l1, l2]),
493 );
494
495 if let Err(err) = result {
496 panic!("{}", err);
497 }
498
499 let (lhs_b, rhs_b, layout) = unsafe { result.unwrap_unchecked() };
500 TensorPromise::with_layout(
501 OpKind::Sub,
502 [lhs_b.to_node(), rhs_b.to_node()].into(),
503 layout,
504 )
505}
506
507fn mul_tensor_impl<T, B, D1, D2>(lhs: &D1, rhs: &D2) -> TensorPromise<T, B>
508where
509 T: Numeric,
510 B: Backend,
511 D1: Operand<T, B>,
512 D2: Operand<T, B>,
513{
514 let target = find_broadcast_target(lhs.layout(), rhs.layout());
515
516 let result = apply_transform_to_pair(
517 lhs,
518 rhs,
519 |l, r| (l.layout().shape() != target, r.layout().shape() != target),
520 |x| broadcast_impl(x, &target),
521 |x| broadcast_impl(x, &target),
522 |l1, l2| compute_layout(&OpKind::<T>::Mul, &[l1, l2]),
523 );
524
525 if let Err(err) = result {
526 panic!("{}", err);
527 }
528
529 let (lhs_b, rhs_b, layout) = unsafe { result.unwrap_unchecked() };
530 TensorPromise::with_layout(
531 OpKind::Mul,
532 [lhs_b.to_node(), rhs_b.to_node()].into(),
533 layout,
534 )
535}
536
537fn div_tensor_impl<T, B, D1, D2>(lhs: &D1, rhs: &D2) -> TensorPromise<T, B>
538where
539 T: Numeric,
540 B: Backend,
541 D1: Operand<T, B>,
542 D2: Operand<T, B>,
543{
544 let target = find_broadcast_target(lhs.layout(), rhs.layout());
545
546 let result = apply_transform_to_pair(
547 lhs,
548 rhs,
549 |l, r| (l.layout().shape() != target, r.layout().shape() != target),
550 |x| broadcast_impl(x, &target),
551 |x| broadcast_impl(x, &target),
552 |l1, l2| compute_layout(&OpKind::<T>::Div, &[l1, l2]),
553 );
554
555 if let Err(err) = result {
556 panic!("{}", err);
557 }
558
559 let (lhs_b, rhs_b, layout) = unsafe { result.unwrap_unchecked() };
560 TensorPromise::with_layout(
561 OpKind::Div,
562 [lhs_b.to_node(), rhs_b.to_node()].into(),
563 layout,
564 )
565}
566
567fn matmul_core<T, B, D1, D2>(lhs: &D1, rhs: &D2) -> Result<TensorPromise<T, B>, OpError>
570where
571 T: Numeric,
572 B: Backend,
573 D1: Operand<T, B>,
574 D2: Operand<T, B>,
575{
576 let (lhs_c, rhs_c, _) = apply_transform_to_pair(
577 lhs,
578 rhs,
579 |l, r| (!is_blas_ready(l), !is_blas_ready(r)),
580 |x| Ok(as_contiguous_impl(x)),
581 |x| Ok(as_contiguous_impl(x)),
582 |l1, _| Ok(l1.clone()),
583 )?;
584
585 let target = find_broadcast_target_until_batch(lhs_c.layout(), rhs_c.layout());
586
587 let (lhs_b, rhs_b, layout) = apply_transform_to_pair(
588 &lhs_c,
589 &rhs_c,
590 |l, r| {
591 (
592 target
593 .as_ref()
594 .is_some_and(|target| l.layout().shape() != target.0),
595 target
596 .as_ref()
597 .is_some_and(|target| r.layout().shape() != target.1),
598 )
599 },
600 |x| broadcast_impl(x, unsafe { &target.as_ref().unwrap_unchecked().0 }),
601 |x| broadcast_impl(x, unsafe { &target.as_ref().unwrap_unchecked().1 }),
602 |l1, l2| compute_layout(&OpKind::<T>::MatMul(T::MUL_NEUTRAL), &[l1, l2]),
603 )?;
604
605 Ok(TensorPromise::with_layout(
606 OpKind::MatMul(T::MUL_NEUTRAL),
607 [lhs_b.to_node(), rhs_b.to_node()].into(),
608 layout,
609 ))
610}
611
612fn drop_dim_from_end<T, B, D>(source: &D, from_end: usize) -> Result<TensorPromise<T, B>, OpError>
616where
617 T: Numeric,
618 B: Backend,
619 D: Operand<T, B>,
620{
621 let mut new_shape: Vec<usize> = source.layout().shape().to_vec();
622 new_shape.remove(new_shape.len() - 1 - from_end);
623 view_impl(source, &new_shape)
624}
625
626fn matmul_tensor_impl<T, B, D1, D2>(lhs: &D1, rhs: &D2) -> Result<TensorPromise<T, B>, OpError>
627where
628 T: Numeric,
629 B: Backend,
630 D1: Operand<T, B>,
631 D2: Operand<T, B>,
632{
633 match (lhs.layout().shape().len(), rhs.layout().shape().len()) {
634 (1, 1) => {
636 let lhs_p = reshape_impl(lhs, &[1, lhs.layout().shape()[0]])?;
637 let rhs_p = reshape_impl(rhs, &[rhs.layout().shape()[0], 1])?;
638 let result = matmul_core(&lhs_p, &rhs_p)?;
639 drop_dim_from_end(&result, 0)
640 }
641 (1, _) => {
643 let lhs_p = reshape_impl(lhs, &[1, lhs.layout().shape()[0]])?;
644 let result = matmul_core(&lhs_p, rhs)?;
645 drop_dim_from_end(&result, 1)
646 }
647 (_, 1) => {
649 let rhs_p = reshape_impl(rhs, &[rhs.layout().shape()[0], 1])?;
650 let result = matmul_core(lhs, &rhs_p)?;
651 drop_dim_from_end(&result, 0)
652 }
653 _ => matmul_core(lhs, rhs),
655 }
656}
657
658fn sum_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
661where
662 T: Numeric,
663 B: Backend,
664 D: Operand<T, B>,
665{
666 let input = Box::new([source.to_node()]);
667
668 unsafe { TensorPromise::new(OpKind::Sum, input).unwrap_unchecked() }
669}
670
671fn sum_axis_impl<T, B, D>(
672 source: &D,
673 axis: isize,
674 keep_dims: bool,
675) -> Result<TensorPromise<T, B>, OpError>
676where
677 T: Numeric,
678 B: Backend,
679 D: Operand<T, B>,
680{
681 let input = Box::new([source.to_node()]);
682 let op = OpKind::<T>::SumAxis(axis, keep_dims);
683 let layout = compute_layout(&op, &[source.layout()])?;
684
685 Ok(TensorPromise::with_layout(op, input, layout))
686}
687
688fn max_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
689where
690 T: Numeric,
691 B: Backend,
692 D: Operand<T, B>,
693{
694 let input = Box::new([source.to_node()]);
695
696 unsafe { TensorPromise::new(OpKind::Max, input).unwrap_unchecked() }
697}
698
699fn max_axis_impl<T, B, D>(
700 source: &D,
701 axis: isize,
702 keep_dims: bool,
703) -> Result<TensorPromise<T, B>, OpError>
704where
705 T: Numeric,
706 B: Backend,
707 D: Operand<T, B>,
708{
709 let input = Box::new([source.to_node()]);
710 let op = OpKind::<T>::MaxAxis(axis, keep_dims);
711 let layout = compute_layout(&op, &[source.layout()])?;
712
713 Ok(TensorPromise::with_layout(op, input, layout))
714}
715
716fn mean_impl<T, B, D>(source: &D) -> TensorPromise<T, B>
717where
718 T: Numeric,
719 B: Backend,
720 D: Operand<T, B>,
721{
722 let input = Box::new([source.to_node()]);
723
724 unsafe { TensorPromise::new(OpKind::Mean, input).unwrap_unchecked() }
725}
726
727fn mean_axis_impl<T, B, D>(
728 source: &D,
729 axis: isize,
730 keep_dims: bool,
731) -> Result<TensorPromise<T, B>, OpError>
732where
733 T: Numeric,
734 B: Backend,
735 D: Operand<T, B>,
736{
737 let input = Box::new([source.to_node()]);
738 let op = OpKind::<T>::MeanAxis(axis, keep_dims);
739 let layout = compute_layout(&op, &[source.layout()])?;
740
741 Ok(TensorPromise::with_layout(op, input, layout))
742}
743
744macro_rules! impl_view {
747 ($ty:ident) => {
748 impl<T, B> $ty<T, B>
749 where
750 T: Numeric,
751 B: Backend,
752 {
753 #[inline]
774 pub fn view(
775 &self,
776 shape: &[usize],
777 ) -> Result<<$ty<T, B> as UnaryResult<T, B>>::Output, OpError> {
778 view_impl(self, shape).map(<$ty<T, B> as UnaryResult<T, B>>::wrap)
779 }
780
781 #[inline]
806 pub fn reshape(
807 &self,
808 shape: &[usize],
809 ) -> Result<<$ty<T, B> as UnaryResult<T, B>>::Output, OpError> {
810 reshape_impl(self, shape).map(<$ty<T, B> as UnaryResult<T, B>>::wrap)
811 }
812 }
813 };
814}
815
816macro_rules! impl_slice {
817 ($ty:ident) => {
818 impl<T, B> $ty<T, B>
819 where
820 T: Numeric,
821 B: Backend,
822 {
823 #[inline]
846 pub fn slice(
847 &self,
848 shape: &[SliceRange],
849 ) -> Result<<$ty<T, B> as UnaryResult<T, B>>::Output, OpError> {
850 slice_impl(self, shape).map(<$ty<T, B> as UnaryResult<T, B>>::wrap)
851 }
852 }
853 };
854}
855
856macro_rules! impl_transpose {
857 ($ty: ident) => {
858 impl<T, B> $ty<T, B>
859 where
860 T: Numeric,
861 B: Backend,
862 {
863 #[inline]
883 pub fn transpose(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
884 <$ty<T, B> as UnaryResult<T, B>>::wrap(transpose_impl(self))
885 }
886 }
887 };
888}
889
890macro_rules! impl_transpose_axes {
891 ($ty:ident) => {
892 impl<T, B> $ty<T, B>
893 where
894 T: Numeric,
895 B: Backend,
896 {
897 #[inline]
923 pub fn transpose_axes(
924 &self,
925 axes: &[usize],
926 ) -> Result<<$ty<T, B> as UnaryResult<T, B>>::Output, OpError> {
927 transpose_axes_impl(self, axes).map(<$ty<T, B> as UnaryResult<T, B>>::wrap)
928 }
929 }
930 };
931}
932
933macro_rules! impl_as_contiguous {
934 ($ty: ident) => {
935 impl<T, B> $ty<T, B>
936 where
937 T: Numeric,
938 B: Backend,
939 {
940 #[inline]
963 pub fn as_contiguous(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
964 <$ty<T, B> as UnaryResult<T, B>>::wrap(as_contiguous_impl(self))
965 }
966 }
967 };
968}
969
970macro_rules! impl_broadcast {
971 ($ty:ident) => {
972 impl<T, B> $ty<T, B>
973 where
974 T: Numeric,
975 B: Backend,
976 {
977 #[inline]
1003 pub fn broadcast(
1004 &self,
1005 shape: &[usize],
1006 ) -> Result<<$ty<T, B> as UnaryResult<T, B>>::Output, OpError> {
1007 broadcast_impl(self, shape).map(<$ty<T, B> as UnaryResult<T, B>>::wrap)
1008 }
1009 }
1010 };
1011}
1012
1013macro_rules! impl_reshape_like {
1014 ($ty:ident) => {
1015 impl_view!($ty);
1016 impl_slice!($ty);
1017 impl_transpose!($ty);
1018 impl_transpose_axes!($ty);
1019 impl_as_contiguous!($ty);
1020 impl_broadcast!($ty);
1021 };
1022}
1023macro_rules! impl_add_scalar {
1026 ($ty:ident) => {
1027 impl<T, B> Add<T> for &$ty<T, B>
1028 where
1029 T: NumericOp,
1030 B: Backend,
1031 {
1032 type Output = <$ty<T, B> as UnaryResult<T, B>>::Output;
1033
1034 #[inline]
1035 fn add(self, rhs: T) -> Self::Output {
1036 <$ty<T, B> as UnaryResult<T, B>>::wrap(add_scalar_impl(self, rhs))
1037 }
1038 }
1039
1040 impl<T, B> Add<T> for $ty<T, B>
1041 where
1042 T: NumericOp,
1043 B: Backend,
1044 {
1045 type Output = <$ty<T, B> as UnaryResult<T, B>>::Output;
1046
1047 #[inline]
1048 fn add(self, rhs: T) -> Self::Output {
1049 (&self).add(rhs)
1050 }
1051 }
1052 };
1053}
1054
1055macro_rules! impl_sub_scalar {
1056 ($ty:ident) => {
1057 impl<T, B> Sub<T> for &$ty<T, B>
1058 where
1059 T: NumericOp + Neg<Output = T>,
1060 B: Backend,
1061 {
1062 type Output = <$ty<T, B> as UnaryResult<T, B>>::Output;
1063
1064 #[inline]
1065 fn sub(self, rhs: T) -> Self::Output {
1066 <$ty<T, B> as UnaryResult<T, B>>::wrap(sub_scalar_impl(self, rhs))
1067 }
1068 }
1069
1070 impl<T, B> Sub<T> for $ty<T, B>
1071 where
1072 T: NumericOp + Neg<Output = T>,
1073 B: Backend,
1074 {
1075 type Output = <$ty<T, B> as UnaryResult<T, B>>::Output;
1076
1077 #[inline]
1078 fn sub(self, rhs: T) -> Self::Output {
1079 (&self).sub(rhs)
1080 }
1081 }
1082 };
1083}
1084
1085macro_rules! impl_mul_scalar {
1086 ($ty:ident) => {
1087 impl<T, B> Mul<T> for &$ty<T, B>
1088 where
1089 T: NumericOp,
1090 B: Backend,
1091 {
1092 type Output = <$ty<T, B> as UnaryResult<T, B>>::Output;
1093
1094 #[inline]
1095 fn mul(self, rhs: T) -> Self::Output {
1096 <$ty<T, B> as UnaryResult<T, B>>::wrap(mul_scalar_impl(self, rhs))
1097 }
1098 }
1099
1100 impl<T, B> Mul<T> for $ty<T, B>
1101 where
1102 T: NumericOp,
1103 B: Backend,
1104 {
1105 type Output = <$ty<T, B> as UnaryResult<T, B>>::Output;
1106
1107 #[inline]
1108 fn mul(self, rhs: T) -> Self::Output {
1109 (&self).mul(rhs)
1110 }
1111 }
1112 };
1113}
1114
1115macro_rules! impl_div_scalar {
1116 ($ty:ident) => {
1117 impl<T, B> Div<T> for &$ty<T, B>
1118 where
1119 T: NumericOp,
1120 B: Backend,
1121 {
1122 type Output = <$ty<T, B> as UnaryResult<T, B>>::Output;
1123
1124 #[inline]
1129 fn div(self, rhs: T) -> Self::Output {
1130 <$ty<T, B> as UnaryResult<T, B>>::wrap(div_scalar_impl(self, rhs))
1131 }
1132 }
1133
1134 impl<T, B> Div<T> for $ty<T, B>
1135 where
1136 T: NumericOp,
1137 B: Backend,
1138 {
1139 type Output = <$ty<T, B> as UnaryResult<T, B>>::Output;
1140
1141 #[inline]
1146 fn div(self, rhs: T) -> Self::Output {
1147 (&self).div(rhs)
1148 }
1149 }
1150 };
1151}
1152
1153macro_rules! impl_exp {
1154 ($ty:ident) => {
1155 impl<T, B> $ty<T, B>
1156 where
1157 T: FloatLike,
1158 B: Backend,
1159 {
1160 #[inline]
1171 pub fn exp(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
1172 <$ty<T, B> as UnaryResult<T, B>>::wrap(exp_impl(self))
1173 }
1174 }
1175 };
1176}
1177
1178macro_rules! impl_ln {
1179 ($ty:ident) => {
1180 impl<T, B> $ty<T, B>
1181 where
1182 T: FloatLike,
1183 B: Backend,
1184 {
1185 #[inline]
1198 pub fn ln(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
1199 <$ty<T, B> as UnaryResult<T, B>>::wrap(ln_impl(self))
1200 }
1201 }
1202 };
1203}
1204
1205macro_rules! impl_log2 {
1206 ($ty:ident) => {
1207 impl<T, B> $ty<T, B>
1208 where
1209 T: FloatLike,
1210 B: Backend,
1211 {
1212 #[inline]
1225 pub fn log2(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
1226 <$ty<T, B> as UnaryResult<T, B>>::wrap(log2_impl(self))
1227 }
1228 }
1229 };
1230}
1231
1232macro_rules! impl_relu {
1233 ($ty:ident) => {
1234 impl<T, B> $ty<T, B>
1235 where
1236 T: FloatLike,
1237 B: Backend,
1238 {
1239 #[inline]
1251 pub fn relu(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
1252 <$ty<T, B> as UnaryResult<T, B>>::wrap(relu_impl(self))
1253 }
1254 }
1255 };
1256}
1257
1258macro_rules! impl_tanh {
1259 ($ty:ident) => {
1260 impl<T, B> $ty<T, B>
1261 where
1262 T: FloatLike,
1263 B: Backend,
1264 {
1265 #[inline]
1276 pub fn tanh(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
1277 <$ty<T, B> as UnaryResult<T, B>>::wrap(tanh_impl(self))
1278 }
1279 }
1280 };
1281}
1282
1283macro_rules! impl_unary_scalar_ops {
1284 ($ty:ident) => {
1285 impl_exp!($ty);
1286 impl_ln!($ty);
1287 impl_log2!($ty);
1288 impl_relu!($ty);
1289 impl_tanh!($ty);
1290 };
1291}
1292
1293macro_rules! impl_op_scalar {
1294 ($ty:ident) => {
1295 impl_add_scalar!($ty);
1296 impl_sub_scalar!($ty);
1297 impl_div_scalar!($ty);
1298 impl_mul_scalar!($ty);
1299 };
1300}
1301
1302macro_rules! impl_add_assign_scalar {
1305 ($ty:ident) => {
1306 impl<T, B> AddAssign<T> for $ty<T, B>
1307 where
1308 T: NumericOp,
1309 B: Backend,
1310 {
1311 #[inline]
1312 fn add_assign(&mut self, rhs: T) {
1313 *self = add_scalar_impl(&*self, rhs);
1314 }
1315 }
1316 };
1317}
1318
1319macro_rules! impl_sub_assign_scalar {
1320 ($ty:ident) => {
1321 impl<T, B> SubAssign<T> for $ty<T, B>
1322 where
1323 T: NumericOp + Neg<Output = T>,
1324 B: Backend,
1325 {
1326 #[inline]
1327 fn sub_assign(&mut self, rhs: T) {
1328 *self = sub_scalar_impl(&*self, rhs);
1329 }
1330 }
1331 };
1332}
1333
1334macro_rules! impl_mul_assign_scalar {
1335 ($ty:ident) => {
1336 impl<T, B> MulAssign<T> for $ty<T, B>
1337 where
1338 T: NumericOp,
1339 B: Backend,
1340 {
1341 #[inline]
1342 fn mul_assign(&mut self, rhs: T) {
1343 *self = mul_scalar_impl(&*self, rhs);
1344 }
1345 }
1346 };
1347}
1348
1349macro_rules! impl_div_assign_scalar {
1350 ($ty:ident) => {
1351 impl<T, B> DivAssign<T> for $ty<T, B>
1352 where
1353 T: NumericOp,
1354 B: Backend,
1355 {
1356 #[inline]
1357 fn div_assign(&mut self, rhs: T) {
1358 *self = div_scalar_impl(&*self, rhs);
1359 }
1360 }
1361 };
1362}
1363
1364macro_rules! impl_op_assign_scalar {
1365 ($ty:ident) => {
1366 impl_add_assign_scalar!($ty);
1367 impl_sub_assign_scalar!($ty);
1368 impl_mul_assign_scalar!($ty);
1369 impl_div_assign_scalar!($ty);
1370 };
1371}
1372
1373macro_rules! impl_tensor_binop {
1376 ($trait:ident, $method:ident, $impl_fn:ident, $lhs:ident, $rhs:ident) => {
1377 impl<T, B> $trait<&$rhs<T, B>> for &$lhs<T, B>
1378 where
1379 T: NumericOp,
1380 B: Backend,
1381 $lhs<T, B>: BinaryResult<$rhs<T, B>, T, B>,
1382 {
1383 type Output = <$lhs<T, B> as BinaryResult<$rhs<T, B>, T, B>>::Output;
1384
1385 #[inline]
1395 fn $method(self, rhs: &$rhs<T, B>) -> Self::Output {
1396 <$lhs<T, B> as BinaryResult<$rhs<T, B>, T, B>>::wrap($impl_fn(self, rhs))
1397 }
1398 }
1399
1400 impl<T, B> $trait<$rhs<T, B>> for &$lhs<T, B>
1401 where
1402 T: NumericOp,
1403 B: Backend,
1404 $lhs<T, B>: BinaryResult<$rhs<T, B>, T, B>,
1405 {
1406 type Output = <$lhs<T, B> as BinaryResult<$rhs<T, B>, T, B>>::Output;
1407
1408 #[inline]
1415 fn $method(self, rhs: $rhs<T, B>) -> Self::Output {
1416 <$lhs<T, B> as BinaryResult<$rhs<T, B>, T, B>>::wrap($impl_fn(self, &rhs))
1417 }
1418 }
1419
1420 impl<T, B> $trait<&$rhs<T, B>> for $lhs<T, B>
1421 where
1422 T: NumericOp,
1423 B: Backend,
1424 $lhs<T, B>: BinaryResult<$rhs<T, B>, T, B>,
1425 {
1426 type Output = <$lhs<T, B> as BinaryResult<$rhs<T, B>, T, B>>::Output;
1427
1428 #[inline]
1435 fn $method(self, rhs: &$rhs<T, B>) -> Self::Output {
1436 <$lhs<T, B> as BinaryResult<$rhs<T, B>, T, B>>::wrap($impl_fn(&self, rhs))
1437 }
1438 }
1439
1440 impl<T, B> $trait<$rhs<T, B>> for $lhs<T, B>
1441 where
1442 T: NumericOp,
1443 B: Backend,
1444 $lhs<T, B>: BinaryResult<$rhs<T, B>, T, B>,
1445 {
1446 type Output = <$lhs<T, B> as BinaryResult<$rhs<T, B>, T, B>>::Output;
1447
1448 #[inline]
1455 fn $method(self, rhs: $rhs<T, B>) -> Self::Output {
1456 <$lhs<T, B> as BinaryResult<$rhs<T, B>, T, B>>::wrap($impl_fn(&self, &rhs))
1457 }
1458 }
1459 };
1460}
1461
1462macro_rules! impl_tensor_ops {
1463 ($lhs:ident, $rhs:ident) => {
1464 impl_tensor_binop!(Add, add, add_tensor_impl, $lhs, $rhs);
1465 impl_tensor_binop!(Sub, sub, sub_tensor_impl, $lhs, $rhs);
1466 impl_tensor_binop!(Mul, mul, mul_tensor_impl, $lhs, $rhs);
1467 impl_tensor_binop!(Div, div, div_tensor_impl, $lhs, $rhs);
1468 };
1469}
1470
1471macro_rules! impl_tensor_ops_cross {
1475 ([$($ty:ident),+ $(,)?]) => {
1476 impl_tensor_ops_cross!(@rows [$($ty),+] [$($ty),+]);
1477 };
1478 (@rows [$($lhs:ident),+] $rhs:tt) => {
1479 $( impl_tensor_ops_cross!(@row $lhs $rhs); )+
1480 };
1481 (@row $lhs:ident [$($rhs:ident),+]) => {
1482 $( impl_tensor_ops!($lhs, $rhs); )+
1483 };
1484}
1485
1486macro_rules! impl_matmul {
1489 ($ty:ident) => {
1490 impl<T, B> $ty<T, B>
1491 where
1492 T: CanMatMul,
1493 B: Backend,
1494 {
1495 #[inline]
1541 pub fn matmul<D>(
1542 &self,
1543 rhs: &D,
1544 ) -> Result<<$ty<T, B> as BinaryResult<D, T, B>>::Output, OpError>
1545 where
1546 D: Operand<T, B>,
1547 $ty<T, B>: BinaryResult<D, T, B>,
1548 {
1549 matmul_tensor_impl(self, rhs).map(<$ty<T, B> as BinaryResult<D, T, B>>::wrap)
1550 }
1551 }
1552 };
1553}
1554
1555macro_rules! impl_sum {
1558 ($ty:ident) => {
1559 impl<T, B> $ty<T, B>
1560 where
1561 T: NumericOp,
1562 B: Backend,
1563 {
1564 #[inline]
1578 pub fn sum(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
1579 <$ty<T, B> as UnaryResult<T, B>>::wrap(sum_impl(self))
1580 }
1581 }
1582 };
1583}
1584
1585macro_rules! impl_sum_axis {
1586 ($ty:ident) => {
1587 impl<T, B> $ty<T, B>
1588 where
1589 T: NumericOp,
1590 B: Backend,
1591 {
1592 #[inline]
1620 pub fn sum_axis(
1621 &self,
1622 axis: isize,
1623 keep_dims: bool,
1624 ) -> Result<<$ty<T, B> as UnaryResult<T, B>>::Output, OpError> {
1625 sum_axis_impl(self, axis, keep_dims).map(<$ty<T, B> as UnaryResult<T, B>>::wrap)
1626 }
1627 }
1628 };
1629}
1630
1631macro_rules! impl_max {
1632 ($ty:ident) => {
1633 impl<T, B> $ty<T, B>
1634 where
1635 T: NumericOp,
1636 B: Backend,
1637 {
1638 #[inline]
1652 pub fn max(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
1653 <$ty<T, B> as UnaryResult<T, B>>::wrap(max_impl(self))
1654 }
1655 }
1656 };
1657}
1658
1659macro_rules! impl_max_axis {
1660 ($ty:ident) => {
1661 impl<T, B> $ty<T, B>
1662 where
1663 T: NumericOp,
1664 B: Backend,
1665 {
1666 #[inline]
1688 pub fn max_axis(
1689 &self,
1690 axis: isize,
1691 keep_dims: bool,
1692 ) -> Result<<$ty<T, B> as UnaryResult<T, B>>::Output, OpError> {
1693 max_axis_impl(self, axis, keep_dims).map(<$ty<T, B> as UnaryResult<T, B>>::wrap)
1694 }
1695 }
1696 };
1697}
1698
1699macro_rules! impl_mean {
1700 ($ty:ident) => {
1701 impl<T, B> $ty<T, B>
1702 where
1703 T: FloatLike,
1704 B: Backend,
1705 {
1706 #[inline]
1720 pub fn mean(&self) -> <$ty<T, B> as UnaryResult<T, B>>::Output {
1721 <$ty<T, B> as UnaryResult<T, B>>::wrap(mean_impl(self))
1722 }
1723 }
1724 };
1725}
1726
1727macro_rules! impl_mean_axis {
1728 ($ty:ident) => {
1729 impl<T, B> $ty<T, B>
1730 where
1731 T: FloatLike,
1732 B: Backend,
1733 {
1734 #[inline]
1756 pub fn mean_axis(
1757 &self,
1758 axis: isize,
1759 keep_dims: bool,
1760 ) -> Result<<$ty<T, B> as UnaryResult<T, B>>::Output, OpError> {
1761 mean_axis_impl(self, axis, keep_dims).map(<$ty<T, B> as UnaryResult<T, B>>::wrap)
1762 }
1763 }
1764 };
1765}
1766
1767macro_rules! impl_tensor_assign_binop {
1770 ($trait:ident, $method:ident, $impl_fn:ident, $rhs:ident) => {
1771 impl<T, B> $trait<$rhs<T, B>> for TensorPromise<T, B>
1772 where
1773 T: NumericOp,
1774 B: Backend,
1775 {
1776 #[inline]
1777 fn $method(&mut self, rhs: $rhs<T, B>) {
1778 *self = $impl_fn(&*self, &rhs);
1779 }
1780 }
1781
1782 impl<T, B> $trait<&$rhs<T, B>> for TensorPromise<T, B>
1783 where
1784 T: NumericOp,
1785 B: Backend,
1786 {
1787 #[inline]
1788 fn $method(&mut self, rhs: &$rhs<T, B>) {
1789 *self = $impl_fn(&*self, rhs);
1790 }
1791 }
1792 };
1793}
1794
1795macro_rules! impl_tensor_assign_ops {
1796 ($rhs:ident) => {
1797 impl_tensor_assign_binop!(AddAssign, add_assign, add_tensor_impl, $rhs);
1798 impl_tensor_assign_binop!(SubAssign, sub_assign, sub_tensor_impl, $rhs);
1799 impl_tensor_assign_binop!(MulAssign, mul_assign, mul_tensor_impl, $rhs);
1800 impl_tensor_assign_binop!(DivAssign, div_assign, div_tensor_impl, $rhs);
1801 };
1802}
1803
1804macro_rules! impl_all_ops {
1807 ($ty:ident) => {
1808 impl_reshape_like!($ty);
1809 impl_unary_scalar_ops!($ty);
1810 impl_op_scalar!($ty);
1811 impl_matmul!($ty);
1812 impl_sum!($ty);
1813 impl_sum_axis!($ty);
1814 impl_max!($ty);
1815 impl_max_axis!($ty);
1816 impl_mean!($ty);
1817 impl_mean_axis!($ty);
1818 };
1819}
1820
1821impl_all_ops!(Tensor);
1822impl_all_ops!(TensorPromise);
1823impl_all_ops!(CachedTensorPromise);
1824impl_all_ops!(BakedPromise);
1825impl_all_ops!(SkeletonSlot);
1826impl_all_ops!(SkeletonPromise);
1827
1828impl_tensor_ops_cross!([
1829 Tensor,
1830 TensorPromise,
1831 CachedTensorPromise,
1832 BakedPromise,
1833 SkeletonSlot,
1834 SkeletonPromise,
1835]);
1836
1837impl_op_assign_scalar!(TensorPromise);
1838
1839impl_tensor_assign_ops!(Tensor);
1840impl_tensor_assign_ops!(TensorPromise);
1841impl_tensor_assign_ops!(CachedTensorPromise);
1842impl_tensor_assign_ops!(BakedPromise);