| 204 | return output |
| 205 | |
| 206 | class 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 | |
| 262 | class FeedForwardNetwork(nn.Module): |
| 263 | def __init__( |