MCPcopy Create free account
hub / github.com/GuoLanqing/ShadowFormer / SIMTransformerBlock

Class SIMTransformerBlock

model.py:758–894  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

756#########################################
757########### SIM Transformer #############
758class 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))

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected