MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / pool_tokens

Method pool_tokens

architecture/embeddings.py:2031–2049  ·  view source on GitHub ↗

Pool tokens in x using mask. NOTE: We assume x does not require gradients. Args: x: (B, L, D) tensor of tokens. mask: (B, L) boolean tensor indicating which tokens are not padding. Returns: pooled: (B, D) tensor of pooled tokens

(x: torch.Tensor, mask: torch.Tensor, *, keepdim=False)

Source from the content-addressed store, hash-verified

2029
2030 @staticmethod
2031 def pool_tokens(x: torch.Tensor, mask: torch.Tensor, *, keepdim=False) -> torch.Tensor:
2032 """
2033 Pool tokens in x using mask.
2034
2035 NOTE: We assume x does not require gradients.
2036
2037 Args:
2038 x: (B, L, D) tensor of tokens.
2039 mask: (B, L) boolean tensor indicating which tokens are not padding.
2040
2041 Returns:
2042 pooled: (B, D) tensor of pooled tokens.
2043 """
2044 assert x.size(1) == mask.size(1) # Expected mask to have same length as tokens.
2045 assert x.size(0) == mask.size(0) # Expected mask to have same batch size as tokens.
2046 mask = mask[:, :, None].to(dtype=x.dtype)
2047 mask = mask / mask.sum(dim=1, keepdim=True).clamp(min=1)
2048 pooled = (x * mask).sum(dim=1, keepdim=keepdim)
2049 return pooled
2050
2051 def forward(self, x: torch.Tensor, mask: torch.BoolTensor) -> torch.Tensor:
2052 r"""

Callers 1

forwardMethod · 0.95

Calls 1

toMethod · 0.45

Tested by

no test coverage detected