Skip to main content

burn_dispatch/ops/
module.rs

1use burn_backend::{
2    IntDType,
3    ops::{
4        DeformConv2dBackward, MaxPool1dBackward, MaxPool1dWithIndices, MaxPool2dBackward,
5        MaxPool2dWithIndices, ModuleOps,
6    },
7    tensor::{FloatTensor, IntTensor},
8};
9
10use crate::Dispatch;
11
12impl ModuleOps<Self> for Dispatch {
13    fn conv2d(
14        x: FloatTensor<Self>,
15        weight: FloatTensor<Self>,
16        bias: Option<FloatTensor<Self>>,
17        options: burn_backend::ops::ConvOptions<2>,
18    ) -> FloatTensor<Self> {
19        multi_op!(
20            inputs[(x, float), (weight, float)],
21            opt_inputs[(bias, float)],
22            => Float,
23            B::conv2d(x, weight, bias, options)
24        )
25    }
26
27    fn deform_conv2d(
28        x: FloatTensor<Self>,
29        offset: FloatTensor<Self>,
30        weight: FloatTensor<Self>,
31        mask: Option<FloatTensor<Self>>,
32        bias: Option<FloatTensor<Self>>,
33        options: burn_backend::ops::DeformConvOptions<2>,
34    ) -> FloatTensor<Self> {
35        multi_op!(
36            inputs[(x, float), (offset, float), (weight, float)],
37            opt_inputs[(mask, float), (bias, float)],
38            => Float,
39            B::deform_conv2d(x, offset, weight, mask, bias, options)
40        )
41    }
42
43    fn deform_conv2d_backward(
44        x: FloatTensor<Self>,
45        offset: FloatTensor<Self>,
46        weight: FloatTensor<Self>,
47        mask: Option<FloatTensor<Self>>,
48        bias: Option<FloatTensor<Self>>,
49        output_grad: FloatTensor<Self>,
50        options: burn_backend::ops::DeformConvOptions<2>,
51    ) -> DeformConv2dBackward<Self> {
52        let (x_grad, offset_grad, weight_grad, mask_grad, bias_grad) = multi_op!(
53            inputs[(x, float), (offset, float), (weight, float), (output_grad, float)],
54            opt_inputs[(mask, float), (bias, float)],
55            outputs[(x_grad, Float), (offset_grad, Float), (weight_grad, Float)],
56            opt_outputs[mask_grad, bias_grad],
57            {
58                let res = B::deform_conv2d_backward(x, offset, weight, mask, bias, output_grad, options);
59                (res.x_grad, res.offset_grad, res.weight_grad, res.mask_grad, res.bias_grad)
60            }
61        );
62        DeformConv2dBackward::new(x_grad, offset_grad, weight_grad, mask_grad, bias_grad)
63    }
64
65    fn conv3d(
66        x: FloatTensor<Self>,
67        weight: FloatTensor<Self>,
68        bias: Option<FloatTensor<Self>>,
69        options: burn_backend::ops::ConvOptions<3>,
70    ) -> FloatTensor<Self> {
71        multi_op!(
72            inputs[(x, float), (weight, float)],
73            opt_inputs[(bias, float)],
74            => Float,
75            B::conv3d(x, weight, bias, options)
76        )
77    }
78
79    fn conv_transpose2d(
80        x: FloatTensor<Self>,
81        weight: FloatTensor<Self>,
82        bias: Option<FloatTensor<Self>>,
83        options: burn_backend::ops::ConvTransposeOptions<2>,
84    ) -> FloatTensor<Self> {
85        multi_op!(
86            inputs[(x, float), (weight, float)],
87            opt_inputs[(bias, float)],
88            => Float,
89            B::conv_transpose2d(x, weight, bias, options)
90        )
91    }
92
93    fn conv_transpose3d(
94        x: FloatTensor<Self>,
95        weight: FloatTensor<Self>,
96        bias: Option<FloatTensor<Self>>,
97        options: burn_backend::ops::ConvTransposeOptions<3>,
98    ) -> FloatTensor<Self> {
99        multi_op!(
100            inputs[(x, float), (weight, float)],
101            opt_inputs[(bias, float)],
102            => Float,
103            B::conv_transpose3d(x, weight, bias, options)
104        )
105    }
106
107    fn avg_pool2d(
108        x: FloatTensor<Self>,
109        kernel_size: [usize; 2],
110        stride: [usize; 2],
111        padding: [usize; 2],
112        count_include_pad: bool,
113        ceil_mode: bool,
114    ) -> FloatTensor<Self> {
115        multi_op!(inputs[(x, float)],
116            => Float,
117            B::avg_pool2d(x, kernel_size, stride, padding, count_include_pad, ceil_mode)
118        )
119    }
120
121    fn avg_pool2d_backward(
122        x: FloatTensor<Self>,
123        grad: FloatTensor<Self>,
124        kernel_size: [usize; 2],
125        stride: [usize; 2],
126        padding: [usize; 2],
127        count_include_pad: bool,
128        ceil_mode: bool,
129    ) -> FloatTensor<Self> {
130        multi_op!(
131            inputs[(x, float), (grad, float)],
132            => Float,
133            B::avg_pool2d_backward(x, grad, kernel_size, stride, padding, count_include_pad, ceil_mode)
134        )
135    }
136
137    fn adaptive_avg_pool2d(x: FloatTensor<Self>, output_size: [usize; 2]) -> FloatTensor<Self> {
138        multi_op!(
139            inputs[(x, float)],
140            => Float,
141            B::adaptive_avg_pool2d(x, output_size)
142        )
143    }
144
145    fn adaptive_avg_pool2d_backward(
146        x: FloatTensor<Self>,
147        grad: FloatTensor<Self>,
148    ) -> FloatTensor<Self> {
149        multi_op!(
150            inputs[(x, float), (grad, float)],
151            => Float,
152            B::adaptive_avg_pool2d_backward(x, grad)
153        )
154    }
155
156    fn adaptive_avg_pool3d(x: FloatTensor<Self>, output_size: [usize; 3]) -> FloatTensor<Self> {
157        multi_op!(
158            inputs[(x, float)],
159            => Float,
160            B::adaptive_avg_pool3d(x, output_size)
161        )
162    }
163
164    fn adaptive_avg_pool3d_backward(
165        x: FloatTensor<Self>,
166        grad: FloatTensor<Self>,
167    ) -> FloatTensor<Self> {
168        multi_op!(
169            inputs[(x, float), (grad, float)],
170            => Float,
171            B::adaptive_avg_pool3d_backward(x, grad)
172        )
173    }
174
175    fn max_pool2d(
176        x: FloatTensor<Self>,
177        kernel_size: [usize; 2],
178        stride: [usize; 2],
179        padding: [usize; 2],
180        dilation: [usize; 2],
181        ceil_mode: bool,
182    ) -> FloatTensor<Self> {
183        multi_op!(
184            inputs[(x, float)],
185            => Float,
186            B::max_pool2d(x, kernel_size, stride, padding, dilation, ceil_mode)
187        )
188    }
189
190    fn max_pool2d_with_indices(
191        x: FloatTensor<Self>,
192        kernel_size: [usize; 2],
193        stride: [usize; 2],
194        padding: [usize; 2],
195        dilation: [usize; 2],
196        ceil_mode: bool,
197        indices_dtype: IntDType,
198    ) -> MaxPool2dWithIndices<Self> {
199        let (out, indices) = multi_op!(
200            inputs[(x, float)],
201            outputs[(out, Float), (indices, Int)],
202            {
203                let res = B::max_pool2d_with_indices(x, kernel_size, stride, padding, dilation, ceil_mode, indices_dtype);
204                (res.output, res.indices)
205            }
206        );
207        MaxPool2dWithIndices::new(out, indices)
208    }
209
210    fn max_pool2d_with_indices_backward(
211        x: FloatTensor<Self>,
212        kernel_size: [usize; 2],
213        stride: [usize; 2],
214        padding: [usize; 2],
215        dilation: [usize; 2],
216        ceil_mode: bool,
217        output_grad: FloatTensor<Self>,
218        indices: IntTensor<Self>,
219    ) -> MaxPool2dBackward<Self> {
220        let x_grad = multi_op!(
221            inputs[(x, float), (output_grad, float), (indices, int)],
222            => Float,
223            {
224                let res = B::max_pool2d_with_indices_backward(x, kernel_size, stride, padding, dilation, ceil_mode, output_grad, indices);
225                res.x_grad
226            }
227        );
228        MaxPool2dBackward::new(x_grad)
229    }
230
231    fn interpolate(
232        x: FloatTensor<Self>,
233        output_size: [usize; 2],
234        options: burn_backend::ops::InterpolateOptions,
235    ) -> FloatTensor<Self> {
236        multi_op!(
237            inputs[(x, float)],
238            => Float,
239            B::interpolate(x, output_size, options)
240        )
241    }
242
243    fn interpolate_backward(
244        x: FloatTensor<Self>,
245        grad: FloatTensor<Self>,
246        output_size: [usize; 2],
247        options: burn_backend::ops::InterpolateOptions,
248    ) -> FloatTensor<Self> {
249        multi_op!(
250            inputs[(x, float), (grad, float)],
251            => Float,
252            B::interpolate_backward(x, grad, output_size, options)
253        )
254    }
255
256    fn embedding(weights: FloatTensor<Self>, indices: IntTensor<Self>) -> FloatTensor<Self> {
257        multi_op!(
258            inputs[(weights, float), (indices, int)],
259            => Float,
260            B::embedding(weights, indices)
261        )
262    }
263
264    fn embedding_backward(
265        weights: FloatTensor<Self>,
266        output_grad: FloatTensor<Self>,
267        indices: IntTensor<Self>,
268    ) -> FloatTensor<Self> {
269        multi_op!(
270            inputs[(weights, float), (output_grad, float), (indices, int)],
271            => Float,
272            B::embedding_backward(weights, output_grad, indices)
273        )
274    }
275
276    fn conv1d(
277        x: FloatTensor<Self>,
278        weight: FloatTensor<Self>,
279        bias: Option<FloatTensor<Self>>,
280        options: burn_backend::ops::ConvOptions<1>,
281    ) -> FloatTensor<Self> {
282        multi_op!(
283            inputs[(x, float), (weight, float)],
284            opt_inputs[(bias, float)],
285            => Float,
286            B::conv1d(x, weight, bias, options)
287        )
288    }
289
290    fn conv1d_x_backward(
291        x: FloatTensor<Self>,
292        weight: FloatTensor<Self>,
293        output_grad: FloatTensor<Self>,
294        options: burn_backend::ops::ConvOptions<1>,
295    ) -> FloatTensor<Self> {
296        multi_op!(
297            inputs[(x, float), (weight, float), (output_grad, float)],
298            => Float,
299            B::conv1d_x_backward(x, weight, output_grad, options)
300        )
301    }
302
303    fn conv1d_weight_backward(
304        x: FloatTensor<Self>,
305        weight: FloatTensor<Self>,
306        output_grad: FloatTensor<Self>,
307        options: burn_backend::ops::ConvOptions<1>,
308    ) -> FloatTensor<Self> {
309        multi_op!(
310            inputs[(x, float), (weight, float), (output_grad, float)],
311            => Float,
312            B::conv1d_weight_backward(x, weight, output_grad, options)
313        )
314    }
315
316    fn conv1d_bias_backward(
317        x: FloatTensor<Self>,
318        bias: FloatTensor<Self>,
319        output_grad: FloatTensor<Self>,
320    ) -> FloatTensor<Self> {
321        multi_op!(
322            inputs[(x, float), (bias, float), (output_grad, float)],
323            => Float,
324            B::conv1d_bias_backward(x, bias, output_grad)
325        )
326    }
327
328    fn conv2d_x_backward(
329        x: FloatTensor<Self>,
330        weight: FloatTensor<Self>,
331        output_grad: FloatTensor<Self>,
332        options: burn_backend::ops::ConvOptions<2>,
333    ) -> FloatTensor<Self> {
334        multi_op!(
335            inputs[(x, float), (weight, float), (output_grad, float)],
336            => Float,
337            B::conv2d_x_backward(x, weight, output_grad, options)
338        )
339    }
340
341    fn conv2d_weight_backward(
342        x: FloatTensor<Self>,
343        weight: FloatTensor<Self>,
344        output_grad: FloatTensor<Self>,
345        options: burn_backend::ops::ConvOptions<2>,
346    ) -> FloatTensor<Self> {
347        multi_op!(
348            inputs[(x, float), (weight, float), (output_grad, float)],
349            => Float,
350            B::conv2d_weight_backward(x, weight, output_grad, options)
351        )
352    }
353
354    fn conv2d_bias_backward(
355        x: FloatTensor<Self>,
356        bias: FloatTensor<Self>,
357        output_grad: FloatTensor<Self>,
358    ) -> FloatTensor<Self> {
359        multi_op!(
360            inputs[(x, float), (bias, float), (output_grad, float)],
361            => Float,
362            B::conv2d_bias_backward(x, bias, output_grad)
363        )
364    }
365
366    fn conv3d_x_backward(
367        x: FloatTensor<Self>,
368        weight: FloatTensor<Self>,
369        output_grad: FloatTensor<Self>,
370        options: burn_backend::ops::ConvOptions<3>,
371    ) -> FloatTensor<Self> {
372        multi_op!(
373            inputs[(x, float), (weight, float), (output_grad, float)],
374            => Float,
375            B::conv3d_x_backward(x, weight, output_grad, options)
376        )
377    }
378
379    fn conv3d_weight_backward(
380        x: FloatTensor<Self>,
381        weight: FloatTensor<Self>,
382        output_grad: FloatTensor<Self>,
383        options: burn_backend::ops::ConvOptions<3>,
384    ) -> FloatTensor<Self> {
385        multi_op!(
386            inputs[(x, float), (weight, float), (output_grad, float)],
387            => Float,
388            B::conv3d_weight_backward(x, weight, output_grad, options)
389        )
390    }
391
392    fn conv3d_bias_backward(
393        x: FloatTensor<Self>,
394        bias: FloatTensor<Self>,
395        output_grad: FloatTensor<Self>,
396    ) -> FloatTensor<Self> {
397        multi_op!(
398            inputs[(x, float), (bias, float), (output_grad, float)],
399            => Float,
400            B::conv3d_bias_backward(x, bias, output_grad)
401        )
402    }
403
404    fn conv_transpose1d(
405        x: FloatTensor<Self>,
406        weight: FloatTensor<Self>,
407        bias: Option<FloatTensor<Self>>,
408        options: burn_backend::ops::ConvTransposeOptions<1>,
409    ) -> FloatTensor<Self> {
410        multi_op!(
411            inputs[(x, float), (weight, float)],
412            opt_inputs[(bias, float)],
413            => Float,
414            B::conv_transpose1d(x, weight, bias, options)
415        )
416    }
417
418    fn conv_transpose1d_x_backward(
419        weight: FloatTensor<Self>,
420        output_grad: FloatTensor<Self>,
421        options: burn_backend::ops::ConvTransposeOptions<1>,
422    ) -> FloatTensor<Self> {
423        multi_op!(
424            inputs[(weight, float), (output_grad, float)],
425            => Float,
426            B::conv_transpose1d_x_backward(weight, output_grad, options)
427        )
428    }
429
430    fn conv_transpose1d_weight_backward(
431        x: FloatTensor<Self>,
432        weight: FloatTensor<Self>,
433        output_grad: FloatTensor<Self>,
434        options: burn_backend::ops::ConvTransposeOptions<1>,
435    ) -> FloatTensor<Self> {
436        multi_op!(
437            inputs[(x, float), (weight, float), (output_grad, float)],
438            => Float,
439            B::conv_transpose1d_weight_backward(x, weight, output_grad, options)
440        )
441    }
442
443    fn conv_transpose1d_bias_backward(
444        x: FloatTensor<Self>,
445        bias: FloatTensor<Self>,
446        output_grad: FloatTensor<Self>,
447    ) -> FloatTensor<Self> {
448        multi_op!(
449            inputs[(x, float), (bias, float), (output_grad, float)],
450            => Float,
451            B::conv_transpose1d_bias_backward(x, bias, output_grad)
452        )
453    }
454
455    fn conv_transpose2d_x_backward(
456        weight: FloatTensor<Self>,
457        output_grad: FloatTensor<Self>,
458        options: burn_backend::ops::ConvTransposeOptions<2>,
459    ) -> FloatTensor<Self> {
460        multi_op!(
461            inputs[(weight, float), (output_grad, float)],
462            => Float,
463            B::conv_transpose2d_x_backward(weight, output_grad, options)
464        )
465    }
466
467    fn conv_transpose2d_weight_backward(
468        x: FloatTensor<Self>,
469        weight: FloatTensor<Self>,
470        output_grad: FloatTensor<Self>,
471        options: burn_backend::ops::ConvTransposeOptions<2>,
472    ) -> FloatTensor<Self> {
473        multi_op!(
474            inputs[(x, float), (weight, float), (output_grad, float)],
475            => Float,
476            B::conv_transpose2d_weight_backward(x, weight, output_grad, options)
477        )
478    }
479
480    fn conv_transpose2d_bias_backward(
481        x: FloatTensor<Self>,
482        bias: FloatTensor<Self>,
483        output_grad: FloatTensor<Self>,
484    ) -> FloatTensor<Self> {
485        multi_op!(
486            inputs[(x, float), (bias, float), (output_grad, float)],
487            => Float,
488            B::conv_transpose2d_bias_backward(x, bias, output_grad)
489        )
490    }
491
492    fn conv_transpose3d_x_backward(
493        weight: FloatTensor<Self>,
494        output_grad: FloatTensor<Self>,
495        options: burn_backend::ops::ConvTransposeOptions<3>,
496    ) -> FloatTensor<Self> {
497        multi_op!(
498            inputs[(weight, float), (output_grad, float)],
499            => Float,
500            B::conv_transpose3d_x_backward(weight, output_grad, options)
501        )
502    }
503
504    fn conv_transpose3d_weight_backward(
505        x: FloatTensor<Self>,
506        weight: FloatTensor<Self>,
507        output_grad: FloatTensor<Self>,
508        options: burn_backend::ops::ConvTransposeOptions<3>,
509    ) -> FloatTensor<Self> {
510        multi_op!(
511            inputs[(x, float), (weight, float), (output_grad, float)],
512            => Float,
513            B::conv_transpose3d_weight_backward(x, weight, output_grad, options)
514        )
515    }
516
517    fn conv_transpose3d_bias_backward(
518        x: FloatTensor<Self>,
519        bias: FloatTensor<Self>,
520        output_grad: FloatTensor<Self>,
521    ) -> FloatTensor<Self> {
522        multi_op!(
523            inputs[(x, float), (bias, float), (output_grad, float)],
524            => Float,
525            B::conv_transpose3d_bias_backward(x, bias, output_grad)
526        )
527    }
528
529    fn unfold4d(
530        x: FloatTensor<Self>,
531        kernel_size: [usize; 2],
532        options: burn_backend::ops::UnfoldOptions,
533    ) -> FloatTensor<Self> {
534        multi_op!(inputs[(x, float)], => Float, B::unfold4d(x, kernel_size, options))
535    }
536
537    fn avg_pool1d(
538        x: FloatTensor<Self>,
539        kernel_size: usize,
540        stride: usize,
541        padding: usize,
542        count_include_pad: bool,
543        ceil_mode: bool,
544    ) -> FloatTensor<Self> {
545        multi_op!(inputs[(x, float)], => Float,
546            B::avg_pool1d(x, kernel_size, stride, padding, count_include_pad, ceil_mode)
547        )
548    }
549
550    fn avg_pool1d_backward(
551        x: FloatTensor<Self>,
552        grad: FloatTensor<Self>,
553        kernel_size: usize,
554        stride: usize,
555        padding: usize,
556        count_include_pad: bool,
557        ceil_mode: bool,
558    ) -> FloatTensor<Self> {
559        multi_op!(
560            inputs[(x, float), (grad, float)],
561            => Float,
562            B::avg_pool1d_backward(x, grad, kernel_size, stride, padding, count_include_pad, ceil_mode)
563        )
564    }
565
566    fn adaptive_avg_pool1d(x: FloatTensor<Self>, output_size: usize) -> FloatTensor<Self> {
567        multi_op!(inputs[(x, float)], => Float, B::adaptive_avg_pool1d(x, output_size))
568    }
569
570    fn adaptive_avg_pool1d_backward(
571        x: FloatTensor<Self>,
572        grad: FloatTensor<Self>,
573    ) -> FloatTensor<Self> {
574        multi_op!(
575            inputs[(x, float), (grad, float)],
576            => Float,
577            B::adaptive_avg_pool1d_backward(x, grad)
578        )
579    }
580
581    fn max_pool1d(
582        x: FloatTensor<Self>,
583        kernel_size: usize,
584        stride: usize,
585        padding: usize,
586        dilation: usize,
587        ceil_mode: bool,
588    ) -> FloatTensor<Self> {
589        multi_op!(inputs[(x, float)], => Float,
590            B::max_pool1d(x, kernel_size, stride, padding, dilation, ceil_mode))
591    }
592
593    fn max_pool1d_with_indices(
594        x: FloatTensor<Self>,
595        kernel_size: usize,
596        stride: usize,
597        padding: usize,
598        dilation: usize,
599        ceil_mode: bool,
600        indices_dtype: IntDType,
601    ) -> MaxPool1dWithIndices<Self> {
602        let (out, indices) = multi_op!(
603            inputs[(x, float)],
604            outputs[(out, Float), (indices, Int)],
605            {
606                let res = B::max_pool1d_with_indices(x, kernel_size, stride, padding, dilation, ceil_mode, indices_dtype);
607                (res.output, res.indices)
608            }
609        );
610        MaxPool1dWithIndices::new(out, indices)
611    }
612
613    fn max_pool1d_with_indices_backward(
614        x: FloatTensor<Self>,
615        kernel_size: usize,
616        stride: usize,
617        padding: usize,
618        dilation: usize,
619        ceil_mode: bool,
620        output_grad: FloatTensor<Self>,
621        indices: IntTensor<Self>,
622    ) -> MaxPool1dBackward<Self> {
623        let x_grad = multi_op!(
624            inputs[(x, float), (output_grad, float), (indices, int)],
625            => Float,
626            {
627                let res = B::max_pool1d_with_indices_backward(x, kernel_size, stride, padding, dilation, ceil_mode, output_grad, indices);
628                res.x_grad
629            }
630        );
631        MaxPool1dBackward::new(x_grad)
632    }
633
634    fn attention(
635        query: FloatTensor<Self>,
636        key: FloatTensor<Self>,
637        value: FloatTensor<Self>,
638        mask: Option<burn_backend::tensor::BoolTensor<Self>>,
639        attn_bias: Option<FloatTensor<Self>>,
640        options: burn_backend::ops::AttentionModuleOptions,
641    ) -> FloatTensor<Self> {
642        multi_op!(
643            inputs[(query, float), (key, float), (value, float)],
644            opt_inputs[(mask, bool), (attn_bias, float)],
645            => Float,
646            B::attention(query, key, value, mask, attn_bias, options)
647        )
648    }
649
650    fn layer_norm(
651        tensor: FloatTensor<Self>,
652        gamma: FloatTensor<Self>,
653        beta: Option<FloatTensor<Self>>,
654        epsilon: f64,
655    ) -> FloatTensor<Self> {
656        multi_op!(
657            inputs[(tensor, float), (gamma, float)],
658            opt_inputs[(beta, float)],
659            => Float,
660            B::layer_norm(tensor, gamma, beta, epsilon)
661        )
662    }
663
664    fn rfft(
665        signal: FloatTensor<Self>,
666        dim: usize,
667        n: Option<usize>,
668    ) -> (FloatTensor<Self>, FloatTensor<Self>) {
669        let (real, imag) = multi_op!(
670            inputs[(signal, float)],
671            outputs[(real, Float), (imag, Float)],
672            {
673                let res = B::rfft(signal, dim, n);
674                (res.0, res.1)
675            }
676        );
677
678        (real, imag)
679    }
680
681    fn irfft(
682        spectrum_re: FloatTensor<Self>,
683        spectrum_im: FloatTensor<Self>,
684        dim: usize,
685        n: Option<usize>,
686    ) -> FloatTensor<Self> {
687        multi_op!(
688            inputs[(spectrum_re, float), (spectrum_im, float)],
689            => Float,
690            {
691                B::irfft(spectrum_re, spectrum_im, dim, n)
692            }
693        )
694    }
695
696    fn has_ctc_loss_backward() -> bool {
697        // Dispatch routes per-tensor at runtime, but autodiff queries this flag
698        // statically. Returning `false` makes autodiff differentiate through
699        // the default decomposed forward, which is safe for every inner
700        // backend regardless of whether it has its own ctc_loss_backward.
701        false
702    }
703
704    fn ctc_loss(
705        log_probs: FloatTensor<Self>,
706        targets: IntTensor<Self>,
707        input_lengths: IntTensor<Self>,
708        target_lengths: IntTensor<Self>,
709        blank: usize,
710    ) -> FloatTensor<Self> {
711        multi_op!(
712            inputs[(log_probs, float), (targets, int), (input_lengths, int), (target_lengths, int)],
713            => Float,
714            B::ctc_loss(log_probs, targets, input_lengths, target_lengths, blank)
715        )
716    }
717
718    fn ctc_loss_backward(
719        log_probs: FloatTensor<Self>,
720        targets: IntTensor<Self>,
721        input_lengths: IntTensor<Self>,
722        target_lengths: IntTensor<Self>,
723        grad_loss: FloatTensor<Self>,
724        blank: usize,
725    ) -> FloatTensor<Self> {
726        multi_op!(
727            inputs[(log_probs, float), (targets, int), (input_lengths, int), (target_lengths, int), (grad_loss, float)],
728            => Float,
729            B::ctc_loss_backward(log_probs, targets, input_lengths, target_lengths, grad_loss, blank)
730        )
731    }
732
733    // TODO: linear ops
734    // fn linear(
735    //         x: FloatTensor<Self>,
736    //         weight: FloatTensor<Self>,
737    //         bias: Option<FloatTensor<Self>>,
738    //     ) -> FloatTensor<Self> {
739
740    // }
741}