(batch_size, seq_len, cap_lens)
| 52 | |
| 53 | |
| 54 | def get_padding_mask(batch_size, seq_len, cap_lens): |
| 55 | cap_lens = cap_lens.data.tolist() |
| 56 | mask_2d = torch.ones((batch_size, seq_len, seq_len), dtype=torch.float32) |
| 57 | for i, cap_len in enumerate(cap_lens): |
| 58 | mask_2d[i, :, :cap_len] = 0 |
| 59 | return mask_2d.bool(), 1 - mask_2d[:, :, 0].clone() |
| 60 | |
| 61 | |
| 62 | class PositionalEncoding(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected