1use super::{conv, ctc, linear, pool};
2use crate::ops::unfold::{create_unfolding_weight, unfold4d_using_conv2d};
3use crate::tensor::{BoolTensor, FloatTensor, IntTensor};
4use crate::{Backend, Scalar, TensorMetadata};
5#[allow(deprecated)]
6pub use burn_std::ops::{
7 AttentionModuleOptions, ConvOptions, ConvTransposeOptions, DeformConvOptions,
8 GridSampleOptions, GridSamplePaddingMode, InterpolateMode, InterpolateOptions, PadMode,
9 PaddedConvOptions, UnfoldOptions,
10};
11use burn_std::{IndexingUpdateOp, IntDType, Shape};
12
13#[derive(new)]
15pub struct Conv2dBackward<B: Backend> {
16 pub x_grad: FloatTensor<B>,
18
19 pub weights_grad: FloatTensor<B>,
21
22 pub bias_grad: Option<FloatTensor<B>>,
24}
25
26#[derive(new)]
28pub struct DeformConv2dBackward<B: Backend> {
29 pub x_grad: FloatTensor<B>,
31
32 pub offset_grad: FloatTensor<B>,
34
35 pub weight_grad: FloatTensor<B>,
37
38 pub mask_grad: Option<FloatTensor<B>>,
40
41 pub bias_grad: Option<FloatTensor<B>>,
43}
44
45#[derive(new)]
47pub struct Conv3dBackward<B: Backend> {
48 pub x_grad: FloatTensor<B>,
50
51 pub weights_grad: FloatTensor<B>,
53
54 pub bias_grad: Option<FloatTensor<B>>,
56}
57
58#[derive(new)]
60pub struct MaxPool1dBackward<B: Backend> {
61 pub x_grad: FloatTensor<B>,
63}
64
65#[derive(new)]
67pub struct MaxPool1dWithIndices<B: Backend> {
68 pub output: FloatTensor<B>,
70
71 pub indices: IntTensor<B>,
73}
74
75#[derive(new)]
77pub struct MaxPool2dBackward<B: Backend> {
78 pub x_grad: FloatTensor<B>,
80}
81
82#[derive(new)]
84pub struct BatchNormTrain<B: Backend> {
85 pub output: FloatTensor<B>,
87
88 pub mean: FloatTensor<B>,
90
91 pub variance: FloatTensor<B>,
93}
94
95#[derive(new)]
98pub struct BatchNormTrainBackward<B: Backend> {
99 pub x_grad: FloatTensor<B>,
101
102 pub gamma_grad: FloatTensor<B>,
104
105 pub beta_grad: FloatTensor<B>,
107}
108
109#[derive(new)]
111pub struct MaxPool2dWithIndices<B: Backend> {
112 pub output: FloatTensor<B>,
114
115 pub indices: IntTensor<B>,
117}
118
119#[derive(new)]
121pub struct InterpolateBackward<B: Backend> {
122 pub x_grad: FloatTensor<B>,
124}
125
126pub trait ModuleOps<B: Backend> {
128 fn batch_norm(
137 x: FloatTensor<B>,
138 gamma: FloatTensor<B>,
139 beta: FloatTensor<B>,
140 mean: FloatTensor<B>,
141 variance: FloatTensor<B>,
142 epsilon: f64,
143 ) -> FloatTensor<B> {
144 let rank = x.shape().num_dims();
145 let channels = x.shape()[1];
146 let mut dimensions = alloc::vec![1; rank];
147 dimensions[1] = channels;
148 let shape = Shape::from(dimensions);
149 let gamma = B::float_reshape(gamma, shape.clone());
150 let beta = B::float_reshape(beta, shape.clone());
151 let mean = B::float_reshape(mean, shape.clone());
152 let variance = B::float_reshape(variance, shape);
153 let std = B::float_sqrt(B::float_add_scalar(variance, Scalar::Float(epsilon)));
154 let normalized = B::float_div(B::float_sub(x, mean), std);
155 B::float_add(B::float_mul(normalized, gamma), beta)
156 }
157
158 fn batch_norm_train(
170 x: FloatTensor<B>,
171 gamma: FloatTensor<B>,
172 beta: FloatTensor<B>,
173 epsilon: f64,
174 ) -> BatchNormTrain<B> {
175 let shape = x.shape();
176 let channels = shape[1];
177 let samples_per_channel = shape.num_elements() / channels;
178 let flattened = Shape::new([channels, samples_per_channel]);
179 let mut per_channel = alloc::vec![1; shape.num_dims()];
180 per_channel[1] = channels;
181
182 let mean = B::float_mean_dim(
184 B::float_reshape(B::float_swap_dims(x.clone(), 0, 1), flattened.clone()),
185 1,
186 );
187 let mean = B::float_reshape(mean, Shape::from(per_channel));
188 let centered = B::float_sub(x.clone(), mean.clone());
189 let variance = B::float_mean_dim(
190 B::float_reshape(
191 B::float_swap_dims(B::float_mul(centered.clone(), centered), 0, 1),
192 flattened,
193 ),
194 1,
195 );
196 let mean = B::float_reshape(mean, Shape::new([channels]));
197 let variance = B::float_reshape(variance, Shape::new([channels]));
198 let output = B::batch_norm(x, gamma, beta, mean.clone(), variance.clone(), epsilon);
199
200 BatchNormTrain::new(output, mean, variance)
201 }
202
203 fn batch_norm_train_backward(
207 x: FloatTensor<B>,
208 gamma: FloatTensor<B>,
209 mean: FloatTensor<B>,
210 variance: FloatTensor<B>,
211 epsilon: f64,
212 output_grad: FloatTensor<B>,
213 ) -> BatchNormTrainBackward<B> {
214 let shape = x.shape();
215 let rank = shape.num_dims();
216 let channels = shape[1];
217 let flattened = Shape::new([channels, shape.num_elements() / channels]);
218 let mut per_channel = alloc::vec![1; rank];
219 per_channel[1] = channels;
220 let per_channel = Shape::from(per_channel);
221
222 let inv_std = B::float_reshape(
223 B::float_recip(B::float_sqrt(B::float_add_scalar(
224 variance,
225 Scalar::Float(epsilon),
226 ))),
227 per_channel.clone(),
228 );
229 let normalized = B::float_mul(
230 B::float_sub(x, B::float_reshape(mean, per_channel.clone())),
231 inv_std.clone(),
232 );
233
234 let output_grad_flat = B::float_reshape(
235 B::float_swap_dims(output_grad.clone(), 0, 1),
236 flattened.clone(),
237 );
238 let normalized_grad_flat = B::float_reshape(
239 B::float_swap_dims(B::float_mul(output_grad.clone(), normalized.clone()), 0, 1),
240 flattened,
241 );
242 let beta_grad = B::float_sum_dim(output_grad_flat.clone(), 1);
243 let gamma_grad = B::float_sum_dim(normalized_grad_flat.clone(), 1);
244
245 let mean_grad =
247 B::float_reshape(B::float_mean_dim(output_grad_flat, 1), per_channel.clone());
248 let mean_normalized_grad = B::float_reshape(
249 B::float_mean_dim(normalized_grad_flat, 1),
250 per_channel.clone(),
251 );
252 let centred_grad = B::float_sub(output_grad, mean_grad);
253 let projected = B::float_mul(normalized, mean_normalized_grad);
254 let scale = B::float_mul(B::float_reshape(gamma, per_channel), inv_std);
255 let x_grad = B::float_mul(scale, B::float_sub(centred_grad, projected));
256
257 BatchNormTrainBackward::new(
258 x_grad,
259 B::float_reshape(gamma_grad, Shape::new([channels])),
260 B::float_reshape(beta_grad, Shape::new([channels])),
261 )
262 }
263
264 fn embedding(weights: FloatTensor<B>, indices: IntTensor<B>) -> FloatTensor<B> {
275 let [batch_size, seq_length] = indices.shape().dims();
276 let [_, d_model] = weights.shape().dims();
277
278 let indices = B::int_reshape(indices, Shape::new([batch_size * seq_length]));
279 let output = B::float_select(weights, 0, indices);
280
281 B::float_reshape(output, Shape::new([batch_size, seq_length, d_model]))
282 }
283
284 fn embedding_backward(
296 weights: FloatTensor<B>,
297 output_grad: FloatTensor<B>,
298 indices: IntTensor<B>,
299 ) -> FloatTensor<B> {
300 let [batch_size, seq_length] = indices.shape().dims();
301 let [n_embeddings, d_model] = weights.shape().dims();
302 let device = weights.device();
303 let dtype = output_grad.dtype();
304
305 let indices = B::int_reshape(indices, Shape::new([batch_size * seq_length]));
306 let output_grad =
307 B::float_reshape(output_grad, Shape::new([batch_size * seq_length, d_model]));
308 let grad = B::float_zeros(Shape::new([n_embeddings, d_model]), &device, dtype.into());
309
310 B::float_select_assign(grad, 0, indices, output_grad, IndexingUpdateOp::Add)
311 }
312
313 fn linear(
321 x: FloatTensor<B>,
322 weight: FloatTensor<B>,
323 bias: Option<FloatTensor<B>>,
324 ) -> FloatTensor<B> {
325 linear::linear::<B>(x, weight, bias)
326 }
327 fn linear_x_backward(weight: FloatTensor<B>, output_grad: FloatTensor<B>) -> FloatTensor<B> {
329 linear::linear_x_backward::<B>(weight, output_grad)
330 }
331 fn linear_weight_backward(x: FloatTensor<B>, output_grad: FloatTensor<B>) -> FloatTensor<B> {
333 linear::linear_weight_backward::<B>(x, output_grad)
334 }
335 fn linear_bias_backward(output_grad: FloatTensor<B>) -> FloatTensor<B> {
337 linear::linear_bias_backward::<B>(output_grad)
338 }
339
340 fn conv1d(
348 x: FloatTensor<B>,
349 weight: FloatTensor<B>,
350 bias: Option<FloatTensor<B>>,
351 options: ConvOptions<1>,
352 ) -> FloatTensor<B> {
353 conv::conv1d_from_conv2d::<B>(x, weight, bias, options)
354 }
355 fn conv1d_x_backward(
357 x: FloatTensor<B>,
358 weight: FloatTensor<B>,
359 output_grad: FloatTensor<B>,
360 options: ConvOptions<1>,
361 ) -> FloatTensor<B> {
362 conv::conv1d_x_backward::<B>(x, weight, output_grad, options)
363 }
364 fn conv1d_weight_backward(
366 x: FloatTensor<B>,
367 weight: FloatTensor<B>,
368 output_grad: FloatTensor<B>,
369 options: ConvOptions<1>,
370 ) -> FloatTensor<B> {
371 conv::conv1d_weight_backward::<B>(x, weight, output_grad, options)
372 }
373 fn conv1d_bias_backward(
375 x: FloatTensor<B>,
376 bias: FloatTensor<B>,
377 output_grad: FloatTensor<B>,
378 ) -> FloatTensor<B> {
379 conv::conv1d_bias_backward::<B>(x, bias, output_grad)
380 }
381 fn conv2d(
389 x: FloatTensor<B>,
390 weight: FloatTensor<B>,
391 bias: Option<FloatTensor<B>>,
392 options: ConvOptions<2>,
393 ) -> FloatTensor<B>;
394 fn conv2d_x_backward(
396 x: FloatTensor<B>,
397 weight: FloatTensor<B>,
398 output_grad: FloatTensor<B>,
399 options: ConvOptions<2>,
400 ) -> FloatTensor<B> {
401 conv::conv2d_x_backward::<B>(x, weight, output_grad, options)
402 }
403 fn conv2d_weight_backward(
405 x: FloatTensor<B>,
406 weight: FloatTensor<B>,
407 output_grad: FloatTensor<B>,
408 options: ConvOptions<2>,
409 ) -> FloatTensor<B> {
410 conv::conv2d_weight_backward::<B>(x, weight, output_grad, options)
411 }
412 fn conv2d_bias_backward(
414 x: FloatTensor<B>,
415 bias: FloatTensor<B>,
416 output_grad: FloatTensor<B>,
417 ) -> FloatTensor<B> {
418 conv::conv2d_bias_backward::<B>(x, bias, output_grad)
419 }
420
421 fn deform_conv2d(
429 x: FloatTensor<B>,
430 offset: FloatTensor<B>,
431 weight: FloatTensor<B>,
432 mask: Option<FloatTensor<B>>,
433 bias: Option<FloatTensor<B>>,
434 options: DeformConvOptions<2>,
435 ) -> FloatTensor<B>;
436 fn deform_conv2d_backward(
438 x: FloatTensor<B>,
439 offset: FloatTensor<B>,
440 weight: FloatTensor<B>,
441 mask: Option<FloatTensor<B>>,
442 bias: Option<FloatTensor<B>>,
443 output_grad: FloatTensor<B>,
444 options: DeformConvOptions<2>,
445 ) -> DeformConv2dBackward<B>;
446
447 fn conv3d(
455 x: FloatTensor<B>,
456 weight: FloatTensor<B>,
457 bias: Option<FloatTensor<B>>,
458 options: ConvOptions<3>,
459 ) -> FloatTensor<B>;
460 fn conv3d_x_backward(
462 x: FloatTensor<B>,
463 weight: FloatTensor<B>,
464 output_grad: FloatTensor<B>,
465 options: ConvOptions<3>,
466 ) -> FloatTensor<B> {
467 conv::conv3d_x_backward::<B>(x, weight, output_grad, options)
468 }
469 fn conv3d_weight_backward(
471 x: FloatTensor<B>,
472 weight: FloatTensor<B>,
473 output_grad: FloatTensor<B>,
474 options: ConvOptions<3>,
475 ) -> FloatTensor<B> {
476 conv::conv3d_weight_backward::<B>(x, weight, output_grad, options)
477 }
478 fn conv3d_bias_backward(
480 x: FloatTensor<B>,
481 bias: FloatTensor<B>,
482 output_grad: FloatTensor<B>,
483 ) -> FloatTensor<B> {
484 conv::conv3d_bias_backward::<B>(x, bias, output_grad)
485 }
486 fn conv_transpose1d(
494 x: FloatTensor<B>,
495 weight: FloatTensor<B>,
496 bias: Option<FloatTensor<B>>,
497 options: ConvTransposeOptions<1>,
498 ) -> FloatTensor<B> {
499 conv::conv_transpose1d_from_conv_transpose2d::<B>(x, weight, bias, options)
500 }
501 fn conv_transpose1d_x_backward(
503 weight: FloatTensor<B>,
504 output_grad: FloatTensor<B>,
505 options: ConvTransposeOptions<1>,
506 ) -> FloatTensor<B> {
507 conv::conv_transpose1d_x_backward::<B>(weight, output_grad, options)
508 }
509 fn conv_transpose1d_weight_backward(
511 x: FloatTensor<B>,
512 weight: FloatTensor<B>,
513 output_grad: FloatTensor<B>,
514 options: ConvTransposeOptions<1>,
515 ) -> FloatTensor<B> {
516 conv::conv_transpose1d_weight_backward::<B>(x, weight, output_grad, options)
517 }
518 fn conv_transpose1d_bias_backward(
520 x: FloatTensor<B>,
521 bias: FloatTensor<B>,
522 output_grad: FloatTensor<B>,
523 ) -> FloatTensor<B> {
524 conv::conv_transpose1d_bias_backward::<B>(x, bias, output_grad)
525 }
526
527 fn conv_transpose2d(
535 x: FloatTensor<B>,
536 weight: FloatTensor<B>,
537 bias: Option<FloatTensor<B>>,
538 options: ConvTransposeOptions<2>,
539 ) -> FloatTensor<B>;
540 fn conv_transpose2d_x_backward(
542 weight: FloatTensor<B>,
543 output_grad: FloatTensor<B>,
544 options: ConvTransposeOptions<2>,
545 ) -> FloatTensor<B> {
546 conv::conv_transpose2d_x_backward::<B>(weight, output_grad, options)
547 }
548 fn conv_transpose2d_weight_backward(
550 x: FloatTensor<B>,
551 weight: FloatTensor<B>,
552 output_grad: FloatTensor<B>,
553 options: ConvTransposeOptions<2>,
554 ) -> FloatTensor<B> {
555 conv::conv_transpose2d_weight_backward::<B>(x, weight, output_grad, options)
556 }
557 fn conv_transpose2d_bias_backward(
559 x: FloatTensor<B>,
560 bias: FloatTensor<B>,
561 output_grad: FloatTensor<B>,
562 ) -> FloatTensor<B> {
563 conv::conv_transpose2d_bias_backward::<B>(x, bias, output_grad)
564 }
565
566 fn conv_transpose3d(
574 x: FloatTensor<B>,
575 weight: FloatTensor<B>,
576 bias: Option<FloatTensor<B>>,
577 options: ConvTransposeOptions<3>,
578 ) -> FloatTensor<B>;
579 fn conv_transpose3d_x_backward(
581 weight: FloatTensor<B>,
582 output_grad: FloatTensor<B>,
583 options: ConvTransposeOptions<3>,
584 ) -> FloatTensor<B> {
585 conv::conv_transpose3d_x_backward::<B>(weight, output_grad, options)
586 }
587 fn conv_transpose3d_weight_backward(
589 x: FloatTensor<B>,
590 weight: FloatTensor<B>,
591 output_grad: FloatTensor<B>,
592 options: ConvTransposeOptions<3>,
593 ) -> FloatTensor<B> {
594 conv::conv_transpose3d_weight_backward::<B>(x, weight, output_grad, options)
595 }
596 fn conv_transpose3d_bias_backward(
598 x: FloatTensor<B>,
599 bias: FloatTensor<B>,
600 output_grad: FloatTensor<B>,
601 ) -> FloatTensor<B> {
602 conv::conv_transpose3d_bias_backward::<B>(x, bias, output_grad)
603 }
604
605 fn unfold4d(
612 x: FloatTensor<B>,
613 kernel_size: [usize; 2],
614 options: UnfoldOptions,
615 ) -> FloatTensor<B> {
616 if options.padding == [0, 0] && options.dilation == [1, 1] {
617 let blocks = B::float_unfold(x, 2, kernel_size[0], options.stride[0]);
618 let blocks = B::float_unfold(blocks, 3, kernel_size[1], options.stride[1]);
619
620 let blocks = B::float_permute(blocks, &[0, 1, 4, 5, 2, 3]);
623 let shape = blocks.shape();
624
625 B::float_reshape(
628 blocks,
629 [
630 shape[0],
631 shape[1] * shape[2] * shape[3],
632 shape[4] * shape[5],
633 ]
634 .into(),
635 )
636 } else {
637 unfold4d_using_conv2d::<B>(x, kernel_size, options)
638 }
639 }
640
641 fn fold4d(
652 x: FloatTensor<B>,
653 output_size: [usize; 2],
654 kernel_size: [usize; 2],
655 options: UnfoldOptions,
656 ) -> FloatTensor<B> {
657 let [batch_size, channels_col, num_blocks] = x.shape().dims();
658 let [kernel_height, kernel_width] = kernel_size;
659 let [output_height, output_width] = output_size;
660 let [stride_height, stride_width] = options.stride;
661 let [padding_height, padding_width] = options.padding;
662 let [dilation_height, dilation_width] = options.dilation;
663
664 let kernel_elems = kernel_height * kernel_width;
665 assert_eq!(
666 channels_col % kernel_elems,
667 0,
668 "fold4d: input channels ({channels_col}) must be divisible by the kernel size product ({kernel_elems})"
669 );
670 let channels = channels_col / kernel_elems;
671
672 let blocks_height =
674 (output_height + 2 * padding_height - dilation_height * (kernel_height - 1) - 1)
675 / stride_height
676 + 1;
677 let blocks_width =
678 (output_width + 2 * padding_width - dilation_width * (kernel_width - 1) - 1)
679 / stride_width
680 + 1;
681 assert_eq!(
682 num_blocks,
683 blocks_height * blocks_width,
684 "fold4d: number of blocks ({num_blocks}) does not match the expected grid ({blocks_height} x {blocks_width}) for the given output size and options"
685 );
686
687 let weight = create_unfolding_weight::<B>(channels, kernel_size, &x.device(), x.dtype());
689
690 let x = B::float_reshape(
692 x,
693 Shape::new([batch_size, channels_col, blocks_height, blocks_width]),
694 );
695
696 let padding_out = [
698 (output_height + 2 * padding_height - dilation_height * (kernel_height - 1) - 1)
699 % stride_height,
700 (output_width + 2 * padding_width - dilation_width * (kernel_width - 1) - 1)
701 % stride_width,
702 ];
703
704 B::conv_transpose2d(
705 x,
706 weight,
707 None,
708 ConvTransposeOptions::new(
709 options.stride,
710 options.padding,
711 padding_out,
712 options.dilation,
713 1,
714 ),
715 )
716 }
717
718 fn avg_pool1d(
724 x: FloatTensor<B>,
725 kernel_size: usize,
726 stride: usize,
727 padding: usize,
728 count_include_pad: bool,
729 ceil_mode: bool,
730 ) -> FloatTensor<B> {
731 pool::avg_pool1d_from_2d::<B>(
732 x,
733 kernel_size,
734 stride,
735 padding,
736 count_include_pad,
737 ceil_mode,
738 )
739 }
740 fn avg_pool1d_backward(
742 x: FloatTensor<B>,
743 grad: FloatTensor<B>,
744 kernel_size: usize,
745 stride: usize,
746 padding: usize,
747 count_include_pad: bool,
748 ceil_mode: bool,
749 ) -> FloatTensor<B> {
750 pool::avg_pool1d_backward_from_2d::<B>(
751 x,
752 grad,
753 kernel_size,
754 stride,
755 padding,
756 count_include_pad,
757 ceil_mode,
758 )
759 }
760 fn avg_pool2d(
766 x: FloatTensor<B>,
767 kernel_size: [usize; 2],
768 stride: [usize; 2],
769 padding: [usize; 2],
770 count_include_pad: bool,
771 ceil_mode: bool,
772 ) -> FloatTensor<B>;
773 fn avg_pool2d_backward(
775 x: FloatTensor<B>,
776 grad: FloatTensor<B>,
777 kernel_size: [usize; 2],
778 stride: [usize; 2],
779 padding: [usize; 2],
780 count_include_pad: bool,
781 ceil_mode: bool,
782 ) -> FloatTensor<B>;
783 fn adaptive_avg_pool2d(x: FloatTensor<B>, output_size: [usize; 2]) -> FloatTensor<B>;
789 fn adaptive_avg_pool2d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B>;
791 fn adaptive_avg_pool3d(x: FloatTensor<B>, output_size: [usize; 3]) -> FloatTensor<B>;
797 fn adaptive_avg_pool3d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B>;
799 fn adaptive_avg_pool1d(x: FloatTensor<B>, output_size: usize) -> FloatTensor<B> {
805 pool::adaptive_avg_pool1d_from_2d::<B>(x, output_size)
806 }
807 fn adaptive_avg_pool1d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
809 pool::adaptive_avg_pool1d_backward_from_2d::<B>(x, grad)
810 }
811 fn max_pool1d(
817 x: FloatTensor<B>,
818 kernel_size: usize,
819 stride: usize,
820 padding: usize,
821 dilation: usize,
822 ceil_mode: bool,
823 ) -> FloatTensor<B> {
824 pool::max_pool1d_from_2d::<B>(x, kernel_size, stride, padding, dilation, ceil_mode)
825 }
826
827 fn max_pool1d_with_indices(
833 x: FloatTensor<B>,
834 kernel_size: usize,
835 stride: usize,
836 padding: usize,
837 dilation: usize,
838 ceil_mode: bool,
839 indices_dtype: IntDType,
840 ) -> MaxPool1dWithIndices<B> {
841 pool::max_pool1d_with_indices_from_2d::<B>(
842 x,
843 kernel_size,
844 stride,
845 padding,
846 dilation,
847 ceil_mode,
848 indices_dtype,
849 )
850 }
851 #[allow(clippy::too_many_arguments)]
853 fn max_pool1d_with_indices_backward(
854 x: FloatTensor<B>,
855 kernel_size: usize,
856 stride: usize,
857 padding: usize,
858 dilation: usize,
859 ceil_mode: bool,
860 output_grad: FloatTensor<B>,
861 indices: IntTensor<B>,
862 ) -> MaxPool1dBackward<B> {
863 pool::max_pool1d_with_indices_backward_from_2d::<B>(
864 x,
865 kernel_size,
866 stride,
867 padding,
868 dilation,
869 ceil_mode,
870 output_grad,
871 indices,
872 )
873 }
874
875 fn max_pool2d(
881 x: FloatTensor<B>,
882 kernel_size: [usize; 2],
883 stride: [usize; 2],
884 padding: [usize; 2],
885 dilation: [usize; 2],
886 ceil_mode: bool,
887 ) -> FloatTensor<B>;
888
889 fn max_pool2d_with_indices(
895 x: FloatTensor<B>,
896 kernel_size: [usize; 2],
897 stride: [usize; 2],
898 padding: [usize; 2],
899 dilation: [usize; 2],
900 ceil_mode: bool,
901 indices_dtype: IntDType,
902 ) -> MaxPool2dWithIndices<B>;
903 #[allow(clippy::too_many_arguments)]
905 fn max_pool2d_with_indices_backward(
906 x: FloatTensor<B>,
907 kernel_size: [usize; 2],
908 stride: [usize; 2],
909 padding: [usize; 2],
910 dilation: [usize; 2],
911 ceil_mode: bool,
912 output_grad: FloatTensor<B>,
913 indices: IntTensor<B>,
914 ) -> MaxPool2dBackward<B>;
915
916 fn interpolate(
922 x: FloatTensor<B>,
923 output_size: [usize; 2],
924 options: InterpolateOptions,
925 ) -> FloatTensor<B>;
926
927 fn interpolate_backward(
929 x: FloatTensor<B>,
930 grad: FloatTensor<B>,
931 output_size: [usize; 2],
932 options: InterpolateOptions,
933 ) -> FloatTensor<B>;
934
935 fn attention(
957 query: FloatTensor<B>,
958 key: FloatTensor<B>,
959 value: FloatTensor<B>,
960 mask: Option<BoolTensor<B>>,
961 attn_bias: Option<FloatTensor<B>>,
962 options: AttentionModuleOptions,
963 ) -> FloatTensor<B>;
964
965 fn layer_norm(
981 tensor: FloatTensor<B>,
982 gamma: FloatTensor<B>,
983 beta: Option<FloatTensor<B>>,
984 epsilon: f64,
985 ) -> FloatTensor<B> {
986 let shape = tensor.shape();
987 let rank = shape.num_dims();
988 let last_dim = rank - 1;
989 let d_model = shape[last_dim];
990
991 let mean = B::float_mean_dim(tensor.clone(), last_dim);
992 let centered = B::float_sub(tensor, mean);
993 let var = B::float_mean_dim(B::float_mul(centered.clone(), centered.clone()), last_dim);
994 let denom = B::float_sqrt(B::float_add_scalar(var, epsilon.into()));
995 let normalized = B::float_div(centered, denom);
996
997 let broadcast_dims: alloc::vec::Vec<usize> = (0..rank)
998 .map(|i| if i == last_dim { d_model } else { 1 })
999 .collect();
1000 let gamma_b = B::float_reshape(gamma, Shape::from(broadcast_dims.clone()));
1001 let scaled = B::float_mul(normalized, gamma_b);
1002
1003 match beta {
1004 Some(beta) => {
1005 let beta_b = B::float_reshape(beta, Shape::from(broadcast_dims));
1006 B::float_add(scaled, beta_b)
1007 }
1008 None => scaled,
1009 }
1010 }
1011
1012 fn ctc_loss(
1029 log_probs: FloatTensor<B>,
1030 targets: IntTensor<B>,
1031 input_lengths: IntTensor<B>,
1032 target_lengths: IntTensor<B>,
1033 blank: usize,
1034 ) -> FloatTensor<B> {
1035 ctc::ctc_loss_default::<B>(log_probs, targets, input_lengths, target_lengths, blank)
1036 }
1037
1038 fn has_ctc_loss_backward() -> bool {
1050 false
1051 }
1052
1053 fn ctc_loss_backward(
1073 _log_probs: FloatTensor<B>,
1074 _targets: IntTensor<B>,
1075 _input_lengths: IntTensor<B>,
1076 _target_lengths: IntTensor<B>,
1077 _grad_loss: FloatTensor<B>,
1078 _blank: usize,
1079 ) -> FloatTensor<B> {
1080 unreachable!(
1081 "ctc_loss_backward called on a backend whose has_ctc_loss_backward() returns false"
1082 )
1083 }
1084}