1use super::{conv, ctc, 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
7#[derive(new)]
9pub struct Conv2dBackward<B: Backend> {
10 pub x_grad: FloatTensor<B>,
12
13 pub weights_grad: FloatTensor<B>,
15
16 pub bias_grad: Option<FloatTensor<B>>,
18}
19
20#[derive(new)]
22pub struct DeformConv2dBackward<B: Backend> {
23 pub x_grad: FloatTensor<B>,
25
26 pub offset_grad: FloatTensor<B>,
28
29 pub weight_grad: FloatTensor<B>,
31
32 pub mask_grad: Option<FloatTensor<B>>,
34
35 pub bias_grad: Option<FloatTensor<B>>,
37}
38
39#[derive(new)]
41pub struct Conv3dBackward<B: Backend> {
42 pub x_grad: FloatTensor<B>,
44
45 pub weights_grad: FloatTensor<B>,
47
48 pub bias_grad: Option<FloatTensor<B>>,
50}
51
52#[derive(new)]
54pub struct MaxPool1dBackward<B: Backend> {
55 pub x_grad: FloatTensor<B>,
57}
58
59#[derive(new)]
61pub struct MaxPool1dWithIndices<B: Backend> {
62 pub output: FloatTensor<B>,
64
65 pub indices: IntTensor<B>,
67}
68
69#[derive(new)]
71pub struct MaxPool2dBackward<B: Backend> {
72 pub x_grad: FloatTensor<B>,
74}
75
76#[derive(new)]
78pub struct MaxPool2dWithIndices<B: Backend> {
79 pub output: FloatTensor<B>,
81
82 pub indices: IntTensor<B>,
84}
85
86pub use ruda_core::tensor::spatial::{ConvOptions, PaddedConvOptions, DeformConvOptions, ConvTransposeOptions, UnfoldOptions};
87
88pub use ruda_core::tensor::spatial::{InterpolateMode, InterpolateOptions};
89
90pub use ruda_core::tensor::spatial::{GridSampleOptions, GridSamplePaddingMode};
91
92#[derive(Debug, Clone, Copy, PartialEq, serde::Deserialize, serde::Serialize)]
103pub enum PadMode {
104 Constant(f32),
110
111 Reflect,
119
120 Edge,
126}
127
128impl Default for PadMode {
129 fn default() -> Self {
130 PadMode::Constant(0.0)
131 }
132}
133
134impl<E: ElementConversion> From<E> for PadMode {
135 fn from(value: E) -> Self {
136 PadMode::Constant(value.elem())
137 }
138}
139
140#[derive(new)]
142pub struct InterpolateBackward<B: Backend> {
143 pub x_grad: FloatTensor<B>,
145}
146
147pub use ruda_core::tensor::spatial::AttentionModuleOptions;
148
149pub trait ModuleOps<B: Backend> {
151 fn embedding(weights: FloatTensor<B>, indices: IntTensor<B>) -> FloatTensor<B> {
162 let [batch_size, seq_length] = indices.shape().dims();
163 let [_, d_model] = weights.shape().dims();
164
165 let indices = B::int_reshape(indices, Shape::new([batch_size * seq_length]));
166 let output = B::float_select(weights, 0, indices);
167
168 B::float_reshape(output, Shape::new([batch_size, seq_length, d_model]))
169 }
170
171 fn embedding_backward(
183 weights: FloatTensor<B>,
184 output_grad: FloatTensor<B>,
185 indices: IntTensor<B>,
186 ) -> FloatTensor<B> {
187 let [batch_size, seq_length] = indices.shape().dims();
188 let [n_embeddings, d_model] = weights.shape().dims();
189 let device = B::float_device(&weights);
190 let dtype = output_grad.dtype();
191
192 let indices = B::int_reshape(indices, Shape::new([batch_size * seq_length]));
193 let output_grad =
194 B::float_reshape(output_grad, Shape::new([batch_size * seq_length, d_model]));
195 let grad = B::float_zeros(Shape::new([n_embeddings, d_model]), &device, dtype.into());
196
197 B::float_select_add(grad, 0, indices, output_grad)
198 }
199
200 fn linear(
208 x: FloatTensor<B>,
209 weight: FloatTensor<B>,
210 bias: Option<FloatTensor<B>>,
211 ) -> FloatTensor<B> {
212 linear::linear::<B>(x, weight, bias)
213 }
214 fn linear_x_backward(weight: FloatTensor<B>, output_grad: FloatTensor<B>) -> FloatTensor<B> {
216 linear::linear_x_backward::<B>(weight, output_grad)
217 }
218 fn linear_weight_backward(x: FloatTensor<B>, output_grad: FloatTensor<B>) -> FloatTensor<B> {
220 linear::linear_weight_backward::<B>(x, output_grad)
221 }
222 fn linear_bias_backward(output_grad: FloatTensor<B>) -> FloatTensor<B> {
224 linear::linear_bias_backward::<B>(output_grad)
225 }
226
227 fn conv1d(
235 x: FloatTensor<B>,
236 weight: FloatTensor<B>,
237 bias: Option<FloatTensor<B>>,
238 options: ConvOptions<1>,
239 ) -> FloatTensor<B> {
240 conv::conv1d_from_conv2d::<B>(x, weight, bias, options)
241 }
242 fn conv1d_x_backward(
244 x: FloatTensor<B>,
245 weight: FloatTensor<B>,
246 output_grad: FloatTensor<B>,
247 options: ConvOptions<1>,
248 ) -> FloatTensor<B> {
249 conv::conv1d_x_backward::<B>(x, weight, output_grad, options)
250 }
251 fn conv1d_weight_backward(
253 x: FloatTensor<B>,
254 weight: FloatTensor<B>,
255 output_grad: FloatTensor<B>,
256 options: ConvOptions<1>,
257 ) -> FloatTensor<B> {
258 conv::conv1d_weight_backward::<B>(x, weight, output_grad, options)
259 }
260 fn conv1d_bias_backward(
262 x: FloatTensor<B>,
263 bias: FloatTensor<B>,
264 output_grad: FloatTensor<B>,
265 ) -> FloatTensor<B> {
266 conv::conv1d_bias_backward::<B>(x, bias, output_grad)
267 }
268 fn conv2d(
276 x: FloatTensor<B>,
277 weight: FloatTensor<B>,
278 bias: Option<FloatTensor<B>>,
279 options: ConvOptions<2>,
280 ) -> FloatTensor<B>;
281 fn conv2d_x_backward(
283 x: FloatTensor<B>,
284 weight: FloatTensor<B>,
285 output_grad: FloatTensor<B>,
286 options: ConvOptions<2>,
287 ) -> FloatTensor<B> {
288 conv::conv2d_x_backward::<B>(x, weight, output_grad, options)
289 }
290 fn conv2d_weight_backward(
292 x: FloatTensor<B>,
293 weight: FloatTensor<B>,
294 output_grad: FloatTensor<B>,
295 options: ConvOptions<2>,
296 ) -> FloatTensor<B> {
297 conv::conv2d_weight_backward::<B>(x, weight, output_grad, options)
298 }
299 fn conv2d_bias_backward(
301 x: FloatTensor<B>,
302 bias: FloatTensor<B>,
303 output_grad: FloatTensor<B>,
304 ) -> FloatTensor<B> {
305 conv::conv2d_bias_backward::<B>(x, bias, output_grad)
306 }
307
308 fn deform_conv2d(
316 x: FloatTensor<B>,
317 offset: FloatTensor<B>,
318 weight: FloatTensor<B>,
319 mask: Option<FloatTensor<B>>,
320 bias: Option<FloatTensor<B>>,
321 options: DeformConvOptions<2>,
322 ) -> FloatTensor<B>;
323 fn deform_conv2d_backward(
325 x: FloatTensor<B>,
326 offset: FloatTensor<B>,
327 weight: FloatTensor<B>,
328 mask: Option<FloatTensor<B>>,
329 bias: Option<FloatTensor<B>>,
330 output_grad: FloatTensor<B>,
331 options: DeformConvOptions<2>,
332 ) -> DeformConv2dBackward<B>;
333
334 fn conv3d(
342 x: FloatTensor<B>,
343 weight: FloatTensor<B>,
344 bias: Option<FloatTensor<B>>,
345 options: ConvOptions<3>,
346 ) -> FloatTensor<B>;
347 fn conv3d_x_backward(
349 x: FloatTensor<B>,
350 weight: FloatTensor<B>,
351 output_grad: FloatTensor<B>,
352 options: ConvOptions<3>,
353 ) -> FloatTensor<B> {
354 conv::conv3d_x_backward::<B>(x, weight, output_grad, options)
355 }
356 fn conv3d_weight_backward(
358 x: FloatTensor<B>,
359 weight: FloatTensor<B>,
360 output_grad: FloatTensor<B>,
361 options: ConvOptions<3>,
362 ) -> FloatTensor<B> {
363 conv::conv3d_weight_backward::<B>(x, weight, output_grad, options)
364 }
365 fn conv3d_bias_backward(
367 x: FloatTensor<B>,
368 bias: FloatTensor<B>,
369 output_grad: FloatTensor<B>,
370 ) -> FloatTensor<B> {
371 conv::conv3d_bias_backward::<B>(x, bias, output_grad)
372 }
373 fn conv_transpose1d(
381 x: FloatTensor<B>,
382 weight: FloatTensor<B>,
383 bias: Option<FloatTensor<B>>,
384 options: ConvTransposeOptions<1>,
385 ) -> FloatTensor<B> {
386 conv::conv_transpose1d_from_conv_transpose2d::<B>(x, weight, bias, options)
387 }
388 fn conv_transpose1d_x_backward(
390 weight: FloatTensor<B>,
391 output_grad: FloatTensor<B>,
392 options: ConvTransposeOptions<1>,
393 ) -> FloatTensor<B> {
394 conv::conv_transpose1d_x_backward::<B>(weight, output_grad, options)
395 }
396 fn conv_transpose1d_weight_backward(
398 x: FloatTensor<B>,
399 weight: FloatTensor<B>,
400 output_grad: FloatTensor<B>,
401 options: ConvTransposeOptions<1>,
402 ) -> FloatTensor<B> {
403 conv::conv_transpose1d_weight_backward::<B>(x, weight, output_grad, options)
404 }
405 fn conv_transpose1d_bias_backward(
407 x: FloatTensor<B>,
408 bias: FloatTensor<B>,
409 output_grad: FloatTensor<B>,
410 ) -> FloatTensor<B> {
411 conv::conv_transpose1d_bias_backward::<B>(x, bias, output_grad)
412 }
413
414 fn conv_transpose2d(
422 x: FloatTensor<B>,
423 weight: FloatTensor<B>,
424 bias: Option<FloatTensor<B>>,
425 options: ConvTransposeOptions<2>,
426 ) -> FloatTensor<B>;
427 fn conv_transpose2d_x_backward(
429 weight: FloatTensor<B>,
430 output_grad: FloatTensor<B>,
431 options: ConvTransposeOptions<2>,
432 ) -> FloatTensor<B> {
433 conv::conv_transpose2d_x_backward::<B>(weight, output_grad, options)
434 }
435 fn conv_transpose2d_weight_backward(
437 x: FloatTensor<B>,
438 weight: FloatTensor<B>,
439 output_grad: FloatTensor<B>,
440 options: ConvTransposeOptions<2>,
441 ) -> FloatTensor<B> {
442 conv::conv_transpose2d_weight_backward::<B>(x, weight, output_grad, options)
443 }
444 fn conv_transpose2d_bias_backward(
446 x: FloatTensor<B>,
447 bias: FloatTensor<B>,
448 output_grad: FloatTensor<B>,
449 ) -> FloatTensor<B> {
450 conv::conv_transpose2d_bias_backward::<B>(x, bias, output_grad)
451 }
452
453 fn conv_transpose3d(
461 x: FloatTensor<B>,
462 weight: FloatTensor<B>,
463 bias: Option<FloatTensor<B>>,
464 options: ConvTransposeOptions<3>,
465 ) -> FloatTensor<B>;
466 fn conv_transpose3d_x_backward(
468 weight: FloatTensor<B>,
469 output_grad: FloatTensor<B>,
470 options: ConvTransposeOptions<3>,
471 ) -> FloatTensor<B> {
472 conv::conv_transpose3d_x_backward::<B>(weight, output_grad, options)
473 }
474 fn conv_transpose3d_weight_backward(
476 x: FloatTensor<B>,
477 weight: FloatTensor<B>,
478 output_grad: FloatTensor<B>,
479 options: ConvTransposeOptions<3>,
480 ) -> FloatTensor<B> {
481 conv::conv_transpose3d_weight_backward::<B>(x, weight, output_grad, options)
482 }
483 fn conv_transpose3d_bias_backward(
485 x: FloatTensor<B>,
486 bias: FloatTensor<B>,
487 output_grad: FloatTensor<B>,
488 ) -> FloatTensor<B> {
489 conv::conv_transpose3d_bias_backward::<B>(x, bias, output_grad)
490 }
491
492 fn unfold4d(
499 x: FloatTensor<B>,
500 kernel_size: [usize; 2],
501 options: UnfoldOptions,
502 ) -> FloatTensor<B> {
503 if options.padding == [0, 0] && options.dilation == [1, 1] {
504 let blocks = B::float_unfold(x, 2, kernel_size[0], options.stride[0]);
505 let blocks = B::float_unfold(blocks, 3, kernel_size[1], options.stride[1]);
506
507 let blocks = B::float_permute(blocks, &[0, 1, 4, 5, 2, 3]);
510 let shape = blocks.shape();
511
512 B::float_reshape(
515 blocks,
516 [
517 shape[0],
518 shape[1] * shape[2] * shape[3],
519 shape[4] * shape[5],
520 ]
521 .into(),
522 )
523 } else {
524 unfold4d_using_conv2d::<B>(x, kernel_size, options)
525 }
526 }
527
528 fn avg_pool1d(
534 x: FloatTensor<B>,
535 kernel_size: usize,
536 stride: usize,
537 padding: usize,
538 count_include_pad: bool,
539 ceil_mode: bool,
540 ) -> FloatTensor<B> {
541 pool::avg_pool1d_from_2d::<B>(
542 x,
543 kernel_size,
544 stride,
545 padding,
546 count_include_pad,
547 ceil_mode,
548 )
549 }
550 fn avg_pool1d_backward(
552 x: FloatTensor<B>,
553 grad: FloatTensor<B>,
554 kernel_size: usize,
555 stride: usize,
556 padding: usize,
557 count_include_pad: bool,
558 ceil_mode: bool,
559 ) -> FloatTensor<B> {
560 pool::avg_pool1d_backward_from_2d::<B>(
561 x,
562 grad,
563 kernel_size,
564 stride,
565 padding,
566 count_include_pad,
567 ceil_mode,
568 )
569 }
570 fn avg_pool2d(
576 x: FloatTensor<B>,
577 kernel_size: [usize; 2],
578 stride: [usize; 2],
579 padding: [usize; 2],
580 count_include_pad: bool,
581 ceil_mode: bool,
582 ) -> FloatTensor<B>;
583 fn avg_pool2d_backward(
585 x: FloatTensor<B>,
586 grad: FloatTensor<B>,
587 kernel_size: [usize; 2],
588 stride: [usize; 2],
589 padding: [usize; 2],
590 count_include_pad: bool,
591 ceil_mode: bool,
592 ) -> FloatTensor<B>;
593 fn adaptive_avg_pool2d(x: FloatTensor<B>, output_size: [usize; 2]) -> FloatTensor<B>;
599 fn adaptive_avg_pool2d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B>;
601 fn adaptive_avg_pool1d(x: FloatTensor<B>, output_size: usize) -> FloatTensor<B> {
607 pool::adaptive_avg_pool1d_from_2d::<B>(x, output_size)
608 }
609 fn adaptive_avg_pool1d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
611 pool::adaptive_avg_pool1d_backward_from_2d::<B>(x, grad)
612 }
613 fn max_pool1d(
619 x: FloatTensor<B>,
620 kernel_size: usize,
621 stride: usize,
622 padding: usize,
623 dilation: usize,
624 ceil_mode: bool,
625 ) -> FloatTensor<B> {
626 pool::max_pool1d_from_2d::<B>(x, kernel_size, stride, padding, dilation, ceil_mode)
627 }
628
629 fn max_pool1d_with_indices(
635 x: FloatTensor<B>,
636 kernel_size: usize,
637 stride: usize,
638 padding: usize,
639 dilation: usize,
640 ceil_mode: bool,
641 ) -> MaxPool1dWithIndices<B> {
642 pool::max_pool1d_with_indices_from_2d::<B>(
643 x,
644 kernel_size,
645 stride,
646 padding,
647 dilation,
648 ceil_mode,
649 )
650 }
651 #[allow(clippy::too_many_arguments)]
653 fn max_pool1d_with_indices_backward(
654 x: FloatTensor<B>,
655 kernel_size: usize,
656 stride: usize,
657 padding: usize,
658 dilation: usize,
659 ceil_mode: bool,
660 output_grad: FloatTensor<B>,
661 indices: IntTensor<B>,
662 ) -> MaxPool1dBackward<B> {
663 pool::max_pool1d_with_indices_backward_from_2d::<B>(
664 x,
665 kernel_size,
666 stride,
667 padding,
668 dilation,
669 ceil_mode,
670 output_grad,
671 indices,
672 )
673 }
674
675 fn max_pool2d(
681 x: FloatTensor<B>,
682 kernel_size: [usize; 2],
683 stride: [usize; 2],
684 padding: [usize; 2],
685 dilation: [usize; 2],
686 ceil_mode: bool,
687 ) -> FloatTensor<B>;
688
689 fn max_pool2d_with_indices(
695 x: FloatTensor<B>,
696 kernel_size: [usize; 2],
697 stride: [usize; 2],
698 padding: [usize; 2],
699 dilation: [usize; 2],
700 ceil_mode: bool,
701 ) -> MaxPool2dWithIndices<B>;
702 #[allow(clippy::too_many_arguments)]
704 fn max_pool2d_with_indices_backward(
705 x: FloatTensor<B>,
706 kernel_size: [usize; 2],
707 stride: [usize; 2],
708 padding: [usize; 2],
709 dilation: [usize; 2],
710 ceil_mode: bool,
711 output_grad: FloatTensor<B>,
712 indices: IntTensor<B>,
713 ) -> MaxPool2dBackward<B>;
714
715 fn interpolate(
721 x: FloatTensor<B>,
722 output_size: [usize; 2],
723 options: InterpolateOptions,
724 ) -> FloatTensor<B>;
725
726 fn interpolate_backward(
728 x: FloatTensor<B>,
729 grad: FloatTensor<B>,
730 output_size: [usize; 2],
731 options: InterpolateOptions,
732 ) -> FloatTensor<B>;
733
734 fn attention(
756 query: FloatTensor<B>,
757 key: FloatTensor<B>,
758 value: FloatTensor<B>,
759 mask: Option<BoolTensor<B>>,
760 attn_bias: Option<FloatTensor<B>>,
761 options: AttentionModuleOptions,
762 ) -> FloatTensor<B>;
763
764 fn layer_norm(
780 tensor: FloatTensor<B>,
781 gamma: FloatTensor<B>,
782 beta: Option<FloatTensor<B>>,
783 epsilon: f64,
784 ) -> FloatTensor<B> {
785 let shape = tensor.shape();
786 let rank = shape.num_dims();
787 let last_dim = rank - 1;
788 let d_model = shape[last_dim];
789
790 let mean = B::float_mean_dim(tensor.clone(), last_dim);
791 let centered = B::float_sub(tensor, mean);
792 let var = B::float_mean_dim(B::float_mul(centered.clone(), centered.clone()), last_dim);
793 let denom = B::float_sqrt(B::float_add_scalar(var, epsilon.into()));
794 let normalized = B::float_div(centered, denom);
795
796 let broadcast_dims: alloc::vec::Vec<usize> = (0..rank)
797 .map(|i| if i == last_dim { d_model } else { 1 })
798 .collect();
799 let gamma_b = B::float_reshape(gamma, Shape::from(broadcast_dims.clone()));
800 let scaled = B::float_mul(normalized, gamma_b);
801
802 match beta {
803 Some(beta) => {
804 let beta_b = B::float_reshape(beta, Shape::from(broadcast_dims));
805 B::float_add(scaled, beta_b)
806 }
807 None => scaled,
808 }
809 }
810
811 fn ctc_loss(
828 log_probs: FloatTensor<B>,
829 targets: IntTensor<B>,
830 input_lengths: IntTensor<B>,
831 target_lengths: IntTensor<B>,
832 blank: usize,
833 ) -> FloatTensor<B> {
834 ctc::ctc_loss_default::<B>(log_probs, targets, input_lengths, target_lengths, blank)
835 }
836
837 fn has_ctc_loss_backward() -> bool {
849 false
850 }
851
852 fn ctc_loss_backward(
872 _log_probs: FloatTensor<B>,
873 _targets: IntTensor<B>,
874 _input_lengths: IntTensor<B>,
875 _target_lengths: IntTensor<B>,
876 _grad_loss: FloatTensor<B>,
877 _blank: usize,
878 ) -> FloatTensor<B> {
879 unreachable!(
880 "ctc_loss_backward called on a backend whose has_ctc_loss_backward() returns false"
881 )
882 }
883
884 fn rfft(
896 signal: FloatTensor<B>,
897 dim: usize,
898 n: Option<usize>,
899 ) -> (FloatTensor<B>, FloatTensor<B>);
900
901 fn irfft(
909 spectrum_re: FloatTensor<B>,
910 spectrum_im: FloatTensor<B>,
911 dim: usize,
912 n: Option<usize>,
913 ) -> FloatTensor<B>;
914}
915
916#[cfg(test)]
917mod tests {
918 use super::*;
919
920 #[test]
921 #[should_panic = "stride must be non-zero"]
922 fn conv_options_stride_zero() {
923 let _opt = ConvOptions::new([0, 1], [0, 0], [1, 1], 1);
924 }
925
926 #[test]
927 #[should_panic = "dilation must be non-zero"]
928 fn conv_options_dilation_zero() {
929 let _opt = ConvOptions::new([1, 1], [0, 0], [0, 0], 1);
930 }
931
932 #[test]
933 #[should_panic = "groups must be non-zero"]
934 fn conv_options_groups_zero() {
935 let _opt = ConvOptions::new([1, 1], [0, 0], [1, 1], 0);
936 }
937
938 #[test]
939 #[should_panic = "stride must be non-zero"]
940 fn conv_transpose_options_stride_zero() {
941 let _opt = ConvTransposeOptions::new([0, 1], [0, 0], [0, 0], [1, 1], 1);
942 }
943
944 #[test]
945 #[should_panic = "dilation must be non-zero"]
946 fn conv_transpose_options_dilation_zero() {
947 let _opt = ConvTransposeOptions::new([1, 1], [0, 0], [0, 0], [0, 0], 1);
948 }
949
950 #[test]
951 #[should_panic = "groups must be non-zero"]
952 fn conv_transpose_options_groups_zero() {
953 let _opt = ConvTransposeOptions::new([1, 1], [0, 0], [0, 0], [1, 1], 0);
954 }
955
956 #[test]
957 #[should_panic = "stride must be non-zero"]
958 fn deform_conv_options_stride_zero() {
959 let _opt = DeformConvOptions::new([0, 1], [0, 0], [1, 1], 1, 1);
960 }
961
962 #[test]
963 #[should_panic = "dilation must be non-zero"]
964 fn deform_conv_options_dilation_zero() {
965 let _opt = DeformConvOptions::new([1, 1], [0, 0], [0, 0], 1, 1);
966 }
967
968 #[test]
969 #[should_panic = "weight groups must be non-zero"]
970 fn deform_conv_options_weights_groups_zero() {
971 let _opt = DeformConvOptions::new([1, 1], [0, 0], [1, 1], 0, 1);
972 }
973
974 #[test]
975 #[should_panic = "offset groups must be non-zero"]
976 fn deform_conv_options_offset_groups_zero() {
977 let _opt = DeformConvOptions::new([1, 1], [0, 0], [1, 1], 1, 0);
978 }
979
980 #[test]
981 #[should_panic = "stride must be non-zero"]
982 fn unfold_options_stride_zero() {
983 let _opt = UnfoldOptions::new([0, 1], [0, 0], [1, 1]);
984 }
985
986 #[test]
987 #[should_panic = "dilation must be non-zero"]
988 fn unfold_options_dilation_zero() {
989 let _opt = UnfoldOptions::new([1, 1], [0, 0], [0, 0]);
990 }
991}