Get Transducer loss I/O. Args: labels: Label ID sequences. (B, L) encoder_out_lens: Encoder output lengths. (B,) ignore_id: Padding symbol ID. blank_id: Blank symbol ID. Returns: decoder_in: Decoder inputs. (B, U) target: Target label ID sequ
(
labels: torch.Tensor,
encoder_out_lens: torch.Tensor,
ignore_id: int = -1,
blank_id: int = 0,
)
| 666 | |
| 667 | |
| 668 | def get_transducer_task_io( |
| 669 | labels: torch.Tensor, |
| 670 | encoder_out_lens: torch.Tensor, |
| 671 | ignore_id: int = -1, |
| 672 | blank_id: int = 0, |
| 673 | ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: |
| 674 | """Get Transducer loss I/O. |
| 675 | |
| 676 | Args: |
| 677 | labels: Label ID sequences. (B, L) |
| 678 | encoder_out_lens: Encoder output lengths. (B,) |
| 679 | ignore_id: Padding symbol ID. |
| 680 | blank_id: Blank symbol ID. |
| 681 | |
| 682 | Returns: |
| 683 | decoder_in: Decoder inputs. (B, U) |
| 684 | target: Target label ID sequences. (B, U) |
| 685 | t_len: Time lengths. (B,) |
| 686 | u_len: Label lengths. (B,) |
| 687 | |
| 688 | """ |
| 689 | |
| 690 | def pad_list(labels: List[torch.Tensor], padding_value: int = 0): |
| 691 | """Create padded batch of labels from a list of labels sequences. |
| 692 | |
| 693 | Args: |
| 694 | labels: Labels sequences. [B x (?)] |
| 695 | padding_value: Padding value. |
| 696 | |
| 697 | Returns: |
| 698 | labels: Batch of padded labels sequences. (B,) |
| 699 | |
| 700 | """ |
| 701 | batch_size = len(labels) |
| 702 | |
| 703 | padded = ( |
| 704 | labels[0] |
| 705 | .new(batch_size, max(x.size(0) for x in labels), *labels[0].size()[1:]) |
| 706 | .fill_(padding_value) |
| 707 | ) |
| 708 | |
| 709 | for i in range(batch_size): |
| 710 | padded[i, : labels[i].size(0)] = labels[i] |
| 711 | |
| 712 | return padded |
| 713 | |
| 714 | device = labels.device |
| 715 | |
| 716 | labels_unpad = [y[y != ignore_id] for y in labels] |
| 717 | blank = labels[0].new([blank_id]) |
| 718 | |
| 719 | decoder_in = pad_list( |
| 720 | [torch.cat([blank, label], dim=0) for label in labels_unpad], blank_id |
| 721 | ).to(device) |
| 722 | |
| 723 | target = pad_list(labels_unpad, blank_id).type(torch.int32).to(device) |
| 724 | |
| 725 | encoder_out_lens = list(map(int, encoder_out_lens)) |