Skip to main content

ruda_tensor/ops/modules/
base.rs

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