Sequence mask. Args: lengths: TODO. maxlen: TODO. dtype: TODO. device: Target device ("cuda:0", "cpu", etc.).
(lengths, maxlen=None, dtype=torch.float32, device=None)
| 324 | |
| 325 | |
| 326 | def sequence_mask(lengths, maxlen=None, dtype=torch.float32, device=None): |
| 327 | """Sequence mask. |
| 328 | |
| 329 | Args: |
| 330 | lengths: TODO. |
| 331 | maxlen: TODO. |
| 332 | dtype: TODO. |
| 333 | device: Target device ("cuda:0", "cpu", etc.). |
| 334 | """ |
| 335 | if maxlen is None: |
| 336 | maxlen = lengths.max() |
| 337 | row_vector = torch.arange(0, maxlen, 1).to(lengths.device) |
| 338 | matrix = torch.unsqueeze(lengths, dim=-1) |
| 339 | mask = row_vector < matrix |
| 340 | mask = mask.detach() |
| 341 | |
| 342 | return mask.type(dtype).to(device) if device is not None else mask.type(dtype) |
| 343 | |
| 344 | |
| 345 | class EncoderLayerSANM(nn.Module): |
no outgoing calls
no test coverage detected
searching dependent graphs…