MCPcopy Create free account
hub / github.com/buaacxf/VIPTR / MaSA

Class MaSA

modules/VIPTRv2.py:206–260  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

204 return output
205
206class MaSA(nn.Module):
207
208 def __init__(self, embed_dim, num_heads, value_factor=1):
209 super().__init__()
210 self.factor = value_factor
211 self.embed_dim = embed_dim
212 self.num_heads = num_heads
213 self.head_dim = self.embed_dim * self.factor // num_heads
214 self.key_dim = self.embed_dim // num_heads
215 self.scaling = self.key_dim ** -0.5
216 self.q_proj = nn.Linear(embed_dim, embed_dim, bias=True)
217 self.k_proj = nn.Linear(embed_dim, embed_dim, bias=True)
218 self.v_proj = nn.Linear(embed_dim, embed_dim * self.factor, bias=True)
219 self.lepe = DWConv2d(embed_dim, 5, 1, 2)
220 self.out_proj = nn.Linear(embed_dim * self.factor, embed_dim, bias=True)
221 self.reset_parameters()
222
223 def reset_parameters(self):
224 nn.init.xavier_normal_(self.q_proj.weight, gain=2 ** -2.5)
225 nn.init.xavier_normal_(self.k_proj.weight, gain=2 ** -2.5)
226 nn.init.xavier_normal_(self.v_proj.weight, gain=2 ** -2.5)
227 nn.init.xavier_normal_(self.out_proj.weight)
228 nn.init.constant_(self.out_proj.bias, 0.0)
229
230 def forward(self, x: torch.Tensor, rel_pos, chunkwise_recurrent=False, incremental_state=None):
231 '''
232 x: (b h w c)
233 rel_pos: mask: (n l l)
234 '''
235 bsz, h, w, _ = x.size()
236 mask = rel_pos
237
238 assert h * w == mask.size(1)
239
240 q = self.q_proj(x)
241 k = self.k_proj(x)
242 v = self.v_proj(x)
243 lepe = self.lepe(v)
244
245 k *= self.scaling
246 qr = q.view(bsz, h, w, self.num_heads, -1).permute(0, 3, 1, 2, 4) # (b n h w d1)
247 kr = k.view(bsz, h, w, self.num_heads, -1).permute(0, 3, 1, 2, 4) # (b n h w d1)
248
249 qr = qr.flatten(2, 3) # (b n l d1)
250 kr = kr.flatten(2, 3) # (b n l d1)
251 vr = v.reshape(bsz, h, w, self.num_heads, -1).permute(0, 3, 1, 2, 4) # (b n h w d2)
252 vr = vr.flatten(2, 3) # (b n l d2)
253 qk_mat = qr @ kr.transpose(-1, -2) # (b n l l)
254 qk_mat = qk_mat + mask # (b n l l)
255 qk_mat = torch.softmax(qk_mat, -1) # (b n l l)
256 output = torch.matmul(qk_mat, vr) # (b n l d2)
257 output = output.transpose(1, 2).reshape(bsz, h, w, -1) # (b h w n*d2)
258 output = output + lepe
259 output = self.out_proj(output)
260 return output
261
262class FeedForwardNetwork(nn.Module):
263 def __init__(

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected