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