MCPcopy Create free account
hub / github.com/modelscope/FunASR / get_transducer_task_io

Function get_transducer_task_io

funasr/models/transformer/utils/nets_utils.py:668–730  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

666
667
668def 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))

Callers 1

forwardMethod · 0.90

Calls 1

pad_listFunction · 0.70

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…