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