x: B, H*W, C
(self, x, size)
| 636 | self.norm2 = norm_layer(dim) |
| 637 | |
| 638 | def forward(self, x, size): |
| 639 | """ |
| 640 | x: B, H*W, C |
| 641 | """ |
| 642 | H, W = size |
| 643 | # H = W = self.patches_resolution |
| 644 | B, L, C = x.shape |
| 645 | assert L == H * W, "flatten img_tokens has wrong size" |
| 646 | img = self.norm1(x) |
| 647 | qkv = self.qkv(img).reshape(B, -1, 3, C).permute(2, 0, 1, 3) |
| 648 | |
| 649 | if self.branch_num == 2: |
| 650 | x1 = self.attns[0](qkv[:, :, :, :C // 2], size) |
| 651 | x2 = self.attns[1](qkv[:, :, :, C // 2:], size) |
| 652 | attened_x = torch.cat([x1, x2], dim=2) |
| 653 | else: |
| 654 | attened_x = self.attns[0](qkv, size) |
| 655 | attened_x = self.proj(attened_x) |
| 656 | x = x + self.drop_path(attened_x) |
| 657 | x = x + self.drop_path(self.mlp(self.norm2(x))) |
| 658 | |
| 659 | return x |
| 660 | |
| 661 | |
| 662 | def img2windows(img, H_sp, W_sp): |
nothing calls this directly
no outgoing calls
no test coverage detected