Skip to main content

ruda_tensor/ops/modules/
base.rs

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