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

Method forward

architecture/embeddings.py:2051–2093  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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
2096def get_fourier_embeds_from_boundingbox(embed_dim, box):

Callers

nothing calls this directly

Calls 1

pool_tokensMethod · 0.95

Tested by

no test coverage detected