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 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 }