Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
swin.py892 linesDownload Raw Back to backbone
1# --------------------------------------------------------2# Swin Transformer3# Copyright (c) 2021 Microsoft4# Licensed under The MIT License [see LICENSE for details]5# Written by Ze Liu, Yutong Lin, Yixuan Wei6# --------------------------------------------------------7 8# Copyright (c) Facebook, Inc. and its affiliates.9# Modified by Bowen Cheng from https://github.com/SwinTransformer/Swin-Transformer-Semantic-Segmentation/blob/main/mmseg/models/backbones/swin_transformer.py10import logging11import numpy as np12import torch13import torch.nn as nn14import torch.nn.functional as F15import torch.utils.checkpoint as checkpoint16from timm.models.layers import DropPath, to_2tuple, trunc_normal_17 18from detectron2.modeling import Backbone, ShapeSpec19from detectron2.utils.file_io import PathManager20 21from .registry import register_backbone22 23logger = logging.getLogger(__name__)24 25 26class Mlp(nn.Module):27    """Multilayer perceptron."""28 29    def __init__(30        self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.031    ):32        super().__init__()33        out_features = out_features or in_features34        hidden_features = hidden_features or in_features35        self.fc1 = nn.Linear(in_features, hidden_features)36        self.act = act_layer()37        self.fc2 = nn.Linear(hidden_features, out_features)38        self.drop = nn.Dropout(drop)39 40    def forward(self, x):41        x = self.fc1(x)42        x = self.act(x)43        x = self.drop(x)44        x = self.fc2(x)45        x = self.drop(x)46        return x47 48 49def window_partition(x, window_size):50    """51    Args:52        x: (B, H, W, C)53        window_size (int): window size54    Returns:55        windows: (num_windows*B, window_size, window_size, C)56    """57    B, H, W, C = x.shape58    x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)59    windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)60    return windows61 62 63def window_reverse(windows, window_size, H, W):64    """65    Args:66        windows: (num_windows*B, window_size, window_size, C)67        window_size (int): Window size68        H (int): Height of image69        W (int): Width of image70    Returns:71        x: (B, H, W, C)72    """73    B = int(windows.shape[0] / (H * W / window_size / window_size))74    x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)75    x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)76    return x77 78 79class WindowAttention(nn.Module):80    """Window based multi-head self attention (W-MSA) module with relative position bias.81    It supports both of shifted and non-shifted window.82    Args:83        dim (int): Number of input channels.84        window_size (tuple[int]): The height and width of the window.85        num_heads (int): Number of attention heads.86        qkv_bias (bool, optional):  If True, add a learnable bias to query, key, value. Default: True87        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set88        attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.089        proj_drop (float, optional): Dropout ratio of output. Default: 0.090    """91 92    def __init__(93        self,94        dim,95        window_size,96        num_heads,97        qkv_bias=True,98        qk_scale=None,99        attn_drop=0.0,100        proj_drop=0.0,101    ):102 103        super().__init__()104        self.dim = dim105        self.window_size = window_size  # Wh, Ww106        self.num_heads = num_heads107        head_dim = dim // num_heads108        self.scale = qk_scale or head_dim ** -0.5109 110        # define a parameter table of relative position bias111        self.relative_position_bias_table = nn.Parameter(112            torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)113        )  # 2*Wh-1 * 2*Ww-1, nH114 115        # get pair-wise relative position index for each token inside the window116        coords_h = torch.arange(self.window_size[0])117        coords_w = torch.arange(self.window_size[1])118        coords = torch.stack(torch.meshgrid([coords_h, coords_w]))  # 2, Wh, Ww119        coords_flatten = torch.flatten(coords, 1)  # 2, Wh*Ww120        relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]  # 2, Wh*Ww, Wh*Ww121        relative_coords = relative_coords.permute(1, 2, 0).contiguous()  # Wh*Ww, Wh*Ww, 2122        relative_coords[:, :, 0] += self.window_size[0] - 1  # shift to start from 0123        relative_coords[:, :, 1] += self.window_size[1] - 1124        relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1125        relative_position_index = relative_coords.sum(-1)  # Wh*Ww, Wh*Ww126        self.register_buffer("relative_position_index", relative_position_index)127 128        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)129        self.attn_drop = nn.Dropout(attn_drop)130        self.proj = nn.Linear(dim, dim)131        self.proj_drop = nn.Dropout(proj_drop)132 133        trunc_normal_(self.relative_position_bias_table, std=0.02)134        self.softmax = nn.Softmax(dim=-1)135 136    def forward(self, x, mask=None):137        """Forward function.138        Args:139            x: input features with shape of (num_windows*B, N, C)140            mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None141        """142        B_, N, C = x.shape143        qkv = (144            self.qkv(x)145            .reshape(B_, N, 3, self.num_heads, C // self.num_heads)146            .permute(2, 0, 3, 1, 4)147        )148        q, k, v = qkv[0], qkv[1], qkv[2]  # make torchscript happy (cannot use tensor as tuple)149 150        q = q * self.scale151        attn = q @ k.transpose(-2, -1)152        153        relative_position_bias = self.relative_position_bias_table[154            self.relative_position_index.view(-1)155        ].view(156            self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1157        )  # Wh*Ww,Wh*Ww,nH158        relative_position_bias = relative_position_bias.permute(159            2, 0, 1160        ).contiguous()  # nH, Wh*Ww, Wh*Ww161        attn = attn + relative_position_bias.unsqueeze(0)162 163        if mask is not None:164            nW = mask.shape[0]165            attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)166            attn = attn.view(-1, self.num_heads, N, N)167            attn = self.softmax(attn)168        else:169            attn = self.softmax(attn)170 171        attn = self.attn_drop(attn)172 173        x = (attn @ v).transpose(1, 2).reshape(B_, N, C)174        x = self.proj(x)175        x = self.proj_drop(x)176 177        return x178 179 180class SwinTransformerBlock(nn.Module):181    """Swin Transformer Block.182    Args:183        dim (int): Number of input channels.184        num_heads (int): Number of attention heads.185        window_size (int): Window size.186        shift_size (int): Shift size for SW-MSA.187        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.188        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True189        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.190        drop (float, optional): Dropout rate. Default: 0.0191        attn_drop (float, optional): Attention dropout rate. Default: 0.0192        drop_path (float, optional): Stochastic depth rate. Default: 0.0193        act_layer (nn.Module, optional): Activation layer. Default: nn.GELU194        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm195    """196 197    def __init__(198        self,199        dim,200        num_heads,201        window_size=7,202        shift_size=0,203        mlp_ratio=4.0,204        qkv_bias=True,205        qk_scale=None,206        drop=0.0,207        attn_drop=0.0,208        drop_path=0.0,209        act_layer=nn.GELU,210        norm_layer=nn.LayerNorm,211    ):212        super().__init__()213        self.dim = dim214        self.num_heads = num_heads215        self.window_size = window_size216        self.shift_size = shift_size217        self.mlp_ratio = mlp_ratio218        assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"219 220        self.norm1 = norm_layer(dim)221        self.attn = WindowAttention(222            dim,223            window_size=to_2tuple(self.window_size),224            num_heads=num_heads,225            qkv_bias=qkv_bias,226            qk_scale=qk_scale,227            attn_drop=attn_drop,228            proj_drop=drop,229        )230 231        self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()232        self.norm2 = norm_layer(dim)233        mlp_hidden_dim = int(dim * mlp_ratio)234        self.mlp = Mlp(235            in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop236        )237 238        self.H = None239        self.W = None240 241    def forward(self, x, mask_matrix):242        """Forward function.243        Args:244            x: Input feature, tensor size (B, H*W, C).245            H, W: Spatial resolution of the input feature.246            mask_matrix: Attention mask for cyclic shift.247        """248        B, L, C = x.shape249        H, W = self.H, self.W250        assert L == H * W, "input feature has wrong size"251 252        # HACK model will not upsampling253        # if min([H, W]) <= self.window_size:254            # if window size is larger than input resolution, we don't partition windows255            # self.shift_size = 0256            # self.window_size = min([H,W])257 258        shortcut = x259        x = self.norm1(x)260        x = x.view(B, H, W, C)261 262        # pad feature maps to multiples of window size263        pad_l = pad_t = 0264        pad_r = (self.window_size - W % self.window_size) % self.window_size265        pad_b = (self.window_size - H % self.window_size) % self.window_size266        x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))267        _, Hp, Wp, _ = x.shape268 269        # cyclic shift270        if self.shift_size > 0:271            shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))272            attn_mask = mask_matrix273        else:274            shifted_x = x275            attn_mask = None276 277        # partition windows278        x_windows = window_partition(279            shifted_x, self.window_size280        )  # nW*B, window_size, window_size, C281        x_windows = x_windows.view(282            -1, self.window_size * self.window_size, C283        )  # nW*B, window_size*window_size, C284 285        # W-MSA/SW-MSA286        attn_windows = self.attn(x_windows, mask=attn_mask)  # nW*B, window_size*window_size, C287 288        # merge windows289        attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)290        shifted_x = window_reverse(attn_windows, self.window_size, Hp, Wp)  # B H' W' C291 292        # reverse cyclic shift293        if self.shift_size > 0:294            x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))295        else:296            x = shifted_x297 298        if pad_r > 0 or pad_b > 0:299            x = x[:, :H, :W, :].contiguous()300 301        x = x.view(B, H * W, C)302 303        # FFN304        x = shortcut + self.drop_path(x)305        x = x + self.drop_path(self.mlp(self.norm2(x)))306        return x307 308 309class PatchMerging(nn.Module):310    """Patch Merging Layer311    Args:312        dim (int): Number of input channels.313        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm314    """315 316    def __init__(self, dim, norm_layer=nn.LayerNorm):317        super().__init__()318        self.dim = dim319        self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)320        self.norm = norm_layer(4 * dim)321 322    def forward(self, x, H, W):323        """Forward function.324        Args:325            x: Input feature, tensor size (B, H*W, C).326            H, W: Spatial resolution of the input feature.327        """328        B, L, C = x.shape329        assert L == H * W, "input feature has wrong size"330 331        x = x.view(B, H, W, C)332 333        # padding334        pad_input = (H % 2 == 1) or (W % 2 == 1)335        if pad_input:336            x = F.pad(x, (0, 0, 0, W % 2, 0, H % 2))337 338        x0 = x[:, 0::2, 0::2, :]  # B H/2 W/2 C339        x1 = x[:, 1::2, 0::2, :]  # B H/2 W/2 C340        x2 = x[:, 0::2, 1::2, :]  # B H/2 W/2 C341        x3 = x[:, 1::2, 1::2, :]  # B H/2 W/2 C342        x = torch.cat([x0, x1, x2, x3], -1)  # B H/2 W/2 4*C343        x = x.view(B, -1, 4 * C)  # B H/2*W/2 4*C344 345        x = self.norm(x)346        x = self.reduction(x)347 348        return x349 350 351class BasicLayer(nn.Module):352    """A basic Swin Transformer layer for one stage.353    Args:354        dim (int): Number of feature channels355        depth (int): Depths of this stage.356        num_heads (int): Number of attention head.357        window_size (int): Local window size. Default: 7.358        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.359        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True360        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.361        drop (float, optional): Dropout rate. Default: 0.0362        attn_drop (float, optional): Attention dropout rate. Default: 0.0363        drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0364        norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm365        downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None366        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.367    """368 369    def __init__(370        self,371        dim,372        depth,373        num_heads,374        window_size=7,375        mlp_ratio=4.0,376        qkv_bias=True,377        qk_scale=None,378        drop=0.0,379        attn_drop=0.0,380        drop_path=0.0,381        norm_layer=nn.LayerNorm,382        downsample=None,383        use_checkpoint=False,384    ):385        super().__init__()386        self.window_size = window_size387        self.shift_size = window_size // 2388        self.depth = depth389        self.use_checkpoint = use_checkpoint390 391        # build blocks392        self.blocks = nn.ModuleList(393            [394                SwinTransformerBlock(395                    dim=dim,396                    num_heads=num_heads,397                    window_size=window_size,398                    shift_size=0 if (i % 2 == 0) else window_size // 2,399                    mlp_ratio=mlp_ratio,400                    qkv_bias=qkv_bias,401                    qk_scale=qk_scale,402                    drop=drop,403                    attn_drop=attn_drop,404                    drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,405                    norm_layer=norm_layer,406                )407                for i in range(depth)408            ]409        )410 411        # patch merging layer412        if downsample is not None:413            self.downsample = downsample(dim=dim, norm_layer=norm_layer)414        else:415            self.downsample = None416 417    def forward(self, x, H, W):418        """Forward function.419        Args:420            x: Input feature, tensor size (B, H*W, C).421            H, W: Spatial resolution of the input feature.422        """423 424        # calculate attention mask for SW-MSA425        Hp = int(np.ceil(H / self.window_size)) * self.window_size426        Wp = int(np.ceil(W / self.window_size)) * self.window_size427        img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device)  # 1 Hp Wp 1428        h_slices = (429            slice(0, -self.window_size),430            slice(-self.window_size, -self.shift_size),431            slice(-self.shift_size, None),432        )433        w_slices = (434            slice(0, -self.window_size),435            slice(-self.window_size, -self.shift_size),436            slice(-self.shift_size, None),437        )438        cnt = 0439        for h in h_slices:440            for w in w_slices:441                img_mask[:, h, w, :] = cnt442                cnt += 1443 444        mask_windows = window_partition(445            img_mask, self.window_size446        )  # nW, window_size, window_size, 1447        mask_windows = mask_windows.view(-1, self.window_size * self.window_size)448        attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)449        attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(450            attn_mask == 0, float(0.0)451        ).type(x.dtype)452        453        for blk in self.blocks:454            blk.H, blk.W = H, W455            if self.use_checkpoint:456                x = checkpoint.checkpoint(blk, x, attn_mask)457            else:458                x = blk(x, attn_mask)459        if self.downsample is not None:460            x_down = self.downsample(x, H, W)461            Wh, Ww = (H + 1) // 2, (W + 1) // 2462            return x, H, W, x_down, Wh, Ww463        else:464            return x, H, W, x, H, W465 466 467class PatchEmbed(nn.Module):468    """Image to Patch Embedding469    Args:470        patch_size (int): Patch token size. Default: 4.471        in_chans (int): Number of input image channels. Default: 3.472        embed_dim (int): Number of linear projection output channels. Default: 96.473        norm_layer (nn.Module, optional): Normalization layer. Default: None474    """475 476    def __init__(self, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None):477        super().__init__()478        patch_size = to_2tuple(patch_size)479        self.patch_size = patch_size480 481        self.in_chans = in_chans482        self.embed_dim = embed_dim483 484        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)485        if norm_layer is not None:486            self.norm = norm_layer(embed_dim)487        else:488            self.norm = None489 490    def forward(self, x):491        """Forward function."""492        # padding493        _, _, H, W = x.size()494        if W % self.patch_size[1] != 0:495            x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1]))496        if H % self.patch_size[0] != 0:497            x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0]))498 499        x = self.proj(x)  # B C Wh Ww500        if self.norm is not None:501            Wh, Ww = x.size(2), x.size(3)502            x = x.flatten(2).transpose(1, 2)503            x = self.norm(x)504            x = x.transpose(1, 2).view(-1, self.embed_dim, Wh, Ww)505 506        return x507 508 509class SwinTransformer(nn.Module):510    """Swin Transformer backbone.511        A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows`  -512          https://arxiv.org/pdf/2103.14030513    Args:514        pretrain_img_size (int): Input image size for training the pretrained model,515            used in absolute postion embedding. Default 224.516        patch_size (int | tuple(int)): Patch size. Default: 4.517        in_chans (int): Number of input image channels. Default: 3.518        embed_dim (int): Number of linear projection output channels. Default: 96.519        depths (tuple[int]): Depths of each Swin Transformer stage.520        num_heads (tuple[int]): Number of attention head of each stage.521        window_size (int): Window size. Default: 7.522        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.523        qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True524        qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.525        drop_rate (float): Dropout rate.526        attn_drop_rate (float): Attention dropout rate. Default: 0.527        drop_path_rate (float): Stochastic depth rate. Default: 0.2.528        norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.529        ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.530        patch_norm (bool): If True, add normalization after patch embedding. Default: True.531        out_indices (Sequence[int]): Output from which stages.532        frozen_stages (int): Stages to be frozen (stop grad and set eval mode).533            -1 means not freezing any parameters.534        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.535    """536 537    def __init__(538        self,539        pretrain_img_size=224,540        patch_size=4,541        in_chans=3,542        embed_dim=96,543        depths=[2, 2, 6, 2],544        num_heads=[3, 6, 12, 24],545        window_size=7,546        mlp_ratio=4.0,547        qkv_bias=True,548        qk_scale=None,549        drop_rate=0.0,550        attn_drop_rate=0.0,551        drop_path_rate=0.2,552        norm_layer=nn.LayerNorm,553        ape=False,554        patch_norm=True,555        out_indices=(0, 1, 2, 3),556        frozen_stages=-1,557        use_checkpoint=False,558    ):559        super().__init__()560 561        self.pretrain_img_size = pretrain_img_size562        self.num_layers = len(depths)563        self.embed_dim = embed_dim564        self.ape = ape565        self.patch_norm = patch_norm566        self.out_indices = out_indices567        self.frozen_stages = frozen_stages568 569        # split image into non-overlapping patches570        self.patch_embed = PatchEmbed(571            patch_size=patch_size,572            in_chans=in_chans,573            embed_dim=embed_dim,574            norm_layer=norm_layer if self.patch_norm else None,575        )576 577        # absolute position embedding578        if self.ape:579            pretrain_img_size = to_2tuple(pretrain_img_size)580            patch_size = to_2tuple(patch_size)581            patches_resolution = [582                pretrain_img_size[0] // patch_size[0],583                pretrain_img_size[1] // patch_size[1],584            ]585 586            self.absolute_pos_embed = nn.Parameter(587                torch.zeros(1, embed_dim, patches_resolution[0], patches_resolution[1])588            )589            trunc_normal_(self.absolute_pos_embed, std=0.02)590 591        self.pos_drop = nn.Dropout(p=drop_rate)592 593        # stochastic depth594        dpr = [595            x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))596        ]  # stochastic depth decay rule597 598        # build layers599        self.layers = nn.ModuleList()600        for i_layer in range(self.num_layers):601            layer = BasicLayer(602                dim=int(embed_dim * 2 ** i_layer),603                depth=depths[i_layer],604                num_heads=num_heads[i_layer],605                window_size=window_size,606                mlp_ratio=mlp_ratio,607                qkv_bias=qkv_bias,608                qk_scale=qk_scale,609                drop=drop_rate,610                attn_drop=attn_drop_rate,611                drop_path=dpr[sum(depths[:i_layer]) : sum(depths[: i_layer + 1])],612                norm_layer=norm_layer,613                downsample=PatchMerging if (i_layer < self.num_layers - 1) else None,614                use_checkpoint=use_checkpoint,615            )616            self.layers.append(layer)617 618        num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]619        self.num_features = num_features620 621        # add a norm layer for each output622        for i_layer in out_indices:623            layer = norm_layer(num_features[i_layer])624            layer_name = f"norm{i_layer}"625            self.add_module(layer_name, layer)626 627        self._freeze_stages()628 629    def _freeze_stages(self):630        if self.frozen_stages >= 0:631            self.patch_embed.eval()632            for param in self.patch_embed.parameters():633                param.requires_grad = False634 635        if self.frozen_stages >= 1 and self.ape:636            self.absolute_pos_embed.requires_grad = False637 638        if self.frozen_stages >= 2:639            self.pos_drop.eval()640            for i in range(0, self.frozen_stages - 1):641                m = self.layers[i]642                m.eval()643                for param in m.parameters():644                    param.requires_grad = False645 646    def init_weights(self, pretrained=None):647        """Initialize the weights in backbone.648        Args:649            pretrained (str, optional): Path to pre-trained weights.650                Defaults to None.651        """652 653        def _init_weights(m):654            if isinstance(m, nn.Linear):655                trunc_normal_(m.weight, std=0.02)656                if isinstance(m, nn.Linear) and m.bias is not None:657                    nn.init.constant_(m.bias, 0)658            elif isinstance(m, nn.LayerNorm):659                nn.init.constant_(m.bias, 0)660                nn.init.constant_(m.weight, 1.0)661 662 663    def load_weights(self, pretrained_dict=None, pretrained_layers=[], verbose=True):664        model_dict = self.state_dict()665        pretrained_dict = {666            k: v for k, v in pretrained_dict.items()667            if k in model_dict.keys()668        }669        need_init_state_dict = {}670        for k, v in pretrained_dict.items():671            need_init = (672                    (673                            k.split('.')[0] in pretrained_layers674                            or pretrained_layers[0] == '*'675                    )676                    and 'relative_position_index' not in k677                    and 'attn_mask' not in k678            )679 680            if need_init:681                # if verbose:682                #     logger.info(f'=> init {k} from {pretrained}')683 684                if 'relative_position_bias_table' in k and v.size() != model_dict[k].size():685                    relative_position_bias_table_pretrained = v686                    relative_position_bias_table_current = model_dict[k]687                    L1, nH1 = relative_position_bias_table_pretrained.size()688                    L2, nH2 = relative_position_bias_table_current.size()689                    if nH1 != nH2:690                        logger.info(f"Error in loading {k}, passing")691                    else:692                        if L1 != L2:693                            logger.info(694                                '=> load_pretrained: resized variant: {} to {}'695                                    .format((L1, nH1), (L2, nH2))696                            )697                            S1 = int(L1 ** 0.5)698                            S2 = int(L2 ** 0.5)699                            relative_position_bias_table_pretrained_resized = torch.nn.functional.interpolate(700                                relative_position_bias_table_pretrained.permute(1, 0).view(1, nH1, S1, S1),701                                size=(S2, S2),702                                mode='bicubic')703                            v = relative_position_bias_table_pretrained_resized.view(nH2, L2).permute(1, 0)704 705                if 'absolute_pos_embed' in k and v.size() != model_dict[k].size():706                    absolute_pos_embed_pretrained = v707                    absolute_pos_embed_current = model_dict[k]708                    _, L1, C1 = absolute_pos_embed_pretrained.size()709                    _, L2, C2 = absolute_pos_embed_current.size()710                    if C1 != C1:711                        logger.info(f"Error in loading {k}, passing")712                    else:713                        if L1 != L2:714                            logger.info(715                                '=> load_pretrained: resized variant: {} to {}'716                                    .format((1, L1, C1), (1, L2, C2))717                            )718                            S1 = int(L1 ** 0.5)719                            S2 = int(L2 ** 0.5)720                            absolute_pos_embed_pretrained = absolute_pos_embed_pretrained.reshape(-1, S1, S1, C1)721                            absolute_pos_embed_pretrained = absolute_pos_embed_pretrained.permute(0, 3, 1, 2)722                            absolute_pos_embed_pretrained_resized = torch.nn.functional.interpolate(723                                absolute_pos_embed_pretrained, size=(S2, S2), mode='bicubic')724                            v = absolute_pos_embed_pretrained_resized.permute(0, 2, 3, 1).flatten(1, 2)725 726                need_init_state_dict[k] = v727        self.load_state_dict(need_init_state_dict, strict=False)728 729 730    def forward(self, x):731        """Forward function."""732        x = self.patch_embed(x)733 734        Wh, Ww = x.size(2), x.size(3)735        if self.ape:736            # interpolate the position embedding to the corresponding size737            absolute_pos_embed = F.interpolate(738                self.absolute_pos_embed, size=(Wh, Ww), mode="bicubic"739            )740            x = (x + absolute_pos_embed).flatten(2).transpose(1, 2)  # B Wh*Ww C741        else:742            x = x.flatten(2).transpose(1, 2)743        x = self.pos_drop(x)744 745        outs = {}746        for i in range(self.num_layers):747            layer = self.layers[i]748            x_out, H, W, x, Wh, Ww = layer(x, Wh, Ww)749 750            if i in self.out_indices:751                norm_layer = getattr(self, f"norm{i}")752                x_out = norm_layer(x_out)753 754                out = x_out.view(-1, H, W, self.num_features[i]).permute(0, 3, 1, 2).contiguous()755                outs["res{}".format(i + 2)] = out756 757        if len(self.out_indices) == 0:758            outs["res5"] = x_out.view(-1, H, W, self.num_features[i]).permute(0, 3, 1, 2).contiguous()759        760 761        return outs762 763    def train(self, mode=True):764        """Convert the model into training mode while keep layers freezed."""765        super(SwinTransformer, self).train(mode)766        self._freeze_stages()767 768 769class D2SwinTransformer(SwinTransformer, Backbone):770    def __init__(self, cfg, pretrain_img_size, patch_size, in_chans, embed_dim, 771                 depths, num_heads, window_size, mlp_ratio, qkv_bias, qk_scale,772                 drop_rate, attn_drop_rate, drop_path_rate, norm_layer, ape, 773                 patch_norm, out_indices, use_checkpoint):774        super().__init__(775            pretrain_img_size,776            patch_size,777            in_chans,778            embed_dim,779            depths,780            num_heads,781            window_size,782            mlp_ratio,783            qkv_bias,784            qk_scale,785            drop_rate,786            attn_drop_rate,787            drop_path_rate,788            norm_layer,789            ape,790            patch_norm,791            out_indices,792            use_checkpoint=use_checkpoint,793        )794 795        self._out_features = cfg['OUT_FEATURES']796 797        self._out_feature_strides = {798            "res2": 4,799            "res3": 8,800            "res4": 16,801            "res5": 32,802        }803        self._out_feature_channels = {804            "res2": self.num_features[0],805            "res3": self.num_features[1],806            "res4": self.num_features[2],807            "res5": self.num_features[3],808        }809 810    def forward(self, x):811        """812        Args:813            x: Tensor of shape (N,C,H,W). H, W must be a multiple of ``self.size_divisibility``.814        Returns:815            dict[str->Tensor]: names and the corresponding features816        """817        assert (818            x.dim() == 4819        ), f"SwinTransformer takes an input of shape (N, C, H, W). Got {x.shape} instead!"820        outputs = {}821        y = super().forward(x)822        for k in y.keys():823            if k in self._out_features:824                outputs[k] = y[k]825        return outputs826 827    def output_shape(self):828        feature_names = list(set(self._out_feature_strides.keys()) & set(self._out_features))829        return {830            name: ShapeSpec(831                channels=self._out_feature_channels[name], stride=self._out_feature_strides[name]832            )833            for name in feature_names834        }835 836    @property837    def size_divisibility(self):838        return 32839 840 841@register_backbone842def get_swin_backbone(cfg):843    swin_cfg = cfg['MODEL']['BACKBONE']['SWIN']844 845    pretrain_img_size = swin_cfg['PRETRAIN_IMG_SIZE']846    patch_size = swin_cfg['PATCH_SIZE']847    in_chans = 3848    embed_dim = swin_cfg['EMBED_DIM']849    depths = swin_cfg['DEPTHS']850    num_heads = swin_cfg['NUM_HEADS']851    window_size = swin_cfg['WINDOW_SIZE']852    mlp_ratio = swin_cfg['MLP_RATIO']853    qkv_bias = swin_cfg['QKV_BIAS']854    qk_scale = swin_cfg['QK_SCALE']855    drop_rate = swin_cfg['DROP_RATE']856    attn_drop_rate = swin_cfg['ATTN_DROP_RATE']857    drop_path_rate = swin_cfg['DROP_PATH_RATE']858    norm_layer = nn.LayerNorm859    ape = swin_cfg['APE']860    patch_norm = swin_cfg['PATCH_NORM']861    use_checkpoint = swin_cfg['USE_CHECKPOINT']862    out_indices = swin_cfg.get('OUT_INDICES', [0,1,2,3])863    864    swin = D2SwinTransformer(865        swin_cfg,866        pretrain_img_size,867        patch_size,868        in_chans,869        embed_dim,870        depths,871        num_heads,872        window_size,873        mlp_ratio,874        qkv_bias,875        qk_scale,876        drop_rate,877        attn_drop_rate,878        drop_path_rate,879        norm_layer,880        ape,881        patch_norm,882        out_indices,883        use_checkpoint=use_checkpoint,884    )    885 886    if cfg['MODEL']['BACKBONE']['LOAD_PRETRAINED'] is True:887        filename = cfg['MODEL']['BACKBONE']['PRETRAINED']888        with PathManager.open(filename, "rb") as f:889            ckpt = torch.load(f, map_location=cfg['device'])['model']890        swin.load_weights(ckpt, swin_cfg.get('PRETRAINED_LAYERS', ['*']), cfg['VERBOSE'])891 892    return swin