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

Class Attention

preprocess/SpaTrackV2_code/models/blocks.py:90–132  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

88 return x
89
90class Attention(nn.Module):
91 def __init__(self, query_dim, context_dim=None,
92 num_heads=8, dim_head=48, qkv_bias=False, flash=False):
93 super().__init__()
94 inner_dim = self.inner_dim = dim_head * num_heads
95 context_dim = default(context_dim, query_dim)
96 self.scale = dim_head**-0.5
97 self.heads = num_heads
98 self.flash = flash
99
100 self.to_q = nn.Linear(query_dim, inner_dim, bias=qkv_bias)
101 self.to_kv = nn.Linear(context_dim, inner_dim * 2, bias=qkv_bias)
102 self.to_out = nn.Linear(inner_dim, query_dim)
103
104 def forward(self, x, context=None, attn_bias=None):
105 B, N1, _ = x.shape
106 C = self.inner_dim
107 h = self.heads
108 q = self.to_q(x).reshape(B, N1, h, C // h).permute(0, 2, 1, 3)
109 context = default(context, x)
110 k, v = self.to_kv(context).chunk(2, dim=-1)
111
112 N2 = context.shape[1]
113 k = k.reshape(B, N2, h, C // h).permute(0, 2, 1, 3)
114 v = v.reshape(B, N2, h, C // h).permute(0, 2, 1, 3)
115
116 with torch.autocast("cuda", enabled=True, dtype=torch.bfloat16):
117 if self.flash==False:
118 sim = (q @ k.transpose(-2, -1)) * self.scale
119 if attn_bias is not None:
120 sim = sim + attn_bias
121 if sim.abs().max()>1e2:
122 import pdb; pdb.set_trace()
123 attn = sim.softmax(dim=-1)
124 x = (attn @ v).transpose(1, 2).reshape(B, N1, C)
125 else:
126 input_args = [x.contiguous() for x in [q, k, v]]
127 x = F.scaled_dot_product_attention(*input_args).permute(0,2,1,3).reshape(B,N1,-1) # type: ignore
128
129 if self.to_out.bias.dtype != x.dtype:
130 x = x.to(self.to_out.bias.dtype)
131
132 return self.to_out(x)
133
134
135class VGG19(nn.Module):

Callers 2

__init__Method · 0.70
__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected