x: B, H*W, C
(self, x, size)
| 707 | self.norm2 = norm_layer(dim) |
| 708 | |
| 709 | def forward(self, x, size): |
| 710 | """ |
| 711 | x: B, H*W, C |
| 712 | """ |
| 713 | H, W = size |
| 714 | B, L, C = x.shape |
| 715 | assert L == H * W, "flatten img_tokens has wrong size" |
| 716 | img = self.norm1(x) |
| 717 | qkv = self.qkv(img).reshape(B, -1, 3, C).permute(2, 0, 1, 3) |
| 718 | |
| 719 | if self.branch_num == 2: |
| 720 | x1 = self.attns[0](qkv[:, :, :, :C // 2], size) |
| 721 | x2 = self.attns[1](qkv[:, :, :, C // 2:], size) |
| 722 | attened_x = torch.cat([x1, x2], dim=2) |
| 723 | else: |
| 724 | attened_x = self.attns[0](qkv, size) |
| 725 | attened_x = self.proj(attened_x) |
| 726 | x = x + self.drop_path(attened_x) |
| 727 | x = x + self.drop_path(self.mlp(self.norm2(x))) |
| 728 | |
| 729 | return x |
| 730 | |
| 731 | |
| 732 | def img2windows(img, H_sp, W_sp): |
nothing calls this directly
no outgoing calls
no test coverage detected