Create source mask for given lengths. Reference: https://github.com/k2-fsa/icefall/blob/master/icefall/utils.py Args: lengths: Sequence lengths. (B,) Returns: : Mask for the sequence lengths. (B, max_len)
(lengths: torch.Tensor)
| 646 | |
| 647 | |
| 648 | def make_source_mask(lengths: torch.Tensor) -> torch.Tensor: |
| 649 | """Create source mask for given lengths. |
| 650 | |
| 651 | Reference: https://github.com/k2-fsa/icefall/blob/master/icefall/utils.py |
| 652 | |
| 653 | Args: |
| 654 | lengths: Sequence lengths. (B,) |
| 655 | |
| 656 | Returns: |
| 657 | : Mask for the sequence lengths. (B, max_len) |
| 658 | |
| 659 | """ |
| 660 | max_len = lengths.max() |
| 661 | batch_size = lengths.size(0) |
| 662 | |
| 663 | expanded_lengths = torch.arange(max_len).expand(batch_size, max_len).to(lengths) |
| 664 | |
| 665 | return expanded_lengths >= lengths.unsqueeze(1) |
| 666 | |
| 667 | |
| 668 | def get_transducer_task_io( |
no outgoing calls
no test coverage detected
searching dependent graphs…