MCPcopy Create free account
hub / github.com/awslabs/gap-text2sql / masked_softmax

Function masked_softmax

relogic/logickit/utils/utils.py:227–245  ·  view source on GitHub ↗

``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)

Source from the content-addressed store, hash-verified

225
226
227def 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
247def masked_log_softmax(vector: torch.Tensor, mask: torch.Tensor, dim: int = -1) -> torch.Tensor:
248 """

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected