``torch.nn.functional.softmax(vector)`` does not work if some elements of ``vector`` should be masked. This performs a softmax on just the non-masked positions of ``vector``. Passing ``None`` in for the mask is also acceptable, which is just the regular softmax.
(vector: torch.Tensor,
mask: torch.Tensor,
dim: int = -1,
mask_fill_value: float = -1e32)
| 225 | |
| 226 | |
| 227 | def masked_softmax(vector: torch.Tensor, |
| 228 | mask: torch.Tensor, |
| 229 | dim: int = -1, |
| 230 | mask_fill_value: float = -1e32) -> torch.Tensor: |
| 231 | """ |
| 232 | ``torch.nn.functional.softmax(vector)`` does not work if some elements of ``vector`` should be |
| 233 | masked. This performs a softmax on just the non-masked positions of ``vector``. Passing ``None`` |
| 234 | in for the mask is also acceptable, which is just the regular softmax. |
| 235 | |
| 236 | """ |
| 237 | if mask is None: |
| 238 | result = torch.softmax(vector, dim=dim) |
| 239 | else: |
| 240 | mask = mask.float() |
| 241 | while mask.dim() < vector.dim(): |
| 242 | mask = mask.unsqueeze(1) |
| 243 | masked_vector = vector.masked_fill((1 - mask).bool(), mask_fill_value) |
| 244 | result = torch.softmax(masked_vector, dim=dim) |
| 245 | return result |
| 246 | |
| 247 | def masked_log_softmax(vector: torch.Tensor, mask: torch.Tensor, dim: int = -1) -> torch.Tensor: |
| 248 | """ |
nothing calls this directly
no outgoing calls
no test coverage detected