1use super::{conv, ctc, embedding, interpolation, linear, pool};
2use crate::ops::unfold::unfold4d_using_conv2d;
3use crate::tensor::{BoolTensor, FloatTensor, IntTensor};
4use crate::{Backend, ElementConversion, TensorMetadata};
5use ruda_core::tensor::Shape;
6
7pub struct SoftmaxOutput<B: Backend> {
9 pub output: FloatTensor<B>,
11 pub working: FloatTensor<B>,
13}
14
15pub struct LayerNormOutput<B: Backend> {
17 pub output: FloatTensor<B>,
19 pub mean: FloatTensor<B>,
21 pub rstd: FloatTensor<B>,
23}
24
25pub struct LayerNormBackward<B: Backend> {
27 pub input: FloatTensor<B>,
29 pub weight: FloatTensor<B>,
31 pub bias: FloatTensor<B>,
33}
34
35pub struct RmsNormOutput<B: Backend> {
37 pub output: FloatTensor<B>,
39 pub rstd: FloatTensor<B>,
41}
42
43pub struct RmsNormBackward<B: Backend> {
45 pub input: FloatTensor<B>,
47 pub weight: FloatTensor<B>,
49}
50
51#[derive(new)]
53pub struct Conv2dBackward<B: Backend> {
54 pub x_grad: FloatTensor<B>,
56
57 pub weights_grad: FloatTensor<B>,
59
60 pub bias_grad: Option<FloatTensor<B>>,
62}
63
64#[derive(new)]
66pub struct DeformConv2dBackward<B: Backend> {
67 pub x_grad: FloatTensor<B>,
69
70 pub offset_grad: FloatTensor<B>,
72
73 pub weight_grad: FloatTensor<B>,
75
76 pub mask_grad: Option<FloatTensor<B>>,
78
79 pub bias_grad: Option<FloatTensor<B>>,
81}
82
83#[derive(new)]
85pub struct Conv3dBackward<B: Backend> {
86 pub x_grad: FloatTensor<B>,
88
89 pub weights_grad: FloatTensor<B>,
91
92 pub bias_grad: Option<FloatTensor<B>>,
94}
95
96#[derive(new)]
98pub struct MaxPool1dBackward<B: Backend> {
99 pub x_grad: FloatTensor<B>,
101}
102
103#[derive(new)]
105pub struct MaxPool1dWithIndices<B: Backend> {
106 pub output: FloatTensor<B>,
108
109 pub indices: IntTensor<B>,
111}
112
113#[derive(new)]
115pub struct MaxPool2dBackward<B: Backend> {
116 pub x_grad: FloatTensor<B>,
118}
119
120#[derive(new)]
122pub struct MaxPool2dWithIndices<B: Backend> {
123 pub output: FloatTensor<B>,
125
126 pub indices: IntTensor<B>,
128}
129
130#[derive(new)]
132pub struct MaxPool3dWithIndices<B: Backend> {
133 pub output: FloatTensor<B>,
135 pub indices: IntTensor<B>,
137}
138
139#[derive(new)]
141pub struct MaxPool3dBackward<B: Backend> {
142 pub x_grad: FloatTensor<B>,
144}
145
146pub use ruda_core::tensor::spatial::{ConvOptions, PaddedConvOptions, DeformConvOptions, ConvTransposeOptions, UnfoldOptions};
147
148pub use ruda_core::tensor::spatial::{InterpolateMode, InterpolateOptions};
149
150pub use ruda_core::tensor::spatial::{GridSampleOptions, GridSamplePaddingMode};
151
152#[derive(Debug, Clone, Copy, PartialEq, serde::Deserialize, serde::Serialize)]
163pub enum PadMode {
164 Constant(f32),
170
171 Reflect,
179
180 Edge,
186}
187
188impl Default for PadMode {
189 fn default() -> Self {
190 PadMode::Constant(0.0)
191 }
192}
193
194impl<E: ElementConversion> From<E> for PadMode {
195 fn from(value: E) -> Self {
196 PadMode::Constant(value.elem())
197 }
198}
199
200#[derive(new)]
202pub struct InterpolateBackward<B: Backend> {
203 pub x_grad: FloatTensor<B>,
205}
206
207pub use ruda_core::tensor::spatial::AttentionModuleOptions;
208
209pub trait ModuleOps<B: Backend> {
211 fn exponential_relu_native(tensor: FloatTensor<B>, alpha: f64, continuous: bool) -> FloatTensor<B> {
213 super::activation_training::exponential_relu_native::<B>(tensor, alpha, continuous)
214 }
215
216 fn exponential_relu_native_backward(tensor: FloatTensor<B>, grad: FloatTensor<B>, alpha: f64, continuous: bool) -> FloatTensor<B> {
218 super::activation_training::exponential_relu_native_backward::<B>(tensor, grad, alpha, continuous)
219 }
220
221 fn leaky_relu_native(tensor: FloatTensor<B>, negative_slope: f64) -> FloatTensor<B> {
223 super::activation_training::leaky_relu_native::<B>(tensor, negative_slope)
224 }
225
226 fn leaky_relu_native_backward(tensor: FloatTensor<B>, grad: FloatTensor<B>, negative_slope: f64) -> FloatTensor<B> {
228 super::activation_training::leaky_relu_native_backward::<B>(tensor, grad, negative_slope)
229 }
230
231 fn prelu_native(tensor: FloatTensor<B>, alpha: FloatTensor<B>) -> FloatTensor<B> {
233 super::prelu_training::prelu_native::<B>(tensor, alpha)
234 }
235
236 fn prelu_native_backward_select(tensor: FloatTensor<B>, alpha: FloatTensor<B>, grad: FloatTensor<B>,
238 mask: [bool; 2]) -> [Option<FloatTensor<B>>; 2] {
239 super::prelu_training::prelu_native_backward_select::<B>(tensor, alpha, grad, mask)
240 }
241
242 fn group_norm(tensor: FloatTensor<B>, gamma: Option<FloatTensor<B>>, beta: Option<FloatTensor<B>>,
244 groups: usize, epsilon: f64) -> FloatTensor<B> {
245 Self::group_norm_with_stats(tensor, gamma, beta, groups, epsilon).output
246 }
247
248 fn group_norm_with_stats(tensor: FloatTensor<B>, gamma: Option<FloatTensor<B>>, beta: Option<FloatTensor<B>>,
250 groups: usize, epsilon: f64) -> LayerNormOutput<B> {
251 super::group_normalization::group_norm_with_stats::<B>(tensor, gamma, beta, groups, epsilon)
252 }
253
254 fn group_norm_backward_select(tensor: FloatTensor<B>, gamma: Option<FloatTensor<B>>, grad: FloatTensor<B>,
256 mean: FloatTensor<B>, rstd: FloatTensor<B>, groups: usize, mask: [bool; 3]) -> [Option<FloatTensor<B>>; 3] {
257 super::group_normalization::group_norm_backward_select::<B>(tensor, gamma, grad, mean, rstd, groups, mask)
258 }
259
260 fn gelu_native(tensor: FloatTensor<B>, approximate: bool) -> FloatTensor<B> {
262 super::activation_training::gelu_native::<B>(tensor, approximate)
263 }
264
265 fn gelu_native_backward(input: FloatTensor<B>, grad: FloatTensor<B>, approximate: bool) -> FloatTensor<B> {
267 super::activation_training::gelu_native_backward::<B>(input, grad, approximate)
268 }
269
270 fn silu_native(tensor: FloatTensor<B>) -> FloatTensor<B> {
272 super::activation_training::silu_native::<B>(tensor)
273 }
274
275 fn silu_native_backward(input: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
277 super::activation_training::silu_native_backward::<B>(input, grad)
278 }
279
280 fn softmax_native(tensor: FloatTensor<B>, dim: usize, logarithmic: bool) -> FloatTensor<B> {
282 Self::softmax_with_stats(tensor, dim, logarithmic).output
283 }
284
285 fn softmax_with_stats(tensor: FloatTensor<B>, dim: usize, logarithmic: bool) -> SoftmaxOutput<B> {
287 super::softmax::softmax_with_stats::<B>(tensor, dim, logarithmic)
288 }
289
290 fn softmax_native_backward(working: FloatTensor<B>, grad: FloatTensor<B>, dim: usize,
292 logarithmic: bool) -> FloatTensor<B> {
293 super::softmax::softmax_backward::<B>(working, grad, dim, logarithmic)
294 }
295
296 fn embedding(weights: FloatTensor<B>, indices: IntTensor<B>) -> FloatTensor<B> {
307 embedding::embedding::<B>(weights, indices)
308 }
309
310 fn embedding_backward(
322 weights: FloatTensor<B>,
323 output_grad: FloatTensor<B>,
324 indices: IntTensor<B>,
325 ) -> FloatTensor<B> {
326 embedding::embedding_backward::<B>(weights, output_grad, indices)
327 }
328
329 fn linear(
337 x: FloatTensor<B>,
338 weight: FloatTensor<B>,
339 bias: Option<FloatTensor<B>>,
340 ) -> FloatTensor<B> {
341 linear::linear::<B>(x, weight, bias)
342 }
343 fn linear_x_backward(weight: FloatTensor<B>, output_grad: FloatTensor<B>) -> FloatTensor<B> {
345 linear::linear_x_backward::<B>(weight, output_grad)
346 }
347 fn linear_weight_backward(x: FloatTensor<B>, output_grad: FloatTensor<B>) -> FloatTensor<B> {
349 linear::linear_weight_backward::<B>(x, output_grad)
350 }
351 fn linear_bias_backward(output_grad: FloatTensor<B>) -> FloatTensor<B> {
353 linear::linear_bias_backward::<B>(output_grad)
354 }
355
356 fn conv1d(
364 x: FloatTensor<B>,
365 weight: FloatTensor<B>,
366 bias: Option<FloatTensor<B>>,
367 options: ConvOptions<1>,
368 ) -> FloatTensor<B> {
369 conv::conv1d_from_conv2d::<B>(x, weight, bias, options)
370 }
371 fn conv1d_x_backward(
373 x: FloatTensor<B>,
374 weight: FloatTensor<B>,
375 output_grad: FloatTensor<B>,
376 options: ConvOptions<1>,
377 ) -> FloatTensor<B> {
378 conv::conv1d_x_backward::<B>(x, weight, output_grad, options)
379 }
380 fn conv1d_weight_backward(
382 x: FloatTensor<B>,
383 weight: FloatTensor<B>,
384 output_grad: FloatTensor<B>,
385 options: ConvOptions<1>,
386 ) -> FloatTensor<B> {
387 conv::conv1d_weight_backward::<B>(x, weight, output_grad, options)
388 }
389 fn conv1d_bias_backward(
391 x: FloatTensor<B>,
392 bias: FloatTensor<B>,
393 output_grad: FloatTensor<B>,
394 ) -> FloatTensor<B> {
395 conv::conv1d_bias_backward::<B>(x, bias, output_grad)
396 }
397 fn conv2d(
405 x: FloatTensor<B>,
406 weight: FloatTensor<B>,
407 bias: Option<FloatTensor<B>>,
408 options: ConvOptions<2>,
409 ) -> FloatTensor<B>;
410 fn conv2d_x_backward(
412 x: FloatTensor<B>,
413 weight: FloatTensor<B>,
414 output_grad: FloatTensor<B>,
415 options: ConvOptions<2>,
416 ) -> FloatTensor<B> {
417 conv::conv2d_x_backward::<B>(x, weight, output_grad, options)
418 }
419 fn conv2d_weight_backward(
421 x: FloatTensor<B>,
422 weight: FloatTensor<B>,
423 output_grad: FloatTensor<B>,
424 options: ConvOptions<2>,
425 ) -> FloatTensor<B> {
426 conv::conv2d_weight_backward::<B>(x, weight, output_grad, options)
427 }
428 fn conv2d_bias_backward(
430 x: FloatTensor<B>,
431 bias: FloatTensor<B>,
432 output_grad: FloatTensor<B>,
433 ) -> FloatTensor<B> {
434 conv::conv2d_bias_backward::<B>(x, bias, output_grad)
435 }
436
437 fn deform_conv2d(
445 x: FloatTensor<B>,
446 offset: FloatTensor<B>,
447 weight: FloatTensor<B>,
448 mask: Option<FloatTensor<B>>,
449 bias: Option<FloatTensor<B>>,
450 options: DeformConvOptions<2>,
451 ) -> FloatTensor<B>;
452 fn deform_conv2d_backward(
454 x: FloatTensor<B>,
455 offset: FloatTensor<B>,
456 weight: FloatTensor<B>,
457 mask: Option<FloatTensor<B>>,
458 bias: Option<FloatTensor<B>>,
459 output_grad: FloatTensor<B>,
460 options: DeformConvOptions<2>,
461 ) -> DeformConv2dBackward<B>;
462
463 fn conv3d(
471 x: FloatTensor<B>,
472 weight: FloatTensor<B>,
473 bias: Option<FloatTensor<B>>,
474 options: ConvOptions<3>,
475 ) -> FloatTensor<B>;
476 fn conv3d_x_backward(
478 x: FloatTensor<B>,
479 weight: FloatTensor<B>,
480 output_grad: FloatTensor<B>,
481 options: ConvOptions<3>,
482 ) -> FloatTensor<B> {
483 conv::conv3d_x_backward::<B>(x, weight, output_grad, options)
484 }
485 fn conv3d_weight_backward(
487 x: FloatTensor<B>,
488 weight: FloatTensor<B>,
489 output_grad: FloatTensor<B>,
490 options: ConvOptions<3>,
491 ) -> FloatTensor<B> {
492 conv::conv3d_weight_backward::<B>(x, weight, output_grad, options)
493 }
494 fn conv3d_bias_backward(
496 x: FloatTensor<B>,
497 bias: FloatTensor<B>,
498 output_grad: FloatTensor<B>,
499 ) -> FloatTensor<B> {
500 conv::conv3d_bias_backward::<B>(x, bias, output_grad)
501 }
502 fn conv_transpose1d(
510 x: FloatTensor<B>,
511 weight: FloatTensor<B>,
512 bias: Option<FloatTensor<B>>,
513 options: ConvTransposeOptions<1>,
514 ) -> FloatTensor<B> {
515 conv::conv_transpose1d_from_conv_transpose2d::<B>(x, weight, bias, options)
516 }
517 fn conv_transpose1d_x_backward(
519 weight: FloatTensor<B>,
520 output_grad: FloatTensor<B>,
521 options: ConvTransposeOptions<1>,
522 ) -> FloatTensor<B> {
523 conv::conv_transpose1d_x_backward::<B>(weight, output_grad, options)
524 }
525 fn conv_transpose1d_weight_backward(
527 x: FloatTensor<B>,
528 weight: FloatTensor<B>,
529 output_grad: FloatTensor<B>,
530 options: ConvTransposeOptions<1>,
531 ) -> FloatTensor<B> {
532 conv::conv_transpose1d_weight_backward::<B>(x, weight, output_grad, options)
533 }
534 fn conv_transpose1d_bias_backward(
536 x: FloatTensor<B>,
537 bias: FloatTensor<B>,
538 output_grad: FloatTensor<B>,
539 ) -> FloatTensor<B> {
540 conv::conv_transpose1d_bias_backward::<B>(x, bias, output_grad)
541 }
542
543 fn conv_transpose2d(
551 x: FloatTensor<B>,
552 weight: FloatTensor<B>,
553 bias: Option<FloatTensor<B>>,
554 options: ConvTransposeOptions<2>,
555 ) -> FloatTensor<B>;
556 fn conv_transpose2d_x_backward(
558 weight: FloatTensor<B>,
559 output_grad: FloatTensor<B>,
560 options: ConvTransposeOptions<2>,
561 ) -> FloatTensor<B> {
562 conv::conv_transpose2d_x_backward::<B>(weight, output_grad, options)
563 }
564 fn conv_transpose2d_weight_backward(
566 x: FloatTensor<B>,
567 weight: FloatTensor<B>,
568 output_grad: FloatTensor<B>,
569 options: ConvTransposeOptions<2>,
570 ) -> FloatTensor<B> {
571 conv::conv_transpose2d_weight_backward::<B>(x, weight, output_grad, options)
572 }
573 fn conv_transpose2d_bias_backward(
575 x: FloatTensor<B>,
576 bias: FloatTensor<B>,
577 output_grad: FloatTensor<B>,
578 ) -> FloatTensor<B> {
579 conv::conv_transpose2d_bias_backward::<B>(x, bias, output_grad)
580 }
581
582 fn conv_transpose3d(
590 x: FloatTensor<B>,
591 weight: FloatTensor<B>,
592 bias: Option<FloatTensor<B>>,
593 options: ConvTransposeOptions<3>,
594 ) -> FloatTensor<B>;
595 fn conv_transpose3d_x_backward(
597 weight: FloatTensor<B>,
598 output_grad: FloatTensor<B>,
599 options: ConvTransposeOptions<3>,
600 ) -> FloatTensor<B> {
601 conv::conv_transpose3d_x_backward::<B>(weight, output_grad, options)
602 }
603 fn conv_transpose3d_weight_backward(
605 x: FloatTensor<B>,
606 weight: FloatTensor<B>,
607 output_grad: FloatTensor<B>,
608 options: ConvTransposeOptions<3>,
609 ) -> FloatTensor<B> {
610 conv::conv_transpose3d_weight_backward::<B>(x, weight, output_grad, options)
611 }
612 fn conv_transpose3d_bias_backward(
614 x: FloatTensor<B>,
615 bias: FloatTensor<B>,
616 output_grad: FloatTensor<B>,
617 ) -> FloatTensor<B> {
618 conv::conv_transpose3d_bias_backward::<B>(x, bias, output_grad)
619 }
620
621 fn unfold4d(
628 x: FloatTensor<B>,
629 kernel_size: [usize; 2],
630 options: UnfoldOptions,
631 ) -> FloatTensor<B> {
632 if options.padding == [0, 0] && options.dilation == [1, 1] {
633 let blocks = B::float_unfold(x, 2, kernel_size[0], options.stride[0]);
634 let blocks = B::float_unfold(blocks, 3, kernel_size[1], options.stride[1]);
635
636 let blocks = B::float_permute(blocks, &[0, 1, 4, 5, 2, 3]);
639 let shape = blocks.shape();
640
641 B::float_reshape(
644 blocks,
645 [
646 shape[0],
647 shape[1] * shape[2] * shape[3],
648 shape[4] * shape[5],
649 ]
650 .into(),
651 )
652 } else {
653 unfold4d_using_conv2d::<B>(x, kernel_size, options)
654 }
655 }
656
657 fn avg_pool1d(
663 x: FloatTensor<B>,
664 kernel_size: usize,
665 stride: usize,
666 padding: usize,
667 count_include_pad: bool,
668 ceil_mode: bool,
669 ) -> FloatTensor<B> {
670 pool::avg_pool1d_from_2d::<B>(
671 x,
672 kernel_size,
673 stride,
674 padding,
675 count_include_pad,
676 ceil_mode,
677 )
678 }
679 fn avg_pool1d_backward(
681 x: FloatTensor<B>,
682 grad: FloatTensor<B>,
683 kernel_size: usize,
684 stride: usize,
685 padding: usize,
686 count_include_pad: bool,
687 ceil_mode: bool,
688 ) -> FloatTensor<B> {
689 pool::avg_pool1d_backward_from_2d::<B>(
690 x,
691 grad,
692 kernel_size,
693 stride,
694 padding,
695 count_include_pad,
696 ceil_mode,
697 )
698 }
699 fn avg_pool2d(
705 x: FloatTensor<B>,
706 kernel_size: [usize; 2],
707 stride: [usize; 2],
708 padding: [usize; 2],
709 count_include_pad: bool,
710 ceil_mode: bool,
711 ) -> FloatTensor<B>;
712 fn avg_pool2d_backward(
714 x: FloatTensor<B>,
715 grad: FloatTensor<B>,
716 kernel_size: [usize; 2],
717 stride: [usize; 2],
718 padding: [usize; 2],
719 count_include_pad: bool,
720 ceil_mode: bool,
721 ) -> FloatTensor<B>;
722 fn avg_pool3d(x: FloatTensor<B>, kernel_size: [usize; 3], stride: [usize; 3],
724 padding: [usize; 3], count_include_pad: bool, ceil_mode: bool) -> FloatTensor<B> {
725 pool::avg_pool3d_from_2d::<B>(x, kernel_size, stride, padding, count_include_pad, ceil_mode)
726 }
727 fn avg_pool3d_native_output_size(_input: [usize; 3], _kernel: [usize; 3],
729 _stride: [usize; 3], _padding: [usize; 3], _ceil: bool) -> Option<[usize; 3]> {
730 None
731 }
732 fn avg_pool3d_backward(x: FloatTensor<B>, grad: FloatTensor<B>, kernel_size: [usize; 3],
734 stride: [usize; 3], padding: [usize; 3], count_include_pad: bool, ceil_mode: bool) -> FloatTensor<B> {
735 pool::avg_pool3d_backward_from_2d::<B>(x, grad, kernel_size, stride, padding, count_include_pad, ceil_mode)
736 }
737 fn adaptive_avg_pool2d(x: FloatTensor<B>, output_size: [usize; 2]) -> FloatTensor<B>;
743 fn adaptive_avg_pool2d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B>;
745 fn adaptive_avg_pool3d(x: FloatTensor<B>, output_size: [usize; 3]) -> FloatTensor<B> {
747 pool::adaptive_avg_pool3d_from_2d::<B>(x, output_size)
748 }
749 fn adaptive_avg_pool3d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
751 pool::adaptive_avg_pool3d_backward_from_2d::<B>(x, grad)
752 }
753 fn max_pool3d(x: FloatTensor<B>, kernel: [usize; 3], stride: [usize; 3], padding: [usize; 3],
755 dilation: [usize; 3], ceil: bool) -> FloatTensor<B> {
756 pool::max_pool3d_from_2d::<B>(x, kernel, stride, padding, dilation, ceil)
757 }
758 fn max_pool3d_with_indices(x: FloatTensor<B>, kernel: [usize; 3], stride: [usize; 3],
760 padding: [usize; 3], dilation: [usize; 3], ceil: bool) -> MaxPool3dWithIndices<B> {
761 pool::max_pool3d_with_indices_from_2d::<B>(x, kernel, stride, padding, dilation, ceil)
762 }
763 fn max_pool3d_with_indices_backward(x: FloatTensor<B>, grad: FloatTensor<B>, indices: IntTensor<B>,
765 _kernel: [usize; 3], _stride: [usize; 3], _padding: [usize; 3], _dilation: [usize; 3],
766 _ceil: bool) -> MaxPool3dBackward<B> {
767 MaxPool3dBackward::new(pool::max_pool3d_backward_from_indices::<B>(x, grad, indices))
768 }
769 fn adaptive_avg_pool1d(x: FloatTensor<B>, output_size: usize) -> FloatTensor<B> {
775 pool::adaptive_avg_pool1d_from_2d::<B>(x, output_size)
776 }
777 fn adaptive_avg_pool1d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
779 pool::adaptive_avg_pool1d_backward_from_2d::<B>(x, grad)
780 }
781 fn max_pool1d(
787 x: FloatTensor<B>,
788 kernel_size: usize,
789 stride: usize,
790 padding: usize,
791 dilation: usize,
792 ceil_mode: bool,
793 ) -> FloatTensor<B> {
794 pool::max_pool1d_from_2d::<B>(x, kernel_size, stride, padding, dilation, ceil_mode)
795 }
796
797 fn max_pool1d_with_indices(
803 x: FloatTensor<B>,
804 kernel_size: usize,
805 stride: usize,
806 padding: usize,
807 dilation: usize,
808 ceil_mode: bool,
809 ) -> MaxPool1dWithIndices<B> {
810 pool::max_pool1d_with_indices_from_2d::<B>(
811 x,
812 kernel_size,
813 stride,
814 padding,
815 dilation,
816 ceil_mode,
817 )
818 }
819 #[allow(clippy::too_many_arguments)]
821 fn max_pool1d_with_indices_backward(
822 x: FloatTensor<B>,
823 kernel_size: usize,
824 stride: usize,
825 padding: usize,
826 dilation: usize,
827 ceil_mode: bool,
828 output_grad: FloatTensor<B>,
829 indices: IntTensor<B>,
830 ) -> MaxPool1dBackward<B> {
831 pool::max_pool1d_with_indices_backward_from_2d::<B>(
832 x,
833 kernel_size,
834 stride,
835 padding,
836 dilation,
837 ceil_mode,
838 output_grad,
839 indices,
840 )
841 }
842
843 fn max_pool2d(
849 x: FloatTensor<B>,
850 kernel_size: [usize; 2],
851 stride: [usize; 2],
852 padding: [usize; 2],
853 dilation: [usize; 2],
854 ceil_mode: bool,
855 ) -> FloatTensor<B>;
856
857 fn max_pool2d_with_indices(
863 x: FloatTensor<B>,
864 kernel_size: [usize; 2],
865 stride: [usize; 2],
866 padding: [usize; 2],
867 dilation: [usize; 2],
868 ceil_mode: bool,
869 ) -> MaxPool2dWithIndices<B>;
870 #[allow(clippy::too_many_arguments)]
872 fn max_pool2d_with_indices_backward(
873 x: FloatTensor<B>,
874 kernel_size: [usize; 2],
875 stride: [usize; 2],
876 padding: [usize; 2],
877 dilation: [usize; 2],
878 ceil_mode: bool,
879 output_grad: FloatTensor<B>,
880 indices: IntTensor<B>,
881 ) -> MaxPool2dBackward<B>;
882
883 fn interpolate(
889 x: FloatTensor<B>,
890 output_size: [usize; 2],
891 options: InterpolateOptions,
892 ) -> FloatTensor<B>;
893
894 fn interpolate_backward(
896 x: FloatTensor<B>,
897 grad: FloatTensor<B>,
898 output_size: [usize; 2],
899 options: InterpolateOptions,
900 ) -> FloatTensor<B>;
901
902 fn interpolate1d(x: FloatTensor<B>, size: usize, options: InterpolateOptions) -> FloatTensor<B> {
904 interpolation::interpolate1d_from_2d::<B>(x, size, options)
905 }
906
907 fn interpolate1d_backward(x: FloatTensor<B>, grad: FloatTensor<B>, size: usize,
909 options: InterpolateOptions) -> FloatTensor<B> {
910 interpolation::interpolate1d_backward_from_2d::<B>(x, grad, size, options)
911 }
912
913 fn interpolate3d(x: FloatTensor<B>, size: [usize; 3], options: InterpolateOptions) -> FloatTensor<B> {
915 interpolation::interpolate3d_from_2d::<B>(x, size, options)
916 }
917
918 fn interpolate3d_backward(x: FloatTensor<B>, grad: FloatTensor<B>, size: [usize; 3],
920 options: InterpolateOptions) -> FloatTensor<B> {
921 interpolation::interpolate3d_backward_from_2d::<B>(x, grad, size, options)
922 }
923
924 fn attention(
946 query: FloatTensor<B>,
947 key: FloatTensor<B>,
948 value: FloatTensor<B>,
949 mask: Option<BoolTensor<B>>,
950 attn_bias: Option<FloatTensor<B>>,
951 options: AttentionModuleOptions,
952 ) -> FloatTensor<B>;
953
954 fn layer_norm(
970 tensor: FloatTensor<B>,
971 gamma: FloatTensor<B>,
972 beta: Option<FloatTensor<B>>,
973 epsilon: f64,
974 ) -> FloatTensor<B> {
975 Self::layer_norm_default(tensor, gamma, beta, epsilon)
976 }
977
978 fn has_layer_norm_backward() -> bool { false }
980
981 fn has_rms_norm_backward() -> bool { false }
983
984 fn rms_norm(tensor: FloatTensor<B>, gamma: FloatTensor<B>, epsilon: f64) -> FloatTensor<B> {
986 Self::rms_norm_with_stats(tensor, gamma, epsilon).output
987 }
988
989 fn rms_norm_with_stats(tensor: FloatTensor<B>, gamma: FloatTensor<B>, epsilon: f64) -> RmsNormOutput<B> {
991 super::normalization::rms_norm_with_stats::<B>(tensor, gamma, epsilon)
992 }
993
994 fn rms_norm_backward(tensor: FloatTensor<B>, gamma: FloatTensor<B>, grad: FloatTensor<B>,
996 rstd: FloatTensor<B>) -> RmsNormBackward<B> {
997 super::normalization::rms_norm_backward::<B>(tensor, gamma, grad, rstd)
998 }
999
1000 fn rms_norm_backward_select(tensor: FloatTensor<B>, gamma: FloatTensor<B>, grad: FloatTensor<B>,
1002 rstd: FloatTensor<B>, mask: [bool; 2]) -> [Option<FloatTensor<B>>; 2] {
1003 super::normalization::rms_norm_backward_select::<B>(tensor, gamma, grad, rstd, mask)
1004 }
1005
1006 fn layer_norm_with_stats(
1008 tensor: FloatTensor<B>, gamma: FloatTensor<B>,
1009 beta: Option<FloatTensor<B>>, epsilon: f64,
1010 ) -> LayerNormOutput<B> {
1011 super::normalization::layer_norm_with_stats::<B>(tensor, gamma, beta, epsilon)
1012 }
1013
1014 fn layer_norm_backward(
1016 tensor: FloatTensor<B>, gamma: FloatTensor<B>, grad: FloatTensor<B>,
1017 mean: FloatTensor<B>, rstd: FloatTensor<B>,
1018 ) -> LayerNormBackward<B> {
1019 super::normalization::layer_norm_backward::<B>(tensor, gamma, grad, mean, rstd)
1020 }
1021
1022 fn layer_norm_backward_select(tensor: FloatTensor<B>, gamma: FloatTensor<B>, grad: FloatTensor<B>,
1024 mean: FloatTensor<B>, rstd: FloatTensor<B>, mask: [bool; 3]) -> [Option<FloatTensor<B>>; 3] {
1025 super::normalization::layer_norm_backward_select::<B>(tensor, gamma, grad, mean, rstd, mask)
1026 }
1027
1028 fn layer_norm_default(
1030 tensor: FloatTensor<B>, gamma: FloatTensor<B>,
1031 beta: Option<FloatTensor<B>>, epsilon: f64,
1032 ) -> FloatTensor<B> {
1033 let shape = tensor.shape();
1034 let rank = shape.num_dims();
1035 let last_dim = rank - 1;
1036 let d_model = shape[last_dim];
1037
1038 let mean = B::float_mean_dim(tensor.clone(), last_dim);
1039 let centered = B::float_sub(tensor, mean);
1040 let var = B::float_mean_dim(B::float_mul(centered.clone(), centered.clone()), last_dim);
1041 let denom = B::float_sqrt(B::float_add_scalar(var, epsilon.into()));
1042 let normalized = B::float_div(centered, denom);
1043
1044 let broadcast_dims: alloc::vec::Vec<usize> = (0..rank)
1045 .map(|i| if i == last_dim { d_model } else { 1 })
1046 .collect();
1047 let gamma_b = B::float_reshape(gamma, Shape::from(broadcast_dims.clone()));
1048 let scaled = B::float_mul(normalized, gamma_b);
1049
1050 match beta {
1051 Some(beta) => {
1052 let beta_b = B::float_reshape(beta, Shape::from(broadcast_dims));
1053 B::float_add(scaled, beta_b)
1054 }
1055 None => scaled,
1056 }
1057 }
1058
1059 fn ctc_loss(
1076 log_probs: FloatTensor<B>,
1077 targets: IntTensor<B>,
1078 input_lengths: IntTensor<B>,
1079 target_lengths: IntTensor<B>,
1080 blank: usize,
1081 ) -> FloatTensor<B> {
1082 ctc::ctc_loss_default::<B>(log_probs, targets, input_lengths, target_lengths, blank)
1083 }
1084
1085 fn has_ctc_loss_backward() -> bool {
1097 false
1098 }
1099
1100 fn ctc_loss_backward(
1120 _log_probs: FloatTensor<B>,
1121 _targets: IntTensor<B>,
1122 _input_lengths: IntTensor<B>,
1123 _target_lengths: IntTensor<B>,
1124 _grad_loss: FloatTensor<B>,
1125 _blank: usize,
1126 ) -> FloatTensor<B> {
1127 unreachable!(
1128 "ctc_loss_backward called on a backend whose has_ctc_loss_backward() returns false"
1129 )
1130 }
1131
1132 fn rfft(
1144 signal: FloatTensor<B>,
1145 dim: usize,
1146 n: Option<usize>,
1147 ) -> (FloatTensor<B>, FloatTensor<B>);
1148
1149 fn irfft(
1157 spectrum_re: FloatTensor<B>,
1158 spectrum_im: FloatTensor<B>,
1159 dim: usize,
1160 n: Option<usize>,
1161 ) -> FloatTensor<B>;
1162}
1163
1164#[cfg(test)]
1165mod tests {
1166 use super::*;
1167
1168 #[test]
1169 #[should_panic = "stride must be non-zero"]
1170 fn conv_options_stride_zero() {
1171 let _opt = ConvOptions::new([0, 1], [0, 0], [1, 1], 1);
1172 }
1173
1174 #[test]
1175 #[should_panic = "dilation must be non-zero"]
1176 fn conv_options_dilation_zero() {
1177 let _opt = ConvOptions::new([1, 1], [0, 0], [0, 0], 1);
1178 }
1179
1180 #[test]
1181 #[should_panic = "groups must be non-zero"]
1182 fn conv_options_groups_zero() {
1183 let _opt = ConvOptions::new([1, 1], [0, 0], [1, 1], 0);
1184 }
1185
1186 #[test]
1187 #[should_panic = "stride must be non-zero"]
1188 fn conv_transpose_options_stride_zero() {
1189 let _opt = ConvTransposeOptions::new([0, 1], [0, 0], [0, 0], [1, 1], 1);
1190 }
1191
1192 #[test]
1193 #[should_panic = "dilation must be non-zero"]
1194 fn conv_transpose_options_dilation_zero() {
1195 let _opt = ConvTransposeOptions::new([1, 1], [0, 0], [0, 0], [0, 0], 1);
1196 }
1197
1198 #[test]
1199 #[should_panic = "groups must be non-zero"]
1200 fn conv_transpose_options_groups_zero() {
1201 let _opt = ConvTransposeOptions::new([1, 1], [0, 0], [0, 0], [1, 1], 0);
1202 }
1203
1204 #[test]
1205 #[should_panic = "stride must be non-zero"]
1206 fn deform_conv_options_stride_zero() {
1207 let _opt = DeformConvOptions::new([0, 1], [0, 0], [1, 1], 1, 1);
1208 }
1209
1210 #[test]
1211 #[should_panic = "dilation must be non-zero"]
1212 fn deform_conv_options_dilation_zero() {
1213 let _opt = DeformConvOptions::new([1, 1], [0, 0], [0, 0], 1, 1);
1214 }
1215
1216 #[test]
1217 #[should_panic = "weight groups must be non-zero"]
1218 fn deform_conv_options_weights_groups_zero() {
1219 let _opt = DeformConvOptions::new([1, 1], [0, 0], [1, 1], 0, 1);
1220 }
1221
1222 #[test]
1223 #[should_panic = "offset groups must be non-zero"]
1224 fn deform_conv_options_offset_groups_zero() {
1225 let _opt = DeformConvOptions::new([1, 1], [0, 0], [1, 1], 1, 0);
1226 }
1227
1228 #[test]
1229 #[should_panic = "stride must be non-zero"]
1230 fn unfold_options_stride_zero() {
1231 let _opt = UnfoldOptions::new([0, 1], [0, 0], [1, 1]);
1232 }
1233
1234 #[test]
1235 #[should_panic = "dilation must be non-zero"]
1236 fn unfold_options_dilation_zero() {
1237 let _opt = UnfoldOptions::new([1, 1], [0, 0], [0, 0]);
1238 }
1239}