| 756 | ######################################### |
| 757 | ########### SIM Transformer ############# |
| 758 | class SIMTransformerBlock(nn.Module): |
| 759 | def __init__(self, dim, input_resolution, num_heads, win_size=10, shift_size=0, |
| 760 | mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0., |
| 761 | act_layer=nn.GELU, norm_layer=nn.LayerNorm,token_projection='linear',token_mlp='leff',se_layer=False): |
| 762 | super().__init__() |
| 763 | self.dim = dim |
| 764 | self.input_resolution = input_resolution |
| 765 | self.num_heads = num_heads |
| 766 | self.win_size = win_size |
| 767 | self.shift_size = shift_size |
| 768 | self.mlp_ratio = mlp_ratio |
| 769 | self.token_mlp = token_mlp |
| 770 | if min(self.input_resolution) <= self.win_size: |
| 771 | self.shift_size = 0 |
| 772 | self.win_size = min(self.input_resolution) |
| 773 | assert 0 <= self.shift_size < self.win_size, "shift_size must in 0-win_size" |
| 774 | |
| 775 | self.norm1 = norm_layer(dim) |
| 776 | self.attn = WindowAttention( |
| 777 | dim, win_size=to_2tuple(self.win_size), num_heads=num_heads, |
| 778 | qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop, |
| 779 | token_projection=token_projection,se_layer=se_layer) |
| 780 | |
| 781 | self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() |
| 782 | self.norm2 = norm_layer(dim) |
| 783 | mlp_hidden_dim = int(dim * mlp_ratio) |
| 784 | self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim,act_layer=act_layer, drop=drop) if token_mlp=='ffn' else LeFF(dim,mlp_hidden_dim,act_layer=act_layer, drop=drop) |
| 785 | self.CAB = CAB(dim, kernel_size=3, reduction=4, bias=False, act=nn.PReLU()) |
| 786 | |
| 787 | |
| 788 | def extra_repr(self) -> str: |
| 789 | return f"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, " \ |
| 790 | f"win_size={self.win_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}" |
| 791 | |
| 792 | def forward(self, x, xm, mask=None, img_size = (128, 128)): |
| 793 | B, L, C = x.shape |
| 794 | H = img_size[0] |
| 795 | W = img_size[1] |
| 796 | assert L == W * H, \ |
| 797 | f"Input image size ({H}*{W} doesn't match model ({L})." |
| 798 | |
| 799 | ## input mask |
| 800 | if mask != None: |
| 801 | input_mask = F.interpolate(mask, size=(H,W)).permute(0,2,3,1) |
| 802 | input_mask_windows = window_partition(input_mask, self.win_size) # nW, win_size, win_size, 1 |
| 803 | attn_mask = input_mask_windows.view(-1, self.win_size * self.win_size) # nW, win_size*win_size |
| 804 | attn_mask = attn_mask.unsqueeze(2)*attn_mask.unsqueeze(1) # nW, win_size*win_size, win_size*win_size |
| 805 | attn_mask = attn_mask.masked_fill(attn_mask!=0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0)) |
| 806 | else: |
| 807 | attn_mask = None |
| 808 | |
| 809 | ## shift mask |
| 810 | if self.shift_size > 0: |
| 811 | # calculate attention mask for SW-MSA |
| 812 | shift_mask = torch.zeros((1, H, W, 1)).type_as(x) |
| 813 | h_slices = (slice(0, -self.win_size), |
| 814 | slice(-self.win_size, -self.shift_size), |
| 815 | slice(-self.shift_size, None)) |