x: (b h w c) mask_h: (n h h) mask_w: (n w w)
(self, x: torch.Tensor, rel_pos, chunkwise_recurrent=False, incremental_state=None)
| 160 | nn.init.constant_(self.out_proj.bias, 0.0) |
| 161 | |
| 162 | def forward(self, x: torch.Tensor, rel_pos, chunkwise_recurrent=False, incremental_state=None): |
| 163 | ''' |
| 164 | x: (b h w c) |
| 165 | mask_h: (n h h) |
| 166 | mask_w: (n w w) |
| 167 | ''' |
| 168 | bsz, h, w, _ = x.size() |
| 169 | |
| 170 | mask_h, mask_w = rel_pos |
| 171 | |
| 172 | q = self.q_proj(x) |
| 173 | k = self.k_proj(x) |
| 174 | v = self.v_proj(x) |
| 175 | lepe = self.lepe(v) |
| 176 | |
| 177 | k *= self.scaling |
| 178 | qr = q.view(bsz, h, w, self.num_heads, self.key_dim).permute(0, 3, 1, 2, 4) # (b n h w d1) |
| 179 | kr = k.view(bsz, h, w, self.num_heads, self.key_dim).permute(0, 3, 1, 2, 4) # (b n h w d1) |
| 180 | |
| 181 | ''' |
| 182 | qr: (b n h w d1) |
| 183 | kr: (b n h w d1) |
| 184 | v: (b h w n*d2) |
| 185 | ''' |
| 186 | |
| 187 | qr_w = qr.transpose(1, 2) # (b h n w d1) |
| 188 | kr_w = kr.transpose(1, 2) # (b h n w d1) |
| 189 | v = v.reshape(bsz, h, w, self.num_heads, -1).permute(0, 1, 3, 2, 4) # (b h n w d2) |
| 190 | |
| 191 | qk_mat_w = qr_w @ kr_w.transpose(-1, -2) # (b h n w w) |
| 192 | qk_mat_w = qk_mat_w + mask_w # (b h n w w) |
| 193 | qk_mat_w = torch.softmax(qk_mat_w, -1) # (b h n w w) |
| 194 | v = torch.matmul(qk_mat_w, v) # (b h n w d2) |
| 195 | |
| 196 | qr_h = qr.permute(0, 3, 1, 2, 4) # (b w n h d1) |
| 197 | kr_h = kr.permute(0, 3, 1, 2, 4) # (b w n h d1) |
| 198 | v = v.permute(0, 3, 2, 1, 4) # (b w n h d2) |
| 199 | |
| 200 | qk_mat_h = qr_h @ kr_h.transpose(-1, -2) # (b w n h h) |
| 201 | qk_mat_h = qk_mat_h + mask_h # (b w n h h) |
| 202 | qk_mat_h = torch.softmax(qk_mat_h, -1) # (b w n h h) |
| 203 | output = torch.matmul(qk_mat_h, v) # (b w n h d2) |
| 204 | |
| 205 | output = output.permute(0, 3, 1, 2, 4).flatten(-2, -1) # (b h w n*d2) |
| 206 | output = output + lepe #################################################################### |
| 207 | output = self.out_proj(output) |
| 208 | return output |
| 209 | |
| 210 | |
| 211 | class MaSA(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected