Skip to main content

ruda_tensor/api/
module.rs

1use crate::api::{
2    Bool, Int, Tensor, TensorPrimitive,
3    backend::Backend,
4    check,
5    check::TensorCheck,
6    ops::{
7        AttentionModuleOptions, ConvOptions, ConvTransposeOptions, InterpolateOptions, PadMode,
8        PaddedConvOptions, UnfoldOptions,
9    },
10};
11
12use super::ops::DeformConvOptions;
13pub use super::spatial_pool::{adaptive_avg_pool3d, avg_pool3d, max_pool3d, max_pool3d_with_indices};
14pub use super::spatial_interpolate::{interpolate1d, interpolate3d};
15pub use super::spatial_pool::{
16    avg_pool1d_padded, avg_pool2d_padded, avg_pool3d_padded,
17    max_pool1d_padded, max_pool2d_padded, max_pool3d_padded,
18    max_pool1d_with_indices_padded, max_pool2d_with_indices_padded,
19    max_pool3d_with_indices_padded,
20};
21
22/// Computes the [CTC loss](crate::api::ops::ModuleOps::ctc_loss).
23///
24/// # Arguments
25///
26/// * `log_probs` - Log-probabilities of shape `[T, N, C]`
27/// * `targets` - Target label indices of shape `[N, S]`
28/// * `input_lengths` - Actual input sequence lengths per batch element `[N]`
29/// * `target_lengths` - Actual target lengths per batch element `[N]`
30/// * `blank` - Index of the blank label
31///
32/// # Returns
33///
34/// Per-sample loss of shape `[N]`
35pub fn ctc_loss<B>(
36    log_probs: Tensor<B, 3>,
37    targets: Tensor<B, 2, Int>,
38    input_lengths: Tensor<B, 1, Int>,
39    target_lengths: Tensor<B, 1, Int>,
40    blank: usize,
41) -> Tensor<B, 1>
42where
43    B: Backend,
44{
45    Tensor::new(TensorPrimitive::Float(B::ctc_loss(
46        log_probs.primitive.tensor(),
47        targets.primitive,
48        input_lengths.primitive,
49        target_lengths.primitive,
50        blank,
51    )))
52}
53
54/// Applies the [embedding module](crate::api::ops::ModuleOps::embedding).
55pub fn embedding<B>(weights: Tensor<B, 2>, indices: Tensor<B, 2, Int>) -> Tensor<B, 3>
56where
57    B: Backend,
58{
59    Tensor::new(TensorPrimitive::Float(B::embedding(
60        weights.primitive.tensor(),
61        indices.primitive,
62    )))
63}
64
65/// Applies a [1D convolution](crate::api::ops::ModuleOps::conv1d).
66///
67/// Accepts [`ConvOptions`] for symmetric padding, or [`PaddedConvOptions`] for
68/// asymmetric padding. When asymmetric padding is specified, an explicit pad
69/// operation is applied before the convolution backend op.
70pub fn conv1d<B>(
71    x: Tensor<B, 3>,
72    weight: Tensor<B, 3>,
73    bias: Option<Tensor<B, 1>>,
74    options: impl Into<PaddedConvOptions<1>>,
75) -> Tensor<B, 3>
76where
77    B: Backend,
78{
79    let padded_options = options.into();
80    check!(TensorCheck::conv(
81        "conv1d",
82        x.dims(),
83        weight.dims(),
84        padded_options.options.groups,
85    ));
86
87    if let Some(padding_end) = padded_options.padding_end {
88        let left = padded_options.options.padding[0];
89        let right = padding_end[0];
90        // For 1D (NCL format), pad the length dimension
91        let padded = x.pad((left, right, 0, 0), PadMode::Constant(0.0));
92        let zero_options = ConvOptions::new(
93            padded_options.options.stride,
94            [0],
95            padded_options.options.dilation,
96            padded_options.options.groups,
97        );
98        Tensor::new(TensorPrimitive::Float(B::conv1d(
99            padded.primitive.tensor(),
100            weight.primitive.tensor(),
101            bias.map(|b| b.primitive.tensor()),
102            zero_options,
103        )))
104    } else {
105        Tensor::new(TensorPrimitive::Float(B::conv1d(
106            x.primitive.tensor(),
107            weight.primitive.tensor(),
108            bias.map(|b| b.primitive.tensor()),
109            padded_options.options,
110        )))
111    }
112}
113
114/// Applies a [2D convolution](crate::api::ops::ModuleOps::conv2d).
115///
116/// Accepts [`ConvOptions`] for symmetric padding, or [`PaddedConvOptions`] for
117/// asymmetric padding. When asymmetric padding is specified, an explicit pad
118/// operation is applied before the convolution backend op.
119pub fn conv2d<B>(
120    x: Tensor<B, 4>,
121    weight: Tensor<B, 4>,
122    bias: Option<Tensor<B, 1>>,
123    options: impl Into<PaddedConvOptions<2>>,
124) -> Tensor<B, 4>
125where
126    B: Backend,
127{
128    let padded_options = options.into();
129    check!(TensorCheck::conv(
130        "conv2d",
131        x.dims(),
132        weight.dims(),
133        padded_options.options.groups,
134    ));
135
136    if let Some(padding_end) = padded_options.padding_end {
137        let top = padded_options.options.padding[0];
138        let left = padded_options.options.padding[1];
139        let bottom = padding_end[0];
140        let right = padding_end[1];
141        // For 2D (NCHW format), pad height and width
142        let padded = x.pad((left, right, top, bottom), PadMode::Constant(0.0));
143        let zero_options = ConvOptions::new(
144            padded_options.options.stride,
145            [0, 0],
146            padded_options.options.dilation,
147            padded_options.options.groups,
148        );
149        Tensor::new(TensorPrimitive::Float(B::conv2d(
150            padded.primitive.tensor(),
151            weight.primitive.tensor(),
152            bias.map(|b| b.primitive.tensor()),
153            zero_options,
154        )))
155    } else {
156        Tensor::new(TensorPrimitive::Float(B::conv2d(
157            x.primitive.tensor(),
158            weight.primitive.tensor(),
159            bias.map(|b| b.primitive.tensor()),
160            padded_options.options,
161        )))
162    }
163}
164
165/// Applies a [3D convolution](crate::api::ops::ModuleOps::conv3d).
166///
167/// Accepts [`ConvOptions`] for symmetric padding, or [`PaddedConvOptions`] for
168/// asymmetric padding. An explicit pad operation handles asymmetric padding
169/// before dispatching to the convolution backend.
170pub fn conv3d<B>(
171    x: Tensor<B, 5>,
172    weight: Tensor<B, 5>,
173    bias: Option<Tensor<B, 1>>,
174    options: impl Into<PaddedConvOptions<3>>,
175) -> Tensor<B, 5>
176where
177    B: Backend,
178{
179    let padded_options = options.into();
180    check!(TensorCheck::conv(
181        "conv3d",
182        x.dims(),
183        weight.dims(),
184        padded_options.options.groups,
185    ));
186
187    let mut options = padded_options.options;
188    let x = if let Some(padding_end) = padded_options.padding_end {
189        let padding: [(usize, usize); 3] =
190            core::array::from_fn(|axis| (options.padding[axis], padding_end[axis]));
191        options.padding = [0; 3];
192        x.pad(padding, PadMode::Constant(0.0))
193    } else {
194        x
195    };
196
197    Tensor::new(TensorPrimitive::Float(B::conv3d(
198        x.primitive.tensor(),
199        weight.primitive.tensor(),
200        bias.map(|b| b.primitive.tensor()),
201        options,
202    )))
203}
204
205/// Applies a [Deformable 2D convolution](crate::api::ops::ModuleOps::deform_conv2d).
206pub fn deform_conv2d<B>(
207    x: Tensor<B, 4>,
208    offset: Tensor<B, 4>,
209    weight: Tensor<B, 4>,
210    mask: Option<Tensor<B, 4>>,
211    bias: Option<Tensor<B, 1>>,
212    options: DeformConvOptions<2>,
213) -> Tensor<B, 4>
214where
215    B: Backend,
216{
217    check!(TensorCheck::conv(
218        "deform_conv2d",
219        x.dims(),
220        weight.dims(),
221        options.weight_groups,
222    ));
223    let [batch, channels, height, width] = x.dims();
224    let [_, _, kernel_height, kernel_width] = weight.dims();
225    assert!(
226        options.offset_groups > 0 && channels.is_multiple_of(options.offset_groups),
227        "deform_conv2d input channels must be divisible by non-zero offset groups"
228    );
229    let [out_height, out_width] = options.output_size(
230        [height, width], [kernel_height, kernel_width],
231    );
232    let mask_channels = options.offset_groups * kernel_height * kernel_width;
233    assert_eq!(
234        offset.dims(), [batch, 2 * mask_channels, out_height, out_width],
235        "deform_conv2d offset shape must match groups, kernel and output"
236    );
237    if let Some(mask) = mask.as_ref() {
238        assert_eq!(
239            mask.dims(), [batch, mask_channels, out_height, out_width],
240            "deform_conv2d mask shape must match groups, kernel and output"
241        );
242    }
243    Tensor::new(TensorPrimitive::Float(B::deform_conv2d(
244        x.primitive.tensor(),
245        offset.primitive.tensor(),
246        weight.primitive.tensor(),
247        mask.map(|m| m.primitive.tensor()),
248        bias.map(|b| b.primitive.tensor()),
249        options,
250    )))
251}
252
253/// Applies a [1D transposed convolution](crate::api::ops::ModuleOps::conv_transpose1d).
254pub fn conv_transpose1d<B>(
255    x: Tensor<B, 3>,
256    weight: Tensor<B, 3>,
257    bias: Option<Tensor<B, 1>>,
258    options: ConvTransposeOptions<1>,
259) -> Tensor<B, 3>
260where
261    B: Backend,
262{
263    check!(TensorCheck::conv_transpose(
264        "conv_transpose1d",
265        x.dims(),
266        weight.dims(),
267    ));
268    Tensor::new(TensorPrimitive::Float(B::conv_transpose1d(
269        x.primitive.tensor(),
270        weight.primitive.tensor(),
271        bias.map(|b| b.primitive.tensor()),
272        options,
273    )))
274}
275
276/// Applies a [2D transposed convolution](crate::api::ops::ModuleOps::conv_transpose2d).
277pub fn conv_transpose2d<B>(
278    x: Tensor<B, 4>,
279    weight: Tensor<B, 4>,
280    bias: Option<Tensor<B, 1>>,
281    options: ConvTransposeOptions<2>,
282) -> Tensor<B, 4>
283where
284    B: Backend,
285{
286    check!(TensorCheck::conv_transpose(
287        "conv_transpose2d",
288        x.dims(),
289        weight.dims(),
290    ));
291    Tensor::new(TensorPrimitive::Float(B::conv_transpose2d(
292        x.primitive.tensor(),
293        weight.primitive.tensor(),
294        bias.map(|b| b.primitive.tensor()),
295        options,
296    )))
297}
298
299/// Applies a 3D transposed convolution](crate::api::ops::ModuleOps::conv_transpose3d).
300pub fn conv_transpose3d<B>(
301    x: Tensor<B, 5>,
302    weight: Tensor<B, 5>,
303    bias: Option<Tensor<B, 1>>,
304    options: ConvTransposeOptions<3>,
305) -> Tensor<B, 5>
306where
307    B: Backend,
308{
309    check!(TensorCheck::conv_transpose(
310        "conv_transpose3d",
311        x.dims(),
312        weight.dims(),
313    ));
314    Tensor::new(TensorPrimitive::Float(B::conv_transpose3d(
315        x.primitive.tensor(),
316        weight.primitive.tensor(),
317        bias.map(|b| b.primitive.tensor()),
318        options,
319    )))
320}
321
322/// Apply a 1D transposed convolution with an explicit output length.
323///
324/// The requested length determines native output padding. The convolution and
325/// its backward pass use the existing backend operation without resizing.
326pub fn conv_transpose1d_with_output_size<B: Backend>(
327    x: Tensor<B, 3>,
328    weight: Tensor<B, 3>,
329    bias: Option<Tensor<B, 1>>,
330    options: ConvTransposeOptions<1>,
331    output_size: usize,
332) -> Tensor<B, 3> {
333    let [_, _, input_length] = x.dims();
334    let [_, _, kernel_length] = weight.dims();
335    let options = options.with_output_size([kernel_length], [input_length], [output_size]);
336    conv_transpose1d(x, weight, bias, options)
337}
338
339/// Apply a 2D transposed convolution with explicit output height and width.
340///
341/// Each spatial axis independently determines native output padding.
342pub fn conv_transpose2d_with_output_size<B: Backend>(
343    x: Tensor<B, 4>,
344    weight: Tensor<B, 4>,
345    bias: Option<Tensor<B, 1>>,
346    options: ConvTransposeOptions<2>,
347    output_size: [usize; 2],
348) -> Tensor<B, 4> {
349    let [_, _, height, width] = x.dims();
350    let [_, _, kernel_height, kernel_width] = weight.dims();
351    let options = options.with_output_size(
352        [kernel_height, kernel_width],
353        [height, width],
354        output_size,
355    );
356    conv_transpose2d(x, weight, bias, options)
357}
358
359/// Apply a 3D transposed convolution with explicit output depth, height and width.
360///
361/// This dispatches the original differentiable convolution operation; neither
362/// the input nor the weights are copied to the host or resampled.
363pub fn conv_transpose3d_with_output_size<B: Backend>(
364    x: Tensor<B, 5>,
365    weight: Tensor<B, 5>,
366    bias: Option<Tensor<B, 1>>,
367    options: ConvTransposeOptions<3>,
368    output_size: [usize; 3],
369) -> Tensor<B, 5> {
370    let [_, _, depth, height, width] = x.dims();
371    let [_, _, kernel_depth, kernel_height, kernel_width] = weight.dims();
372    let options = options.with_output_size(
373        [kernel_depth, kernel_height, kernel_width],
374        [depth, height, width],
375        output_size,
376    );
377    conv_transpose3d(x, weight, bias, options)
378}
379
380/// Applies a [4D to 3D unfold](crate::api::ops::ModuleOps::unfold4d).
381pub fn unfold4d<B>(x: Tensor<B, 4>, kernel_size: [usize; 2], options: UnfoldOptions) -> Tensor<B, 3>
382where
383    B: Backend,
384{
385    Tensor::new(TensorPrimitive::Float(B::unfold4d(
386        x.primitive.tensor(),
387        kernel_size,
388        options,
389    )))
390}
391
392/// Applies a [1D max pooling](crate::api::ops::ModuleOps::max_pool1d).
393pub fn max_pool1d<B>(
394    x: Tensor<B, 3>,
395    kernel_size: usize,
396    stride: usize,
397    padding: usize,
398    dilation: usize,
399    ceil_mode: bool,
400) -> Tensor<B, 3>
401where
402    B: Backend,
403{
404    Tensor::new(TensorPrimitive::Float(B::max_pool1d(
405        x.primitive.tensor(),
406        kernel_size,
407        stride,
408        padding,
409        dilation,
410        ceil_mode,
411    )))
412}
413
414/// Applies a [2D max pooling](crate::api::ops::ModuleOps::max_pool2d).
415pub fn max_pool2d<B>(
416    x: Tensor<B, 4>,
417    kernel_size: [usize; 2],
418    stride: [usize; 2],
419    padding: [usize; 2],
420    dilation: [usize; 2],
421    ceil_mode: bool,
422) -> Tensor<B, 4>
423where
424    B: Backend,
425{
426    Tensor::new(TensorPrimitive::Float(B::max_pool2d(
427        x.primitive.tensor(),
428        kernel_size,
429        stride,
430        padding,
431        dilation,
432        ceil_mode,
433    )))
434}
435
436/// Applies a [2D avg pooling](crate::api::ops::ModuleOps::avg_pool2d).
437pub fn avg_pool2d<B>(
438    x: Tensor<B, 4>,
439    kernel_size: [usize; 2],
440    stride: [usize; 2],
441    padding: [usize; 2],
442    count_include_pad: bool,
443    ceil_mode: bool,
444) -> Tensor<B, 4>
445where
446    B: Backend,
447{
448    Tensor::new(TensorPrimitive::Float(B::avg_pool2d(
449        x.primitive.tensor(),
450        kernel_size,
451        stride,
452        padding,
453        count_include_pad,
454        ceil_mode,
455    )))
456}
457
458/// Applies a [1D avg pooling](crate::api::ops::ModuleOps::avg_pool1d).
459pub fn avg_pool1d<B>(
460    x: Tensor<B, 3>,
461    kernel_size: usize,
462    stride: usize,
463    padding: usize,
464    count_include_pad: bool,
465    ceil_mode: bool,
466) -> Tensor<B, 3>
467where
468    B: Backend,
469{
470    Tensor::new(TensorPrimitive::Float(B::avg_pool1d(
471        x.primitive.tensor(),
472        kernel_size,
473        stride,
474        padding,
475        count_include_pad,
476        ceil_mode,
477    )))
478}
479
480/// Applies a [1D max pooling](crate::api::ops::ModuleOps::max_pool1d).
481pub fn max_pool1d_with_indices<B>(
482    x: Tensor<B, 3>,
483    kernel_size: usize,
484    stride: usize,
485    padding: usize,
486    dilation: usize,
487    ceil_mode: bool,
488) -> (Tensor<B, 3>, Tensor<B, 3, Int>)
489where
490    B: Backend,
491{
492    let output = B::max_pool1d_with_indices(
493        x.primitive.tensor(),
494        kernel_size,
495        stride,
496        padding,
497        dilation,
498        ceil_mode,
499    );
500
501    (
502        Tensor::new(TensorPrimitive::Float(output.output)),
503        Tensor::new(output.indices),
504    )
505}
506
507/// Applies a [2D max pooling with indices](crate::api::ops::ModuleOps::max_pool2d_with_indices).
508pub fn max_pool2d_with_indices<B>(
509    x: Tensor<B, 4>,
510    kernel_size: [usize; 2],
511    stride: [usize; 2],
512    padding: [usize; 2],
513    dilation: [usize; 2],
514    ceil_mode: bool,
515) -> (Tensor<B, 4>, Tensor<B, 4, Int>)
516where
517    B: Backend,
518{
519    let output = B::max_pool2d_with_indices(
520        x.primitive.tensor(),
521        kernel_size,
522        stride,
523        padding,
524        dilation,
525        ceil_mode,
526    );
527
528    (
529        Tensor::new(TensorPrimitive::Float(output.output)),
530        Tensor::new(output.indices),
531    )
532}
533
534/// Applies a [2D adaptive avg pooling](crate::api::ops::ModuleOps::adaptive_avg_pool2d).
535pub fn adaptive_avg_pool2d<B>(x: Tensor<B, 4>, output_size: [usize; 2]) -> Tensor<B, 4>
536where
537    B: Backend,
538{
539    Tensor::new(TensorPrimitive::Float(B::adaptive_avg_pool2d(
540        x.primitive.tensor(),
541        output_size,
542    )))
543}
544
545/// Applies a [1D adaptive avg pooling](crate::api::ops::ModuleOps::adaptive_avg_pool1d).
546pub fn adaptive_avg_pool1d<B>(x: Tensor<B, 3>, output_size: usize) -> Tensor<B, 3>
547where
548    B: Backend,
549{
550    Tensor::new(TensorPrimitive::Float(B::adaptive_avg_pool1d(
551        x.primitive.tensor(),
552        output_size,
553    )))
554}
555
556/// Applies a [2D interpolation](crate::api::ops::ModuleOps::interpolate).
557pub fn interpolate<B>(
558    x: Tensor<B, 4>,
559    output_size: [usize; 2],
560    options: InterpolateOptions,
561) -> Tensor<B, 4>
562where
563    B: Backend,
564{
565    Tensor::new(TensorPrimitive::Float(B::interpolate(
566        x.primitive.tensor(),
567        output_size,
568        options,
569    )))
570}
571
572/// Applies a linear transformation to the input tensor using the given weight and bias.
573///
574/// ```math
575/// y = x @ weight + [bias]
576/// ```
577///
578/// # Arguments:
579///
580/// - `input` is the input tensor, ``[..., d_input]``.
581/// - `weight` is the weight tensor, ``[d_input, d_output]``.
582/// - `bias` is the bias tensor (optional), ``[d_output]``.
583///
584/// # Returns:
585///
586/// The transformed tensor, ``[..., d_output]``.
587///
588/// # Compatibility
589///
590/// This function differs from PyTorch's ``torch.nn.functional.linear`` in that it does not
591/// transpose the weight matrix. In PyTorch, the weight matrix is transposed before
592/// multiplication:
593///
594/// ```math
595/// y = x @ weight^T + [bias]
596/// ```
597pub fn linear<B: Backend, const D: usize>(
598    input: Tensor<B, D>,
599    weight: Tensor<B, 2>,
600    bias: Option<Tensor<B, 1>>,
601) -> Tensor<B, D> {
602    if D == 1 {
603        // Insert and remove an extra batch dimension for the batch matmul to work.
604        let input = input.unsqueeze::<2>();
605        let output = linear(input, weight, bias);
606        return output.squeeze_dim(0);
607    }
608
609    Tensor::new(TensorPrimitive::Float(B::linear(
610        input.primitive.tensor(),
611        weight.primitive.tensor(),
612        bias.map(|b| b.primitive.tensor()),
613    )))
614}
615
616/// Computes scaled dot-product attention: softmax(QKᵗ * scale) · V,
617/// where scale defaults to 1/sqrt(head_dim) (configurable via `options.scale`).
618/// Optionally applies masking, additive bias, causal masking, and softcap.
619///
620/// # Arguments
621/// - `query`: Query tensor of shape `[batch_size, num_heads, seq_len_q, head_dim]`
622/// - `key`: Key tensor of shape `[batch_size, num_heads, seq_len_k, head_dim]`
623/// - `value`: Value tensor of shape `[batch_size, num_heads, seq_len_k, val_dim]`
624/// - `mask`: Optional boolean mask of shape `[batch_size, num_heads, seq_len_q, seq_len_k]`,
625///   where `true` indicates positions to mask (i.e. set to -inf before softmax).
626/// - `attn_bias`: Optional float tensor of shape `[batch_size, num_heads, seq_len_q, seq_len_k]`
627///   added to the attention scores before softmax (e.g. ALiBi, relative position biases).
628/// - `options`: Additional attention options (custom scale, softcap, causal masking).
629///
630/// # Returns
631/// A tensor of shape `[batch_size, num_heads, seq_len_q, val_dim]`
632/// representing the attended context per head.
633///
634/// # Note
635/// This implementation does not support dropout and is intended for inference or
636/// use cases where dropout is not needed.
637pub fn attention<B: Backend>(
638    query: Tensor<B, 4>,
639    key: Tensor<B, 4>,
640    value: Tensor<B, 4>,
641    mask: Option<Tensor<B, 4, Bool>>,
642    attn_bias: Option<Tensor<B, 4>>,
643    options: AttentionModuleOptions,
644) -> Tensor<B, 4> {
645    Tensor::new(TensorPrimitive::Float(B::attention(
646        query.primitive.tensor(),
647        key.primitive.tensor(),
648        value.primitive.tensor(),
649        mask.map(|mask| mask.primitive),
650        attn_bias.map(|bias| bias.primitive.tensor()),
651        options,
652    )))
653}
654
655/// Exports attention fallback to test backend's attention against.
656pub fn attention_fallback<B: Backend>(
657    query: Tensor<B, 4>,
658    key: Tensor<B, 4>,
659    value: Tensor<B, 4>,
660    mask: Option<Tensor<B, 4, Bool>>,
661    attn_bias: Option<Tensor<B, 4>>,
662    options: AttentionModuleOptions,
663) -> Tensor<B, 4> {
664    Tensor::new(TensorPrimitive::Float(
665        crate::api::ops::attention::attention_fallback::<B>(
666            query.primitive.tensor(),
667            key.primitive.tensor(),
668            value.primitive.tensor(),
669            mask.map(|mask| mask.primitive),
670            attn_bias.map(|bias| bias.primitive.tensor()),
671            options,
672        ),
673    ))
674}