Convert additive mask to binary [B, seq, 1] and zero out padding.
(encoded, encoded_mask)
| 31 | |
| 32 | |
| 33 | def _to_binary_mask(encoded, encoded_mask): |
| 34 | """Convert additive mask to binary [B, seq, 1] and zero out padding.""" |
| 35 | binary_mask = (encoded_mask < 0.000001).to(torch.int64) |
| 36 | binary_mask = binary_mask.reshape(encoded.shape[0], encoded.shape[1], 1) |
| 37 | return encoded * binary_mask, binary_mask |
| 38 | |
| 39 | |
| 40 | def _rescale_norm(x: torch.Tensor, target_dim: int, source_dim: int) -> torch.Tensor: |
no outgoing calls
no test coverage detected