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 max_pool2d(
157        x: FloatTensor<Self>,
158        kernel_size: [usize; 2],
159        stride: [usize; 2],
160        padding: [usize; 2],
161        dilation: [usize; 2],
162        ceil_mode: bool,
163    ) -> FloatTensor<Self> {
164        multi_op!(
165            inputs[(x, float)],
166            => Float,
167            B::max_pool2d(x, kernel_size, stride, padding, dilation, ceil_mode)
168        )
169    }
170
171    fn max_pool2d_with_indices(
172        x: FloatTensor<Self>,
173        kernel_size: [usize; 2],
174        stride: [usize; 2],
175        padding: [usize; 2],
176        dilation: [usize; 2],
177        ceil_mode: bool,
178        indices_dtype: IntDType,
179    ) -> MaxPool2dWithIndices<Self> {
180        let (out, indices) = multi_op!(
181            inputs[(x, float)],
182            outputs[(out, Float), (indices, Int)],
183            {
184                let res = B::max_pool2d_with_indices(x, kernel_size, stride, padding, dilation, ceil_mode, indices_dtype);
185                (res.output, res.indices)
186            }
187        );
188        MaxPool2dWithIndices::new(out, indices)
189    }
190
191    fn max_pool2d_with_indices_backward(
192        x: FloatTensor<Self>,
193        kernel_size: [usize; 2],
194        stride: [usize; 2],
195        padding: [usize; 2],
196        dilation: [usize; 2],
197        ceil_mode: bool,
198        output_grad: FloatTensor<Self>,
199        indices: IntTensor<Self>,
200    ) -> MaxPool2dBackward<Self> {
201        let x_grad = multi_op!(
202            inputs[(x, float), (output_grad, float), (indices, int)],
203            => Float,
204            {
205                let res = B::max_pool2d_with_indices_backward(x, kernel_size, stride, padding, dilation, ceil_mode, output_grad, indices);
206                res.x_grad
207            }
208        );
209        MaxPool2dBackward::new(x_grad)
210    }
211
212    fn interpolate(
213        x: FloatTensor<Self>,
214        output_size: [usize; 2],
215        options: burn_backend::ops::InterpolateOptions,
216    ) -> FloatTensor<Self> {
217        multi_op!(
218            inputs[(x, float)],
219            => Float,
220            B::interpolate(x, output_size, options)
221        )
222    }
223
224    fn interpolate_backward(
225        x: FloatTensor<Self>,
226        grad: FloatTensor<Self>,
227        output_size: [usize; 2],
228        options: burn_backend::ops::InterpolateOptions,
229    ) -> FloatTensor<Self> {
230        multi_op!(
231            inputs[(x, float), (grad, float)],
232            => Float,
233            B::interpolate_backward(x, grad, output_size, options)
234        )
235    }
236
237    fn embedding(weights: FloatTensor<Self>, indices: IntTensor<Self>) -> FloatTensor<Self> {
238        multi_op!(
239            inputs[(weights, float), (indices, int)],
240            => Float,
241            B::embedding(weights, indices)
242        )
243    }
244
245    fn embedding_backward(
246        weights: FloatTensor<Self>,
247        output_grad: FloatTensor<Self>,
248        indices: IntTensor<Self>,
249    ) -> FloatTensor<Self> {
250        multi_op!(
251            inputs[(weights, float), (output_grad, float), (indices, int)],
252            => Float,
253            B::embedding_backward(weights, output_grad, indices)
254        )
255    }
256
257    fn conv1d(
258        x: FloatTensor<Self>,
259        weight: FloatTensor<Self>,
260        bias: Option<FloatTensor<Self>>,
261        options: burn_backend::ops::ConvOptions<1>,
262    ) -> FloatTensor<Self> {
263        multi_op!(
264            inputs[(x, float), (weight, float)],
265            opt_inputs[(bias, float)],
266            => Float,
267            B::conv1d(x, weight, bias, options)
268        )
269    }
270
271    fn conv1d_x_backward(
272        x: FloatTensor<Self>,
273        weight: FloatTensor<Self>,
274        output_grad: FloatTensor<Self>,
275        options: burn_backend::ops::ConvOptions<1>,
276    ) -> FloatTensor<Self> {
277        multi_op!(
278            inputs[(x, float), (weight, float), (output_grad, float)],
279            => Float,
280            B::conv1d_x_backward(x, weight, output_grad, options)
281        )
282    }
283
284    fn conv1d_weight_backward(
285        x: FloatTensor<Self>,
286        weight: FloatTensor<Self>,
287        output_grad: FloatTensor<Self>,
288        options: burn_backend::ops::ConvOptions<1>,
289    ) -> FloatTensor<Self> {
290        multi_op!(
291            inputs[(x, float), (weight, float), (output_grad, float)],
292            => Float,
293            B::conv1d_weight_backward(x, weight, output_grad, options)
294        )
295    }
296
297    fn conv1d_bias_backward(
298        x: FloatTensor<Self>,
299        bias: FloatTensor<Self>,
300        output_grad: FloatTensor<Self>,
301    ) -> FloatTensor<Self> {
302        multi_op!(
303            inputs[(x, float), (bias, float), (output_grad, float)],
304            => Float,
305            B::conv1d_bias_backward(x, bias, output_grad)
306        )
307    }
308
309    fn conv2d_x_backward(
310        x: FloatTensor<Self>,
311        weight: FloatTensor<Self>,
312        output_grad: FloatTensor<Self>,
313        options: burn_backend::ops::ConvOptions<2>,
314    ) -> FloatTensor<Self> {
315        multi_op!(
316            inputs[(x, float), (weight, float), (output_grad, float)],
317            => Float,
318            B::conv2d_x_backward(x, weight, output_grad, options)
319        )
320    }
321
322    fn conv2d_weight_backward(
323        x: FloatTensor<Self>,
324        weight: FloatTensor<Self>,
325        output_grad: FloatTensor<Self>,
326        options: burn_backend::ops::ConvOptions<2>,
327    ) -> FloatTensor<Self> {
328        multi_op!(
329            inputs[(x, float), (weight, float), (output_grad, float)],
330            => Float,
331            B::conv2d_weight_backward(x, weight, output_grad, options)
332        )
333    }
334
335    fn conv2d_bias_backward(
336        x: FloatTensor<Self>,
337        bias: FloatTensor<Self>,
338        output_grad: FloatTensor<Self>,
339    ) -> FloatTensor<Self> {
340        multi_op!(
341            inputs[(x, float), (bias, float), (output_grad, float)],
342            => Float,
343            B::conv2d_bias_backward(x, bias, output_grad)
344        )
345    }
346
347    fn conv3d_x_backward(
348        x: FloatTensor<Self>,
349        weight: FloatTensor<Self>,
350        output_grad: FloatTensor<Self>,
351        options: burn_backend::ops::ConvOptions<3>,
352    ) -> FloatTensor<Self> {
353        multi_op!(
354            inputs[(x, float), (weight, float), (output_grad, float)],
355            => Float,
356            B::conv3d_x_backward(x, weight, output_grad, options)
357        )
358    }
359
360    fn conv3d_weight_backward(
361        x: FloatTensor<Self>,
362        weight: FloatTensor<Self>,
363        output_grad: FloatTensor<Self>,
364        options: burn_backend::ops::ConvOptions<3>,
365    ) -> FloatTensor<Self> {
366        multi_op!(
367            inputs[(x, float), (weight, float), (output_grad, float)],
368            => Float,
369            B::conv3d_weight_backward(x, weight, output_grad, options)
370        )
371    }
372
373    fn conv3d_bias_backward(
374        x: FloatTensor<Self>,
375        bias: FloatTensor<Self>,
376        output_grad: FloatTensor<Self>,
377    ) -> FloatTensor<Self> {
378        multi_op!(
379            inputs[(x, float), (bias, float), (output_grad, float)],
380            => Float,
381            B::conv3d_bias_backward(x, bias, output_grad)
382        )
383    }
384
385    fn conv_transpose1d(
386        x: FloatTensor<Self>,
387        weight: FloatTensor<Self>,
388        bias: Option<FloatTensor<Self>>,
389        options: burn_backend::ops::ConvTransposeOptions<1>,
390    ) -> FloatTensor<Self> {
391        multi_op!(
392            inputs[(x, float), (weight, float)],
393            opt_inputs[(bias, float)],
394            => Float,
395            B::conv_transpose1d(x, weight, bias, options)
396        )
397    }
398
399    fn conv_transpose1d_x_backward(
400        weight: FloatTensor<Self>,
401        output_grad: FloatTensor<Self>,
402        options: burn_backend::ops::ConvTransposeOptions<1>,
403    ) -> FloatTensor<Self> {
404        multi_op!(
405            inputs[(weight, float), (output_grad, float)],
406            => Float,
407            B::conv_transpose1d_x_backward(weight, output_grad, options)
408        )
409    }
410
411    fn conv_transpose1d_weight_backward(
412        x: FloatTensor<Self>,
413        weight: FloatTensor<Self>,
414        output_grad: FloatTensor<Self>,
415        options: burn_backend::ops::ConvTransposeOptions<1>,
416    ) -> FloatTensor<Self> {
417        multi_op!(
418            inputs[(x, float), (weight, float), (output_grad, float)],
419            => Float,
420            B::conv_transpose1d_weight_backward(x, weight, output_grad, options)
421        )
422    }
423
424    fn conv_transpose1d_bias_backward(
425        x: FloatTensor<Self>,
426        bias: FloatTensor<Self>,
427        output_grad: FloatTensor<Self>,
428    ) -> FloatTensor<Self> {
429        multi_op!(
430            inputs[(x, float), (bias, float), (output_grad, float)],
431            => Float,
432            B::conv_transpose1d_bias_backward(x, bias, output_grad)
433        )
434    }
435
436    fn conv_transpose2d_x_backward(
437        weight: FloatTensor<Self>,
438        output_grad: FloatTensor<Self>,
439        options: burn_backend::ops::ConvTransposeOptions<2>,
440    ) -> FloatTensor<Self> {
441        multi_op!(
442            inputs[(weight, float), (output_grad, float)],
443            => Float,
444            B::conv_transpose2d_x_backward(weight, output_grad, options)
445        )
446    }
447
448    fn conv_transpose2d_weight_backward(
449        x: FloatTensor<Self>,
450        weight: FloatTensor<Self>,
451        output_grad: FloatTensor<Self>,
452        options: burn_backend::ops::ConvTransposeOptions<2>,
453    ) -> FloatTensor<Self> {
454        multi_op!(
455            inputs[(x, float), (weight, float), (output_grad, float)],
456            => Float,
457            B::conv_transpose2d_weight_backward(x, weight, output_grad, options)
458        )
459    }
460
461    fn conv_transpose2d_bias_backward(
462        x: FloatTensor<Self>,
463        bias: FloatTensor<Self>,
464        output_grad: FloatTensor<Self>,
465    ) -> FloatTensor<Self> {
466        multi_op!(
467            inputs[(x, float), (bias, float), (output_grad, float)],
468            => Float,
469            B::conv_transpose2d_bias_backward(x, bias, output_grad)
470        )
471    }
472
473    fn conv_transpose3d_x_backward(
474        weight: FloatTensor<Self>,
475        output_grad: FloatTensor<Self>,
476        options: burn_backend::ops::ConvTransposeOptions<3>,
477    ) -> FloatTensor<Self> {
478        multi_op!(
479            inputs[(weight, float), (output_grad, float)],
480            => Float,
481            B::conv_transpose3d_x_backward(weight, output_grad, options)
482        )
483    }
484
485    fn conv_transpose3d_weight_backward(
486        x: FloatTensor<Self>,
487        weight: FloatTensor<Self>,
488        output_grad: FloatTensor<Self>,
489        options: burn_backend::ops::ConvTransposeOptions<3>,
490    ) -> FloatTensor<Self> {
491        multi_op!(
492            inputs[(x, float), (weight, float), (output_grad, float)],
493            => Float,
494            B::conv_transpose3d_weight_backward(x, weight, output_grad, options)
495        )
496    }
497
498    fn conv_transpose3d_bias_backward(
499        x: FloatTensor<Self>,
500        bias: FloatTensor<Self>,
501        output_grad: FloatTensor<Self>,
502    ) -> FloatTensor<Self> {
503        multi_op!(
504            inputs[(x, float), (bias, float), (output_grad, float)],
505            => Float,
506            B::conv_transpose3d_bias_backward(x, bias, output_grad)
507        )
508    }
509
510    fn unfold4d(
511        x: FloatTensor<Self>,
512        kernel_size: [usize; 2],
513        options: burn_backend::ops::UnfoldOptions,
514    ) -> FloatTensor<Self> {
515        multi_op!(inputs[(x, float)], => Float, B::unfold4d(x, kernel_size, options))
516    }
517
518    fn avg_pool1d(
519        x: FloatTensor<Self>,
520        kernel_size: usize,
521        stride: usize,
522        padding: usize,
523        count_include_pad: bool,
524        ceil_mode: bool,
525    ) -> FloatTensor<Self> {
526        multi_op!(inputs[(x, float)], => Float,
527            B::avg_pool1d(x, kernel_size, stride, padding, count_include_pad, ceil_mode)
528        )
529    }
530
531    fn avg_pool1d_backward(
532        x: FloatTensor<Self>,
533        grad: FloatTensor<Self>,
534        kernel_size: usize,
535        stride: usize,
536        padding: usize,
537        count_include_pad: bool,
538        ceil_mode: bool,
539    ) -> FloatTensor<Self> {
540        multi_op!(
541            inputs[(x, float), (grad, float)],
542            => Float,
543            B::avg_pool1d_backward(x, grad, kernel_size, stride, padding, count_include_pad, ceil_mode)
544        )
545    }
546
547    fn adaptive_avg_pool1d(x: FloatTensor<Self>, output_size: usize) -> FloatTensor<Self> {
548        multi_op!(inputs[(x, float)], => Float, B::adaptive_avg_pool1d(x, output_size))
549    }
550
551    fn adaptive_avg_pool1d_backward(
552        x: FloatTensor<Self>,
553        grad: FloatTensor<Self>,
554    ) -> FloatTensor<Self> {
555        multi_op!(
556            inputs[(x, float), (grad, float)],
557            => Float,
558            B::adaptive_avg_pool1d_backward(x, grad)
559        )
560    }
561
562    fn max_pool1d(
563        x: FloatTensor<Self>,
564        kernel_size: usize,
565        stride: usize,
566        padding: usize,
567        dilation: usize,
568        ceil_mode: bool,
569    ) -> FloatTensor<Self> {
570        multi_op!(inputs[(x, float)], => Float,
571            B::max_pool1d(x, kernel_size, stride, padding, dilation, ceil_mode))
572    }
573
574    fn max_pool1d_with_indices(
575        x: FloatTensor<Self>,
576        kernel_size: usize,
577        stride: usize,
578        padding: usize,
579        dilation: usize,
580        ceil_mode: bool,
581        indices_dtype: IntDType,
582    ) -> MaxPool1dWithIndices<Self> {
583        let (out, indices) = multi_op!(
584            inputs[(x, float)],
585            outputs[(out, Float), (indices, Int)],
586            {
587                let res = B::max_pool1d_with_indices(x, kernel_size, stride, padding, dilation, ceil_mode, indices_dtype);
588                (res.output, res.indices)
589            }
590        );
591        MaxPool1dWithIndices::new(out, indices)
592    }
593
594    fn max_pool1d_with_indices_backward(
595        x: FloatTensor<Self>,
596        kernel_size: usize,
597        stride: usize,
598        padding: usize,
599        dilation: usize,
600        ceil_mode: bool,
601        output_grad: FloatTensor<Self>,
602        indices: IntTensor<Self>,
603    ) -> MaxPool1dBackward<Self> {
604        let x_grad = multi_op!(
605            inputs[(x, float), (output_grad, float), (indices, int)],
606            => Float,
607            {
608                let res = B::max_pool1d_with_indices_backward(x, kernel_size, stride, padding, dilation, ceil_mode, output_grad, indices);
609                res.x_grad
610            }
611        );
612        MaxPool1dBackward::new(x_grad)
613    }
614
615    fn attention(
616        query: FloatTensor<Self>,
617        key: FloatTensor<Self>,
618        value: FloatTensor<Self>,
619        mask: Option<burn_backend::tensor::BoolTensor<Self>>,
620        attn_bias: Option<FloatTensor<Self>>,
621        options: burn_backend::ops::AttentionModuleOptions,
622    ) -> FloatTensor<Self> {
623        multi_op!(
624            inputs[(query, float), (key, float), (value, float)],
625            opt_inputs[(mask, bool), (attn_bias, float)],
626            => Float,
627            B::attention(query, key, value, mask, attn_bias, options)
628        )
629    }
630
631    fn layer_norm(
632        tensor: FloatTensor<Self>,
633        gamma: FloatTensor<Self>,
634        beta: Option<FloatTensor<Self>>,
635        epsilon: f64,
636    ) -> FloatTensor<Self> {
637        multi_op!(
638            inputs[(tensor, float), (gamma, float)],
639            opt_inputs[(beta, float)],
640            => Float,
641            B::layer_norm(tensor, gamma, beta, epsilon)
642        )
643    }
644
645    fn rfft(
646        signal: FloatTensor<Self>,
647        dim: usize,
648        n: Option<usize>,
649    ) -> (FloatTensor<Self>, FloatTensor<Self>) {
650        let (real, imag) = multi_op!(
651            inputs[(signal, float)],
652            outputs[(real, Float), (imag, Float)],
653            {
654                let res = B::rfft(signal, dim, n);
655                (res.0, res.1)
656            }
657        );
658
659        (real, imag)
660    }
661
662    fn irfft(
663        spectrum_re: FloatTensor<Self>,
664        spectrum_im: FloatTensor<Self>,
665        dim: usize,
666        n: Option<usize>,
667    ) -> FloatTensor<Self> {
668        multi_op!(
669            inputs[(spectrum_re, float), (spectrum_im, float)],
670            => Float,
671            {
672                B::irfft(spectrum_re, spectrum_im, dim, n)
673            }
674        )
675    }
676
677    fn has_ctc_loss_backward() -> bool {
678        // Dispatch routes per-tensor at runtime, but autodiff queries this flag
679        // statically. Returning `false` makes autodiff differentiate through
680        // the default decomposed forward, which is safe for every inner
681        // backend regardless of whether it has its own ctc_loss_backward.
682        false
683    }
684
685    fn ctc_loss(
686        log_probs: FloatTensor<Self>,
687        targets: IntTensor<Self>,
688        input_lengths: IntTensor<Self>,
689        target_lengths: IntTensor<Self>,
690        blank: usize,
691    ) -> FloatTensor<Self> {
692        multi_op!(
693            inputs[(log_probs, float), (targets, int), (input_lengths, int), (target_lengths, int)],
694            => Float,
695            B::ctc_loss(log_probs, targets, input_lengths, target_lengths, blank)
696        )
697    }
698
699    fn ctc_loss_backward(
700        log_probs: FloatTensor<Self>,
701        targets: IntTensor<Self>,
702        input_lengths: IntTensor<Self>,
703        target_lengths: IntTensor<Self>,
704        grad_loss: FloatTensor<Self>,
705        blank: usize,
706    ) -> FloatTensor<Self> {
707        multi_op!(
708            inputs[(log_probs, float), (targets, int), (input_lengths, int), (target_lengths, int), (grad_loss, float)],
709            => Float,
710            B::ctc_loss_backward(log_probs, targets, input_lengths, target_lengths, grad_loss, blank)
711        )
712    }
713
714    // TODO: linear ops
715    // fn linear(
716    //         x: FloatTensor<Self>,
717    //         weight: FloatTensor<Self>,
718    //         bias: Option<FloatTensor<Self>>,
719    //     ) -> FloatTensor<Self> {
720
721    // }
722}