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}