r""" Args: x (`torch.Tensor`): Tensor of shape `(B, S, D)` of input tokens. mask (`torch.Tensor`): Boolean ensor of shape `(B, S)` indicating which tokens are not padding. Returns: `torch.Tensor`: `(
(self, x: torch.Tensor, mask: torch.BoolTensor)
| 2049 | return pooled |
| 2050 | |
| 2051 | def forward(self, x: torch.Tensor, mask: torch.BoolTensor) -> torch.Tensor: |
| 2052 | r""" |
| 2053 | Args: |
| 2054 | x (`torch.Tensor`): |
| 2055 | Tensor of shape `(B, S, D)` of input tokens. |
| 2056 | mask (`torch.Tensor`): |
| 2057 | Boolean ensor of shape `(B, S)` indicating which tokens are not padding. |
| 2058 | |
| 2059 | Returns: |
| 2060 | `torch.Tensor`: |
| 2061 | `(B, D)` tensor of pooled tokens. |
| 2062 | """ |
| 2063 | D = x.size(2) |
| 2064 | |
| 2065 | # Construct attention mask, shape: (B, 1, num_queries=1, num_keys=1+L). |
| 2066 | attn_mask = mask[:, None, None, :].bool() # (B, 1, 1, L). |
| 2067 | attn_mask = F.pad(attn_mask, (1, 0), value=True) # (B, 1, 1, 1+L). |
| 2068 | |
| 2069 | # Average non-padding token features. These will be used as the query. |
| 2070 | x_pool = self.pool_tokens(x, mask, keepdim=True) # (B, 1, D) |
| 2071 | |
| 2072 | # Concat pooled features to input sequence. |
| 2073 | x = torch.cat([x_pool, x], dim=1) # (B, L+1, D) |
| 2074 | |
| 2075 | # Compute queries, keys, values. Only the mean token is used to create a query. |
| 2076 | kv = self.to_kv(x) # (B, L+1, 2 * D) |
| 2077 | q = self.to_q(x[:, 0]) # (B, D) |
| 2078 | |
| 2079 | # Extract heads. |
| 2080 | head_dim = D // self.num_attention_heads |
| 2081 | kv = kv.unflatten(2, (2, self.num_attention_heads, head_dim)) # (B, 1+L, 2, H, head_dim) |
| 2082 | kv = kv.transpose(1, 3) # (B, H, 2, 1+L, head_dim) |
| 2083 | k, v = kv.unbind(2) # (B, H, 1+L, head_dim) |
| 2084 | q = q.unflatten(1, (self.num_attention_heads, head_dim)) # (B, H, head_dim) |
| 2085 | q = q.unsqueeze(2) # (B, H, 1, head_dim) |
| 2086 | |
| 2087 | # Compute attention. |
| 2088 | x = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=0.0) # (B, H, 1, head_dim) |
| 2089 | |
| 2090 | # Concatenate heads and run output. |
| 2091 | x = x.squeeze(2).flatten(1, 2) # (B, D = H * head_dim) |
| 2092 | x = self.to_out(x) |
| 2093 | return x |
| 2094 | |
| 2095 | |
| 2096 | def get_fourier_embeds_from_boundingbox(embed_dim, box): |
nothing calls this directly
no test coverage detected