MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / AttnBlock

Class AttnBlock

sat/sgm/modules/autoencoding/vqvae/vqvae_blocks.py:114–155  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

112
113
114class AttnBlock(nn.Module):
115 def __init__(self, in_channels):
116 super().__init__()
117 self.in_channels = in_channels
118
119 self.norm = Normalize(in_channels)
120 self.q = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
121 self.k = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
122 self.v = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
123 self.proj_out = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
124
125 def forward(self, x):
126 h_ = x
127 h_ = self.norm(h_)
128 q = self.q(h_)
129 k = self.k(h_)
130 v = self.v(h_)
131
132 # compute attention
133 b, c, h, w = q.shape
134 q = q.reshape(b, c, h * w)
135 q = q.permute(0, 2, 1) # b,hw,c
136 k = k.reshape(b, c, h * w) # b,c,hw
137
138 # # original version, nan in fp16
139 # w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
140 # w_ = w_ * (int(c)**(-0.5))
141 # # implement c**-0.5 on q
142 q = q * (int(c) ** (-0.5))
143 w_ = torch.bmm(q, k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
144
145 w_ = torch.nn.functional.softmax(w_, dim=2)
146
147 # attend to values
148 v = v.reshape(b, c, h * w)
149 w_ = w_.permute(0, 2, 1) # b,hw,hw (first hw of k, second of q)
150 h_ = torch.bmm(v, w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
151 h_ = h_.reshape(b, c, h, w)
152
153 h_ = self.proj_out(h_)
154
155 return x + h_
156
157
158class Encoder(nn.Module):

Callers 2

__init__Method · 0.70
__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected