| 3 | |
| 4 | |
| 5 | def make_pad_mask(lengths: torch.Tensor, max_len: int = 0) -> torch.Tensor: |
| 6 | batch_size = lengths.size(0) |
| 7 | max_len = max_len if max_len > 0 else lengths.max().item() |
| 8 | seq_range = torch.arange(0, max_len, dtype=torch.int64, device=lengths.device) |
| 9 | seq_range_expand = seq_range.unsqueeze(0).expand(batch_size, max_len) |
| 10 | seq_length_expand = lengths.unsqueeze(-1) |
| 11 | mask = seq_range_expand >= seq_length_expand |
| 12 | return mask # (b, t) |
| 13 | |
| 14 | |
| 15 | def make_nonpad_mask(lengths: torch.Tensor, max_len: int = 0) -> torch.Tensor: |