(lengths, max_len=None, is_2d=True)
| 34 | |
| 35 | |
| 36 | def sequence_mask(lengths, max_len=None, is_2d=True): |
| 37 | batch_size = lengths.numel() |
| 38 | max_len = max_len or lengths.max() |
| 39 | mask = (torch.arange(0, max_len, device=lengths.device) |
| 40 | .type_as(lengths) |
| 41 | .repeat(batch_size, 1) |
| 42 | .lt(lengths.unsqueeze(1))) |
| 43 | if is_2d: |
| 44 | return mask |
| 45 | else: |
| 46 | mask = mask.view(-1, 1, 1, max_len) |
| 47 | m2 = mask.transpose(2, 3) |
| 48 | return mask * m2 |
| 49 | |
| 50 | |
| 51 | def main(): |
no test coverage detected