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