Skip to main content

ruda_tensor/ops/modules/
base.rs

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/// LayerNorm output and saved statistics used by its native backward operation.
8pub struct LayerNormOutput<B: Backend> {
9    /// Affine-normalized output.
10    pub output: FloatTensor<B>,
11    /// Per-row mean.
12    pub mean: FloatTensor<B>,
13    /// Per-row reciprocal standard deviation, including epsilon.
14    pub rstd: FloatTensor<B>,
15}
16
17/// First-order LayerNorm gradients.
18pub struct LayerNormBackward<B: Backend> {
19    /// Input gradient.
20    pub input: FloatTensor<B>,
21    /// Scale gradient, summed over all leading dimensions.
22    pub weight: FloatTensor<B>,
23    /// Bias gradient, summed over all leading dimensions.
24    pub bias: FloatTensor<B>,
25}
26
27/// Gradient computed during the backward pass for each tensor used by [conv2d](ModuleOps::conv2d).
28#[derive(new)]
29pub struct Conv2dBackward<B: Backend> {
30    /// Gradient.
31    pub x_grad: FloatTensor<B>,
32
33    /// Weights gradient.
34    pub weights_grad: FloatTensor<B>,
35
36    /// Bias gradient.
37    pub bias_grad: Option<FloatTensor<B>>,
38}
39
40/// Gradient computed during the backward pass for each tensor used by [deform_conv2d](ModuleOps::deform_conv2d).
41#[derive(new)]
42pub struct DeformConv2dBackward<B: Backend> {
43    /// Gradient.
44    pub x_grad: FloatTensor<B>,
45
46    /// Offset gradient.
47    pub offset_grad: FloatTensor<B>,
48
49    /// Weights gradient.
50    pub weight_grad: FloatTensor<B>,
51
52    /// Mask gradient.
53    pub mask_grad: Option<FloatTensor<B>>,
54
55    /// Bias gradient.
56    pub bias_grad: Option<FloatTensor<B>>,
57}
58
59/// Gradient computed during the backward pass for each tensor used by [conv3d](ModuleOps::conv3d).
60#[derive(new)]
61pub struct Conv3dBackward<B: Backend> {
62    /// Gradient.
63    pub x_grad: FloatTensor<B>,
64
65    /// Weights gradient.
66    pub weights_grad: FloatTensor<B>,
67
68    /// Bias gradient.
69    pub bias_grad: Option<FloatTensor<B>>,
70}
71
72/// Gradient computed during the backward pass for each tensor used by [max_pool1d](ModuleOps::max_pool1d).
73#[derive(new)]
74pub struct MaxPool1dBackward<B: Backend> {
75    /// Gradient.
76    pub x_grad: FloatTensor<B>,
77}
78
79/// Results from [max_pool1d](ModuleOps::max_pool1d_with_indices).
80#[derive(new)]
81pub struct MaxPool1dWithIndices<B: Backend> {
82    /// The output tensor.
83    pub output: FloatTensor<B>,
84
85    /// The indices tensor.
86    pub indices: IntTensor<B>,
87}
88
89/// Gradient computed during the backward pass for each tensor used by [max_pool2d](ModuleOps::max_pool2d).
90#[derive(new)]
91pub struct MaxPool2dBackward<B: Backend> {
92    /// Gradient.
93    pub x_grad: FloatTensor<B>,
94}
95
96/// Results from [max_pool2d](ModuleOps::max_pool2d_with_indices).
97#[derive(new)]
98pub struct MaxPool2dWithIndices<B: Backend> {
99    /// The output tensor.
100    pub output: FloatTensor<B>,
101
102    /// The indices tensor.
103    pub indices: IntTensor<B>,
104}
105
106pub use ruda_core::tensor::spatial::{ConvOptions, PaddedConvOptions, DeformConvOptions, ConvTransposeOptions, UnfoldOptions};
107
108pub use ruda_core::tensor::spatial::{InterpolateMode, InterpolateOptions};
109
110pub use ruda_core::tensor::spatial::{GridSampleOptions, GridSamplePaddingMode};
111
112/// Padding mode for tensor pad operations.
113///
114/// Defines how values are filled when padding a tensor beyond its original boundaries.
115/// Padding can be applied to any dimension of a tensor.
116///
117/// # Modes
118///
119/// - [`Constant`](PadMode::Constant): Fill with a specified value (default: 0.0)
120/// - [`Reflect`](PadMode::Reflect): Mirror values at boundary, excluding edge (requires padding < dim_size)
121/// - [`Edge`](PadMode::Edge): Replicate boundary values
122#[derive(Debug, Clone, Copy, PartialEq, serde::Deserialize, serde::Serialize)]
123pub enum PadMode {
124    /// Fill padded regions with a constant value.
125    ///
126    /// # Example
127    /// For tensor `[1, 2, 3]` with padding 2 on the left and value 0:
128    /// Result: `[0, 0, 1, 2, 3]`
129    Constant(f32),
130
131    /// Reflect values at the boundary, excluding the edge value.
132    ///
133    /// Padding must be less than the dimension size (i.e., `padding < dim_size`).
134    ///
135    /// # Example
136    /// For tensor `[1, 2, 3, 4]` with padding 2 on the left:
137    /// Result: `[3, 2, 1, 2, 3, 4]` (reflects from index 1, not 0)
138    Reflect,
139
140    /// Replicate the edge values.
141    ///
142    /// # Example
143    /// For tensor `[1, 2, 3, 4]` with padding 2 on the left:
144    /// Result: `[1, 1, 1, 2, 3, 4]`
145    Edge,
146}
147
148impl Default for PadMode {
149    fn default() -> Self {
150        PadMode::Constant(0.0)
151    }
152}
153
154impl<E: ElementConversion> From<E> for PadMode {
155    fn from(value: E) -> Self {
156        PadMode::Constant(value.elem())
157    }
158}
159
160/// Gradient computed during the backward pass for each tensor used by [interpolate](ModuleOps::interpolate).
161#[derive(new)]
162pub struct InterpolateBackward<B: Backend> {
163    /// Gradient.
164    pub x_grad: FloatTensor<B>,
165}
166
167pub use ruda_core::tensor::spatial::AttentionModuleOptions;
168
169/// Module operations trait.
170pub trait ModuleOps<B: Backend> {
171    /// Embedding operation.
172    ///
173    /// # Arguments
174    ///
175    /// * `weights` - The embedding weights.
176    /// * `indices` - The indices tensor.
177    ///
178    /// # Returns
179    ///
180    /// The output tensor.
181    fn embedding(weights: FloatTensor<B>, indices: IntTensor<B>) -> FloatTensor<B> {
182        let [batch_size, seq_length] = indices.shape().dims();
183        let [_, d_model] = weights.shape().dims();
184
185        let indices = B::int_reshape(indices, Shape::new([batch_size * seq_length]));
186        let output = B::float_select(weights, 0, indices);
187
188        B::float_reshape(output, Shape::new([batch_size, seq_length, d_model]))
189    }
190
191    /// Embedding backward operation.
192    ///
193    /// # Arguments
194    ///
195    /// * `weights` - The embedding weights.
196    /// * `output_grad` - The output gradient.
197    /// * `indices` - The indices tensor.
198    ///
199    /// # Returns
200    ///
201    /// The gradient.
202    fn embedding_backward(
203        weights: FloatTensor<B>,
204        output_grad: FloatTensor<B>,
205        indices: IntTensor<B>,
206    ) -> FloatTensor<B> {
207        let [batch_size, seq_length] = indices.shape().dims();
208        let [n_embeddings, d_model] = weights.shape().dims();
209        let device = B::float_device(&weights);
210        let dtype = output_grad.dtype();
211
212        let indices = B::int_reshape(indices, Shape::new([batch_size * seq_length]));
213        let output_grad =
214            B::float_reshape(output_grad, Shape::new([batch_size * seq_length, d_model]));
215        let grad = B::float_zeros(Shape::new([n_embeddings, d_model]), &device, dtype.into());
216
217        B::float_select_add(grad, 0, indices, output_grad)
218    }
219
220    /// Linear transformation.
221    ///
222    /// # Shapes
223    ///
224    /// x:      `[..., d_input]`,
225    /// weight: `[d_input, d_output]`,
226    /// bias:   `[d_output]`,
227    fn linear(
228        x: FloatTensor<B>,
229        weight: FloatTensor<B>,
230        bias: Option<FloatTensor<B>>,
231    ) -> FloatTensor<B> {
232        linear::linear::<B>(x, weight, bias)
233    }
234    /// Backward pass for [linear](ModuleOps::linear), returning the gradient for `x`.
235    fn linear_x_backward(weight: FloatTensor<B>, output_grad: FloatTensor<B>) -> FloatTensor<B> {
236        linear::linear_x_backward::<B>(weight, output_grad)
237    }
238    /// Backward pass for [linear](ModuleOps::linear), returning the gradient for `weight`.
239    fn linear_weight_backward(x: FloatTensor<B>, output_grad: FloatTensor<B>) -> FloatTensor<B> {
240        linear::linear_weight_backward::<B>(x, output_grad)
241    }
242    /// Backward pass for [linear](ModuleOps::linear), returning the gradient for `bias`.
243    fn linear_bias_backward(output_grad: FloatTensor<B>) -> FloatTensor<B> {
244        linear::linear_bias_backward::<B>(output_grad)
245    }
246
247    /// One dimensional convolution.
248    ///
249    /// # Shapes
250    ///
251    /// x:      `[batch_size, channels_in, length]`,
252    /// weight: `[channels_out, channels_in, kernel_size]`,
253    /// bias:   `[channels_out]`,
254    fn conv1d(
255        x: FloatTensor<B>,
256        weight: FloatTensor<B>,
257        bias: Option<FloatTensor<B>>,
258        options: ConvOptions<1>,
259    ) -> FloatTensor<B> {
260        conv::conv1d_from_conv2d::<B>(x, weight, bias, options)
261    }
262    /// Backward pass for the [conv1d](ModuleOps::conv1d) operation, returning the gradient for `x`.
263    fn conv1d_x_backward(
264        x: FloatTensor<B>,
265        weight: FloatTensor<B>,
266        output_grad: FloatTensor<B>,
267        options: ConvOptions<1>,
268    ) -> FloatTensor<B> {
269        conv::conv1d_x_backward::<B>(x, weight, output_grad, options)
270    }
271    /// Backward pass for the [conv1d](ModuleOps::conv1d) operation, returning the gradient for `weight`.
272    fn conv1d_weight_backward(
273        x: FloatTensor<B>,
274        weight: FloatTensor<B>,
275        output_grad: FloatTensor<B>,
276        options: ConvOptions<1>,
277    ) -> FloatTensor<B> {
278        conv::conv1d_weight_backward::<B>(x, weight, output_grad, options)
279    }
280    /// Backward pass for the [conv1d](ModuleOps::conv1d) operation, returning the gradient for `bias`.
281    fn conv1d_bias_backward(
282        x: FloatTensor<B>,
283        bias: FloatTensor<B>,
284        output_grad: FloatTensor<B>,
285    ) -> FloatTensor<B> {
286        conv::conv1d_bias_backward::<B>(x, bias, output_grad)
287    }
288    /// Two dimensional convolution.
289    ///
290    /// # Shapes
291    ///
292    /// x:      `[batch_size, channels_in, height, width]`,
293    /// weight: `[channels_out, channels_in, kernel_size_1, kernel_size_2]`,
294    /// bias:   `[channels_out]`,
295    fn conv2d(
296        x: FloatTensor<B>,
297        weight: FloatTensor<B>,
298        bias: Option<FloatTensor<B>>,
299        options: ConvOptions<2>,
300    ) -> FloatTensor<B>;
301    /// Backward pass for the [conv2d](ModuleOps::conv2d) operation, returning the gradient for `x`.
302    fn conv2d_x_backward(
303        x: FloatTensor<B>,
304        weight: FloatTensor<B>,
305        output_grad: FloatTensor<B>,
306        options: ConvOptions<2>,
307    ) -> FloatTensor<B> {
308        conv::conv2d_x_backward::<B>(x, weight, output_grad, options)
309    }
310    /// Backward pass for the [conv2d](ModuleOps::conv2d) operation, returning the gradient for `weight`.
311    fn conv2d_weight_backward(
312        x: FloatTensor<B>,
313        weight: FloatTensor<B>,
314        output_grad: FloatTensor<B>,
315        options: ConvOptions<2>,
316    ) -> FloatTensor<B> {
317        conv::conv2d_weight_backward::<B>(x, weight, output_grad, options)
318    }
319    /// Backward pass for the [conv2d](ModuleOps::conv2d) operation, returning the gradient for `bias`.
320    fn conv2d_bias_backward(
321        x: FloatTensor<B>,
322        bias: FloatTensor<B>,
323        output_grad: FloatTensor<B>,
324    ) -> FloatTensor<B> {
325        conv::conv2d_bias_backward::<B>(x, bias, output_grad)
326    }
327
328    /// Two dimensional deformable convolution.
329    ///
330    /// # Shapes
331    ///
332    /// x:      `[batch_size, channels_in, height, width]`,
333    /// weight: `[channels_out, channels_in, kernel_size_1, kernel_size_2]`,
334    /// bias:   `[channels_out]`,
335    fn deform_conv2d(
336        x: FloatTensor<B>,
337        offset: FloatTensor<B>,
338        weight: FloatTensor<B>,
339        mask: Option<FloatTensor<B>>,
340        bias: Option<FloatTensor<B>>,
341        options: DeformConvOptions<2>,
342    ) -> FloatTensor<B>;
343    /// Backward pass for the [deform_conv2d](ModuleOps::deform_conv2d) operation.
344    fn deform_conv2d_backward(
345        x: FloatTensor<B>,
346        offset: FloatTensor<B>,
347        weight: FloatTensor<B>,
348        mask: Option<FloatTensor<B>>,
349        bias: Option<FloatTensor<B>>,
350        output_grad: FloatTensor<B>,
351        options: DeformConvOptions<2>,
352    ) -> DeformConv2dBackward<B>;
353
354    /// Three dimensional convolution.
355    ///
356    /// # Shapes
357    ///
358    /// x:      `[batch_size, channels_in, depth, height, width]`,
359    /// weight: `[channels_out, channels_in, kernel_size_1, kernel_size_2, kernel_size_3]`,
360    /// bias:   `[channels_out]`,
361    fn conv3d(
362        x: FloatTensor<B>,
363        weight: FloatTensor<B>,
364        bias: Option<FloatTensor<B>>,
365        options: ConvOptions<3>,
366    ) -> FloatTensor<B>;
367    /// Backward pass for the [conv3d](ModuleOps::conv3d) operation, returning the gradient for `x`.
368    fn conv3d_x_backward(
369        x: FloatTensor<B>,
370        weight: FloatTensor<B>,
371        output_grad: FloatTensor<B>,
372        options: ConvOptions<3>,
373    ) -> FloatTensor<B> {
374        conv::conv3d_x_backward::<B>(x, weight, output_grad, options)
375    }
376    /// Backward pass for the [conv3d](ModuleOps::conv3d) operation, returning the gradient for `weight`.
377    fn conv3d_weight_backward(
378        x: FloatTensor<B>,
379        weight: FloatTensor<B>,
380        output_grad: FloatTensor<B>,
381        options: ConvOptions<3>,
382    ) -> FloatTensor<B> {
383        conv::conv3d_weight_backward::<B>(x, weight, output_grad, options)
384    }
385    /// Backward pass for the [conv3d](ModuleOps::conv3d) operation, returning the gradient for `bias`.
386    fn conv3d_bias_backward(
387        x: FloatTensor<B>,
388        bias: FloatTensor<B>,
389        output_grad: FloatTensor<B>,
390    ) -> FloatTensor<B> {
391        conv::conv3d_bias_backward::<B>(x, bias, output_grad)
392    }
393    /// One dimensional transposed convolution.
394    ///
395    /// # Shapes
396    ///
397    /// x:      `[batch_size, channels_in, length]`,
398    /// weight: `[channels_in, channels_out, length]`,
399    /// bias:   `[channels_out]`,
400    fn conv_transpose1d(
401        x: FloatTensor<B>,
402        weight: FloatTensor<B>,
403        bias: Option<FloatTensor<B>>,
404        options: ConvTransposeOptions<1>,
405    ) -> FloatTensor<B> {
406        conv::conv_transpose1d_from_conv_transpose2d::<B>(x, weight, bias, options)
407    }
408    /// Backward pass for the [conv transpose 1d](ModuleOps::conv_transpose1d) operation, returning the gradient for `x`.
409    fn conv_transpose1d_x_backward(
410        weight: FloatTensor<B>,
411        output_grad: FloatTensor<B>,
412        options: ConvTransposeOptions<1>,
413    ) -> FloatTensor<B> {
414        conv::conv_transpose1d_x_backward::<B>(weight, output_grad, options)
415    }
416    /// Backward pass for the [conv transpose 1d](ModuleOps::conv_transpose1d) operation, returning the gradient for `weight`.
417    fn conv_transpose1d_weight_backward(
418        x: FloatTensor<B>,
419        weight: FloatTensor<B>,
420        output_grad: FloatTensor<B>,
421        options: ConvTransposeOptions<1>,
422    ) -> FloatTensor<B> {
423        conv::conv_transpose1d_weight_backward::<B>(x, weight, output_grad, options)
424    }
425    /// Backward pass for the [conv transpose 1d](ModuleOps::conv_transpose1d) operation, returning the gradient for `bias`.
426    fn conv_transpose1d_bias_backward(
427        x: FloatTensor<B>,
428        bias: FloatTensor<B>,
429        output_grad: FloatTensor<B>,
430    ) -> FloatTensor<B> {
431        conv::conv_transpose1d_bias_backward::<B>(x, bias, output_grad)
432    }
433
434    /// Two dimensional transposed convolution.
435    ///
436    /// # Shapes
437    ///
438    /// x:      `[batch_size, channels_in, height, width]`,
439    /// weight: `[channels_in, channels_out, kernel_size_1, kernel_size_2]`,
440    /// bias:   `[channels_out]`,
441    fn conv_transpose2d(
442        x: FloatTensor<B>,
443        weight: FloatTensor<B>,
444        bias: Option<FloatTensor<B>>,
445        options: ConvTransposeOptions<2>,
446    ) -> FloatTensor<B>;
447    /// Backward pass for the [conv transpose 2d](ModuleOps::conv_transpose2d) operation, returning the gradient for `x`.
448    fn conv_transpose2d_x_backward(
449        weight: FloatTensor<B>,
450        output_grad: FloatTensor<B>,
451        options: ConvTransposeOptions<2>,
452    ) -> FloatTensor<B> {
453        conv::conv_transpose2d_x_backward::<B>(weight, output_grad, options)
454    }
455    /// Backward pass for the [conv transpose 2d](ModuleOps::conv_transpose2d) operation, returning the gradient for `weight`.
456    fn conv_transpose2d_weight_backward(
457        x: FloatTensor<B>,
458        weight: FloatTensor<B>,
459        output_grad: FloatTensor<B>,
460        options: ConvTransposeOptions<2>,
461    ) -> FloatTensor<B> {
462        conv::conv_transpose2d_weight_backward::<B>(x, weight, output_grad, options)
463    }
464    /// Backward pass for the [conv transpose 2d](ModuleOps::conv_transpose2d) operation, returning the gradient for `bias`.
465    fn conv_transpose2d_bias_backward(
466        x: FloatTensor<B>,
467        bias: FloatTensor<B>,
468        output_grad: FloatTensor<B>,
469    ) -> FloatTensor<B> {
470        conv::conv_transpose2d_bias_backward::<B>(x, bias, output_grad)
471    }
472
473    /// Three dimensional transposed convolution.
474    ///
475    /// # Shapes
476    ///
477    /// x:      `[batch_size, channels_in, height, width]`,
478    /// weight: `[channels_in, channels_out, kernel_size_1, kernel_size_2, kernel_size_3]`,
479    /// bias:   `[channels_out]`,
480    fn conv_transpose3d(
481        x: FloatTensor<B>,
482        weight: FloatTensor<B>,
483        bias: Option<FloatTensor<B>>,
484        options: ConvTransposeOptions<3>,
485    ) -> FloatTensor<B>;
486    /// Backward pass for the [conv transpose 3d](ModuleOps::conv_transpose3d) operation, returning the gradient for `x`.
487    fn conv_transpose3d_x_backward(
488        weight: FloatTensor<B>,
489        output_grad: FloatTensor<B>,
490        options: ConvTransposeOptions<3>,
491    ) -> FloatTensor<B> {
492        conv::conv_transpose3d_x_backward::<B>(weight, output_grad, options)
493    }
494    /// Backward pass for the [conv transpose 3d](ModuleOps::conv_transpose3d) operation, returning the gradient for `weight`.
495    fn conv_transpose3d_weight_backward(
496        x: FloatTensor<B>,
497        weight: FloatTensor<B>,
498        output_grad: FloatTensor<B>,
499        options: ConvTransposeOptions<3>,
500    ) -> FloatTensor<B> {
501        conv::conv_transpose3d_weight_backward::<B>(x, weight, output_grad, options)
502    }
503    /// Backward pass for the [conv transpose 3d](ModuleOps::conv_transpose3d) operation, returning the gradient for `bias`.
504    fn conv_transpose3d_bias_backward(
505        x: FloatTensor<B>,
506        bias: FloatTensor<B>,
507        output_grad: FloatTensor<B>,
508    ) -> FloatTensor<B> {
509        conv::conv_transpose3d_bias_backward::<B>(x, bias, output_grad)
510    }
511
512    /// Four-dimensional unfolding.
513    ///
514    /// # Shapes
515    ///
516    /// * x:      ``[batch_size, channels_in, height, width]``,
517    /// * returns: ``[batch_size, channels_in * kernel_size_1 * kernel_size_2, number of blocks]``,
518    fn unfold4d(
519        x: FloatTensor<B>,
520        kernel_size: [usize; 2],
521        options: UnfoldOptions,
522    ) -> FloatTensor<B> {
523        if options.padding == [0, 0] && options.dilation == [1, 1] {
524            let blocks = B::float_unfold(x, 2, kernel_size[0], options.stride[0]);
525            let blocks = B::float_unfold(blocks, 3, kernel_size[1], options.stride[1]);
526
527            // batch, channels, h_blocks, w_blocks, h_kern, w_kern
528
529            let blocks = B::float_permute(blocks, &[0, 1, 4, 5, 2, 3]);
530            let shape = blocks.shape();
531
532            // batch, channels, h_kern, w_kern, h_blocks, w_blocks
533
534            B::float_reshape(
535                blocks,
536                [
537                    shape[0],
538                    shape[1] * shape[2] * shape[3],
539                    shape[4] * shape[5],
540                ]
541                .into(),
542            )
543        } else {
544            unfold4d_using_conv2d::<B>(x, kernel_size, options)
545        }
546    }
547
548    /// One dimensional avg pooling.
549    ///
550    /// # Shapes
551    ///
552    /// x: [batch_size, channels, length],
553    fn avg_pool1d(
554        x: FloatTensor<B>,
555        kernel_size: usize,
556        stride: usize,
557        padding: usize,
558        count_include_pad: bool,
559        ceil_mode: bool,
560    ) -> FloatTensor<B> {
561        pool::avg_pool1d_from_2d::<B>(
562            x,
563            kernel_size,
564            stride,
565            padding,
566            count_include_pad,
567            ceil_mode,
568        )
569    }
570    /// Backward pass for the [avg pooling 1d](ModuleOps::avg_pool1d) operation.
571    fn avg_pool1d_backward(
572        x: FloatTensor<B>,
573        grad: FloatTensor<B>,
574        kernel_size: usize,
575        stride: usize,
576        padding: usize,
577        count_include_pad: bool,
578        ceil_mode: bool,
579    ) -> FloatTensor<B> {
580        pool::avg_pool1d_backward_from_2d::<B>(
581            x,
582            grad,
583            kernel_size,
584            stride,
585            padding,
586            count_include_pad,
587            ceil_mode,
588        )
589    }
590    /// Two dimensional avg pooling.
591    ///
592    /// # Shapes
593    ///
594    /// x: [batch_size, channels, height, width],
595    fn avg_pool2d(
596        x: FloatTensor<B>,
597        kernel_size: [usize; 2],
598        stride: [usize; 2],
599        padding: [usize; 2],
600        count_include_pad: bool,
601        ceil_mode: bool,
602    ) -> FloatTensor<B>;
603    /// Backward pass for the [avg pooling 2d](ModuleOps::avg_pool2d) operation.
604    fn avg_pool2d_backward(
605        x: FloatTensor<B>,
606        grad: FloatTensor<B>,
607        kernel_size: [usize; 2],
608        stride: [usize; 2],
609        padding: [usize; 2],
610        count_include_pad: bool,
611        ceil_mode: bool,
612    ) -> FloatTensor<B>;
613    /// Two dimensional adaptive avg pooling.
614    ///
615    /// # Shapes
616    ///
617    /// x: [batch_size, channels, height, width],
618    fn adaptive_avg_pool2d(x: FloatTensor<B>, output_size: [usize; 2]) -> FloatTensor<B>;
619    /// Backward pass for the [adaptive avg pooling 2d](ModuleOps::adaptive_avg_pool2d) operation.
620    fn adaptive_avg_pool2d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B>;
621    /// One dimensional adaptive avg pooling.
622    ///
623    /// # Shapes
624    ///
625    /// x: [batch_size, channels, length],
626    fn adaptive_avg_pool1d(x: FloatTensor<B>, output_size: usize) -> FloatTensor<B> {
627        pool::adaptive_avg_pool1d_from_2d::<B>(x, output_size)
628    }
629    /// Backward pass for the [adaptive avg pooling 1d](ModuleOps::adaptive_avg_pool1d) operation.
630    fn adaptive_avg_pool1d_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
631        pool::adaptive_avg_pool1d_backward_from_2d::<B>(x, grad)
632    }
633    /// One dimensional max pooling.
634    ///
635    /// # Shapes
636    ///
637    /// x: [batch_size, channels, length],
638    fn max_pool1d(
639        x: FloatTensor<B>,
640        kernel_size: usize,
641        stride: usize,
642        padding: usize,
643        dilation: usize,
644        ceil_mode: bool,
645    ) -> FloatTensor<B> {
646        pool::max_pool1d_from_2d::<B>(x, kernel_size, stride, padding, dilation, ceil_mode)
647    }
648
649    /// One dimensional max pooling with indices.
650    ///
651    /// # Shapes
652    ///
653    /// x: [batch_size, channels, height, width],
654    fn max_pool1d_with_indices(
655        x: FloatTensor<B>,
656        kernel_size: usize,
657        stride: usize,
658        padding: usize,
659        dilation: usize,
660        ceil_mode: bool,
661    ) -> MaxPool1dWithIndices<B> {
662        pool::max_pool1d_with_indices_from_2d::<B>(
663            x,
664            kernel_size,
665            stride,
666            padding,
667            dilation,
668            ceil_mode,
669        )
670    }
671    /// Backward pass for the [max pooling 1d](ModuleOps::max_pool1d_with_indices) operation.
672    #[allow(clippy::too_many_arguments)]
673    fn max_pool1d_with_indices_backward(
674        x: FloatTensor<B>,
675        kernel_size: usize,
676        stride: usize,
677        padding: usize,
678        dilation: usize,
679        ceil_mode: bool,
680        output_grad: FloatTensor<B>,
681        indices: IntTensor<B>,
682    ) -> MaxPool1dBackward<B> {
683        pool::max_pool1d_with_indices_backward_from_2d::<B>(
684            x,
685            kernel_size,
686            stride,
687            padding,
688            dilation,
689            ceil_mode,
690            output_grad,
691            indices,
692        )
693    }
694
695    /// Two dimensional max pooling.
696    ///
697    /// # Shapes
698    ///
699    /// x: [batch_size, channels, height, width],
700    fn max_pool2d(
701        x: FloatTensor<B>,
702        kernel_size: [usize; 2],
703        stride: [usize; 2],
704        padding: [usize; 2],
705        dilation: [usize; 2],
706        ceil_mode: bool,
707    ) -> FloatTensor<B>;
708
709    /// Two dimensional max pooling with indices.
710    ///
711    /// # Shapes
712    ///
713    /// x: [batch_size, channels, height, width],
714    fn max_pool2d_with_indices(
715        x: FloatTensor<B>,
716        kernel_size: [usize; 2],
717        stride: [usize; 2],
718        padding: [usize; 2],
719        dilation: [usize; 2],
720        ceil_mode: bool,
721    ) -> MaxPool2dWithIndices<B>;
722    /// Backward pass for the [max pooling 2d](ModuleOps::max_pool2d_with_indices) operation.
723    #[allow(clippy::too_many_arguments)]
724    fn max_pool2d_with_indices_backward(
725        x: FloatTensor<B>,
726        kernel_size: [usize; 2],
727        stride: [usize; 2],
728        padding: [usize; 2],
729        dilation: [usize; 2],
730        ceil_mode: bool,
731        output_grad: FloatTensor<B>,
732        indices: IntTensor<B>,
733    ) -> MaxPool2dBackward<B>;
734
735    /// Down/up samples the input.
736    ///
737    /// # Shapes
738    ///
739    /// x: `[batch_size, channels, height, width]`,
740    fn interpolate(
741        x: FloatTensor<B>,
742        output_size: [usize; 2],
743        options: InterpolateOptions,
744    ) -> FloatTensor<B>;
745
746    /// Backward pass for the [interpolate](ModuleOps::interpolate) operation.
747    fn interpolate_backward(
748        x: FloatTensor<B>,
749        grad: FloatTensor<B>,
750        output_size: [usize; 2],
751        options: InterpolateOptions,
752    ) -> FloatTensor<B>;
753
754    /// Computes scaled dot-product attention: softmax(QKᵗ * scale) · V,
755    /// where scale defaults to 1/sqrt(head_dim). Optionally applies masking,
756    /// additive bias, causal masking, and softcap to the attention scores.
757    ///
758    /// # Arguments
759    /// - `query`: Query tensor of shape `[batch_size, num_heads, seq_len_q, head_dim]`
760    /// - `key`: Key tensor of shape `[batch_size, num_heads, seq_len_k, head_dim]`
761    /// - `value`: Value tensor of shape `[batch_size, num_heads, seq_len_k, val_dim]`
762    /// - `mask`: Optional boolean mask of shape `[batch_size, num_heads, seq_len_q, seq_len_k]`,
763    ///   where `true` indicates positions to mask (i.e. set to -inf before softmax).
764    /// - `attn_bias`: Optional float tensor of shape `[batch_size, num_heads, seq_len_q, seq_len_k]`
765    ///   added to the attention scores before softmax (e.g. ALiBi, relative position biases).
766    /// - `options`: Additional attention options (custom scale, softcap, causal masking).
767    ///
768    /// # Returns
769    /// A tensor of shape `[batch_size, num_heads, seq_len_q, val_dim]`
770    /// representing the attended context per head.
771    ///
772    /// # Note
773    /// This implementation does not support dropout and is intended for inference or
774    /// use cases where dropout is not needed.
775    fn attention(
776        query: FloatTensor<B>,
777        key: FloatTensor<B>,
778        value: FloatTensor<B>,
779        mask: Option<BoolTensor<B>>,
780        attn_bias: Option<FloatTensor<B>>,
781        options: AttentionModuleOptions,
782    ) -> FloatTensor<B>;
783
784    /// Applies Layer Normalization over the last dimension of the input tensor.
785    ///
786    /// Computes `(x - mean) / sqrt(var + epsilon) * gamma + beta`, where `mean` and
787    /// (biased) `var` are reduced over the last axis.
788    ///
789    /// # Arguments
790    ///
791    /// * `tensor` - Input tensor of shape `[..., d_model]`.
792    /// * `gamma` - Scale tensor of shape `[d_model]`.
793    /// * `beta` - Optional bias tensor of shape `[d_model]`.
794    /// * `epsilon` - Numerical stability term added to the variance before the square root.
795    ///
796    /// # Returns
797    ///
798    /// A tensor with the same shape as `tensor`.
799    fn layer_norm(
800        tensor: FloatTensor<B>,
801        gamma: FloatTensor<B>,
802        beta: Option<FloatTensor<B>>,
803        epsilon: f64,
804    ) -> FloatTensor<B> {
805        Self::layer_norm_default(tensor, gamma, beta, epsilon)
806    }
807
808    /// Whether native forward statistics and complete first-order backward are available.
809    fn has_layer_norm_backward() -> bool { false }
810
811    /// Native forward with statistics retained for backward.
812    fn layer_norm_with_stats(
813        _tensor: FloatTensor<B>, _gamma: FloatTensor<B>,
814        _beta: Option<FloatTensor<B>>, _epsilon: f64,
815    ) -> LayerNormOutput<B> {
816        unimplemented!("native LayerNorm statistics unavailable")
817    }
818
819    /// Native backward using the exact statistics returned by forward.
820    fn layer_norm_backward(
821        _tensor: FloatTensor<B>, _gamma: FloatTensor<B>, _grad: FloatTensor<B>,
822        _mean: FloatTensor<B>, _rstd: FloatTensor<B>,
823    ) -> LayerNormBackward<B> {
824        unimplemented!("native LayerNorm backward unavailable")
825    }
826
827    /// Differentiable primitive composition for backends without native LayerNorm backward.
828    fn layer_norm_default(
829        tensor: FloatTensor<B>, gamma: FloatTensor<B>,
830        beta: Option<FloatTensor<B>>, epsilon: f64,
831    ) -> FloatTensor<B> {
832        let shape = tensor.shape();
833        let rank = shape.num_dims();
834        let last_dim = rank - 1;
835        let d_model = shape[last_dim];
836
837        let mean = B::float_mean_dim(tensor.clone(), last_dim);
838        let centered = B::float_sub(tensor, mean);
839        let var = B::float_mean_dim(B::float_mul(centered.clone(), centered.clone()), last_dim);
840        let denom = B::float_sqrt(B::float_add_scalar(var, epsilon.into()));
841        let normalized = B::float_div(centered, denom);
842
843        let broadcast_dims: alloc::vec::Vec<usize> = (0..rank)
844            .map(|i| if i == last_dim { d_model } else { 1 })
845            .collect();
846        let gamma_b = B::float_reshape(gamma, Shape::from(broadcast_dims.clone()));
847        let scaled = B::float_mul(normalized, gamma_b);
848
849        match beta {
850            Some(beta) => {
851                let beta_b = B::float_reshape(beta, Shape::from(broadcast_dims));
852                B::float_add(scaled, beta_b)
853            }
854            None => scaled,
855        }
856    }
857
858    /// Computes the Connectionist Temporal Classification (CTC) loss.
859    ///
860    /// Sums over all valid alignments between the input and target sequences
861    /// using the forward (alpha) algorithm.
862    ///
863    /// # Arguments
864    ///
865    /// * `log_probs` - Log-probabilities of shape `[T, N, C]`
866    /// * `targets` - Target label indices of shape `[N, S]`
867    /// * `input_lengths` - Actual input sequence lengths per batch element `[N]`
868    /// * `target_lengths` - Actual target lengths per batch element `[N]`
869    /// * `blank` - Index of the blank label
870    ///
871    /// # Returns
872    ///
873    /// Per-sample loss of shape `[N]`
874    fn ctc_loss(
875        log_probs: FloatTensor<B>,
876        targets: IntTensor<B>,
877        input_lengths: IntTensor<B>,
878        target_lengths: IntTensor<B>,
879        blank: usize,
880    ) -> FloatTensor<B> {
881        ctc::ctc_loss_default::<B>(log_probs, targets, input_lengths, target_lengths, blank)
882    }
883
884    /// Returns `true` if this backend implements [ctc_loss_backward](ModuleOps::ctc_loss_backward)
885    /// natively.
886    ///
887    /// Autodiff queries this flag to decide between two paths:
888    /// - `true`: use the backend's [ctc_loss](ModuleOps::ctc_loss) and
889    ///   [ctc_loss_backward](ModuleOps::ctc_loss_backward) directly.
890    /// - `false`: call [ctc::ctc_loss_default] for the forward pass; autodiff
891    ///   then differentiates through the decomposed tensor ops.
892    ///
893    /// Backends that override `ctc_loss_backward` must also override this to
894    /// return `true`.
895    fn has_ctc_loss_backward() -> bool {
896        false
897    }
898
899    /// Backward pass for [ctc_loss](ModuleOps::ctc_loss): gradient w.r.t. `log_probs`.
900    ///
901    /// Only called when [has_ctc_loss_backward](ModuleOps::has_ctc_loss_backward)
902    /// returns `true`. Backends without a native implementation should leave
903    /// both methods at their defaults; the gradient is computed automatically by
904    /// autodiff against the decomposed [ctc::ctc_loss_default] forward.
905    ///
906    /// # Arguments
907    ///
908    /// * `log_probs` - Log-probabilities of shape `[T, N, C]`
909    /// * `targets` - Target label indices of shape `[N, S]`
910    /// * `input_lengths` - Actual input sequence lengths per batch element `[N]`
911    /// * `target_lengths` - Actual target lengths per batch element `[N]`
912    /// * `grad_loss` - Upstream gradient w.r.t. the per-sample loss `[N]`
913    /// * `blank` - Index of the blank label
914    ///
915    /// # Returns
916    ///
917    /// Gradient w.r.t. `log_probs` of shape `[T, N, C]`
918    fn ctc_loss_backward(
919        _log_probs: FloatTensor<B>,
920        _targets: IntTensor<B>,
921        _input_lengths: IntTensor<B>,
922        _target_lengths: IntTensor<B>,
923        _grad_loss: FloatTensor<B>,
924        _blank: usize,
925    ) -> FloatTensor<B> {
926        unreachable!(
927            "ctc_loss_backward called on a backend whose has_ctc_loss_backward() returns false"
928        )
929    }
930
931    /// Real-valued FFT with optional size parameter.
932    ///
933    /// When `n` is `None`, the signal must be a power of two along `dim`, and the output has
934    /// `signal_len / 2 + 1` frequency bins.
935    ///
936    /// When `n` is `Some(size)`, `size` must also be a power of two. The signal is truncated
937    /// or zero-padded to `size` and the output has `size / 2 + 1` frequency bins. Non-power-
938    /// of-two sizes are currently rejected at the public API boundary; true arbitrary-`n` DFT
939    /// support (Bluestein's algorithm) is tracked as a follow-up.
940    ///
941    /// Returns two tensors: the real part and the imaginary part.
942    fn rfft(
943        signal: FloatTensor<B>,
944        dim: usize,
945        n: Option<usize>,
946    ) -> (FloatTensor<B>, FloatTensor<B>);
947
948    /// Inverse real-valued FFT with optional output size.
949    ///
950    /// When `n` is `None`, the reconstructed signal length `2 * (spectrum_size - 1)` must be
951    /// a power of two.
952    ///
953    /// When `n` is `Some(size)`, `size` must also be a power of two. Output has exactly
954    /// `size` samples.
955    fn irfft(
956        spectrum_re: FloatTensor<B>,
957        spectrum_im: FloatTensor<B>,
958        dim: usize,
959        n: Option<usize>,
960    ) -> FloatTensor<B>;
961}
962
963#[cfg(test)]
964mod tests {
965    use super::*;
966
967    #[test]
968    #[should_panic = "stride must be non-zero"]
969    fn conv_options_stride_zero() {
970        let _opt = ConvOptions::new([0, 1], [0, 0], [1, 1], 1);
971    }
972
973    #[test]
974    #[should_panic = "dilation must be non-zero"]
975    fn conv_options_dilation_zero() {
976        let _opt = ConvOptions::new([1, 1], [0, 0], [0, 0], 1);
977    }
978
979    #[test]
980    #[should_panic = "groups must be non-zero"]
981    fn conv_options_groups_zero() {
982        let _opt = ConvOptions::new([1, 1], [0, 0], [1, 1], 0);
983    }
984
985    #[test]
986    #[should_panic = "stride must be non-zero"]
987    fn conv_transpose_options_stride_zero() {
988        let _opt = ConvTransposeOptions::new([0, 1], [0, 0], [0, 0], [1, 1], 1);
989    }
990
991    #[test]
992    #[should_panic = "dilation must be non-zero"]
993    fn conv_transpose_options_dilation_zero() {
994        let _opt = ConvTransposeOptions::new([1, 1], [0, 0], [0, 0], [0, 0], 1);
995    }
996
997    #[test]
998    #[should_panic = "groups must be non-zero"]
999    fn conv_transpose_options_groups_zero() {
1000        let _opt = ConvTransposeOptions::new([1, 1], [0, 0], [0, 0], [1, 1], 0);
1001    }
1002
1003    #[test]
1004    #[should_panic = "stride must be non-zero"]
1005    fn deform_conv_options_stride_zero() {
1006        let _opt = DeformConvOptions::new([0, 1], [0, 0], [1, 1], 1, 1);
1007    }
1008
1009    #[test]
1010    #[should_panic = "dilation must be non-zero"]
1011    fn deform_conv_options_dilation_zero() {
1012        let _opt = DeformConvOptions::new([1, 1], [0, 0], [0, 0], 1, 1);
1013    }
1014
1015    #[test]
1016    #[should_panic = "weight groups must be non-zero"]
1017    fn deform_conv_options_weights_groups_zero() {
1018        let _opt = DeformConvOptions::new([1, 1], [0, 0], [1, 1], 0, 1);
1019    }
1020
1021    #[test]
1022    #[should_panic = "offset groups must be non-zero"]
1023    fn deform_conv_options_offset_groups_zero() {
1024        let _opt = DeformConvOptions::new([1, 1], [0, 0], [1, 1], 1, 0);
1025    }
1026
1027    #[test]
1028    #[should_panic = "stride must be non-zero"]
1029    fn unfold_options_stride_zero() {
1030        let _opt = UnfoldOptions::new([0, 1], [0, 0], [1, 1]);
1031    }
1032
1033    #[test]
1034    #[should_panic = "dilation must be non-zero"]
1035    fn unfold_options_dilation_zero() {
1036        let _opt = UnfoldOptions::new([1, 1], [0, 0], [0, 0]);
1037    }
1038}