{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!conda install --offline ../input/pyvips-dependency/*.tar.bz2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-08T06:45:58.069543Z","iopub.execute_input":"2022-09-08T06:45:58.070657Z","iopub.status.idle":"2022-09-08T06:46:08.713568Z","shell.execute_reply.started":"2022-09-08T06:45:58.070606Z","shell.execute_reply":"2022-09-08T06:46:08.712313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False","metadata":{"execution":{"iopub.status.busy":"2022-09-08T06:46:08.720364Z","iopub.execute_input":"2022-09-08T06:46:08.723042Z","iopub.status.idle":"2022-09-08T06:46:08.730567Z","shell.execute_reply.started":"2022-09-08T06:46:08.722966Z","shell.execute_reply":"2022-09-08T06:46:08.729463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile coat.py\n\n\"\"\" \nCoaT architecture.\n\nModified from timm/models/vision_transformer.py\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom timm.data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD\nfrom timm.models.layers import DropPath, to_2tuple, trunc_normal_\nfrom timm.models.registry import register_model\n\nfrom einops import rearrange\nfrom functools import partial\nfrom torch import nn, einsum\n\n__all__ = [\n    \"coat_tiny\",\n    \"coat_mini\",\n    \"coat_small\",\n    \"coat_lite_tiny\",\n    \"coat_lite_mini\",\n    \"coat_lite_small\"\n]\n\n\ndef _cfg_coat(url='', **kwargs):\n    return {\n        'url': url,\n        'num_classes': 1000, 'input_size': (3, 224, 224), 'pool_size': None,\n        'crop_pct': .9, 'interpolation': 'bicubic',\n        'mean': IMAGENET_DEFAULT_MEAN, 'std': IMAGENET_DEFAULT_STD,\n        'first_conv': 'patch_embed.proj', 'classifier': 'head',\n        **kwargs\n    }\n\n\nclass Mlp(nn.Module):\n    \"\"\" Feed-forward network (FFN, a.k.a. MLP) class. \"\"\"\n    def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):\n        super().__init__()\n        out_features = out_features or in_features\n        hidden_features = hidden_features or in_features\n        self.fc1 = nn.Linear(in_features, hidden_features)\n        self.act = act_layer()\n        self.fc2 = nn.Linear(hidden_features, out_features)\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.act(x)\n        x = self.drop(x)\n        x = self.fc2(x)\n        x = self.drop(x)\n        return x\n\n\nclass ConvRelPosEnc(nn.Module):\n    \"\"\" Convolutional relative position encoding. \"\"\"\n    def __init__(self, Ch, h, window):\n        \"\"\"\n        Initialization.\n            Ch: Channels per head.\n            h: Number of heads.\n            window: Window size(s) in convolutional relative positional encoding. It can have two forms:\n                    1. An integer of window size, which assigns all attention heads with the same window size in ConvRelPosEnc.\n                    2. A dict mapping window size to #attention head splits (e.g. {window size 1: #attention head split 1, window size 2: #attention head split 2})\n                       It will apply different window size to the attention head splits.\n        \"\"\"\n        super().__init__()\n\n        if isinstance(window, int):\n            window = {window: h}                                                         # Set the same window size for all attention heads.\n            self.window = window\n        elif isinstance(window, dict):\n            self.window = window\n        else:\n            raise ValueError()            \n        \n        self.conv_list = nn.ModuleList()\n        self.head_splits = []\n        for cur_window, cur_head_split in window.items():\n            dilation = 1                                                                 # Use dilation=1 at default.\n            padding_size = (cur_window + (cur_window - 1) * (dilation - 1)) // 2         # Determine padding size. Ref: https://discuss.pytorch.org/t/how-to-keep-the-shape-of-input-and-output-same-when-dilation-conv/14338\n            cur_conv = nn.Conv2d(cur_head_split*Ch, cur_head_split*Ch,\n                kernel_size=(cur_window, cur_window), \n                padding=(padding_size, padding_size),\n                dilation=(dilation, dilation),                          \n                groups=cur_head_split*Ch,\n            )\n            self.conv_list.append(cur_conv)\n            self.head_splits.append(cur_head_split)\n        self.channel_splits = [x*Ch for x in self.head_splits]\n\n    def forward(self, q, v, size):\n        B, h, N, Ch = q.shape\n        H, W = size\n        assert N == 1 + H * W\n\n        # Convolutional relative position encoding.\n        q_img = q[:,:,1:,:]                                                              # Shape: [B, h, H*W, Ch].\n        v_img = v[:,:,1:,:]                                                              # Shape: [B, h, H*W, Ch].\n        \n        v_img = rearrange(v_img, 'B h (H W) Ch -> B (h Ch) H W', H=H, W=W)               # Shape: [B, h, H*W, Ch] -> [B, h*Ch, H, W].\n        v_img_list = torch.split(v_img, self.channel_splits, dim=1)                      # Split according to channels.\n        conv_v_img_list = [conv(x) for conv, x in zip(self.conv_list, v_img_list)]\n        conv_v_img = torch.cat(conv_v_img_list, dim=1)\n        conv_v_img = rearrange(conv_v_img, 'B (h Ch) H W -> B h (H W) Ch', h=h)          # Shape: [B, h*Ch, H, W] -> [B, h, H*W, Ch].\n\n        EV_hat_img = q_img * conv_v_img\n        zero = torch.zeros((B, h, 1, Ch), dtype=q.dtype, layout=q.layout, device=q.device)\n        EV_hat = torch.cat((zero, EV_hat_img), dim=2)                                # Shape: [B, h, N, Ch].\n\n        return EV_hat\n\n\nclass FactorAtt_ConvRelPosEnc(nn.Module):\n    \"\"\" Factorized attention with convolutional relative position encoding class. \"\"\"\n    def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., shared_crpe=None):\n        super().__init__()\n        self.num_heads = num_heads\n        head_dim = dim // num_heads\n        self.scale = qk_scale or head_dim ** -0.5\n\n        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(attn_drop)                                           # Note: attn_drop is actually not used.\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n        # Shared convolutional relative position encoding.\n        self.crpe = shared_crpe\n\n    def forward(self, x, size):\n        B, N, C = x.shape\n\n        # Generate Q, K, V.\n        qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)  # Shape: [3, B, h, N, Ch].\n        q, k, v = qkv[0], qkv[1], qkv[2]                                                 # Shape: [B, h, N, Ch].\n\n        # Factorized attention.\n        k_softmax = k.softmax(dim=2)                                                     # Softmax on dim N.\n        k_softmax_T_dot_v = einsum('b h n k, b h n v -> b h k v', k_softmax, v)          # Shape: [B, h, Ch, Ch].\n        factor_att        = einsum('b h n k, b h k v -> b h n v', q, k_softmax_T_dot_v)  # Shape: [B, h, N, Ch].\n\n        # Convolutional relative position encoding.\n        crpe = self.crpe(q, v, size=size)                                                # Shape: [B, h, N, Ch].\n\n        # Merge and reshape.\n        x = self.scale * factor_att + crpe\n        x = x.transpose(1, 2).reshape(B, N, C)                                           # Shape: [B, h, N, Ch] -> [B, N, h, Ch] -> [B, N, C].\n\n        # Output projection.\n        x = self.proj(x)\n        x = self.proj_drop(x)\n\n        return x                                                                         # Shape: [B, N, C].\n\n\nclass ConvPosEnc(nn.Module):\n    \"\"\" Convolutional Position Encoding. \n        Note: This module is similar to the conditional position encoding in CPVT.\n    \"\"\"\n    def __init__(self, dim, k=3):\n        super(ConvPosEnc, self).__init__()\n        self.proj = nn.Conv2d(dim, dim, k, 1, k//2, groups=dim) \n    \n    def forward(self, x, size):\n        B, N, C = x.shape\n        H, W = size\n        assert N == 1 + H * W\n\n        # Extract CLS token and image tokens.\n        cls_token, img_tokens = x[:, :1], x[:, 1:]                                       # Shape: [B, 1, C], [B, H*W, C].\n        \n        # Depthwise convolution.\n        feat = img_tokens.transpose(1, 2).view(B, C, H, W)\n        x = self.proj(feat) + feat\n        x = x.flatten(2).transpose(1, 2)\n\n        # Combine with CLS token.\n        x = torch.cat((cls_token, x), dim=1)\n\n        return x\n\n\nclass SerialBlock(nn.Module):\n    \"\"\" Serial block class.\n        Note: In this implementation, each serial block only contains a conv-attention and a FFN (MLP) module. \"\"\"\n    def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,\n                 drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm,\n                 shared_cpe=None, shared_crpe=None):\n        super().__init__()\n\n        # Conv-Attention.\n        self.cpe = shared_cpe\n\n        self.norm1 = norm_layer(dim)\n        self.factoratt_crpe = FactorAtt_ConvRelPosEnc(\n            dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop, \n            shared_crpe=shared_crpe)\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n        # MLP.\n        self.norm2 = norm_layer(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n\n    def forward(self, x, size):\n        # Conv-Attention.\n        x = self.cpe(x, size)                  # Apply convolutional position encoding.\n        cur = self.norm1(x)\n        cur = self.factoratt_crpe(cur, size)   # Apply factorized attention and convolutional relative position encoding.\n        x = x + self.drop_path(cur) \n\n        # MLP. \n        cur = self.norm2(x)\n        cur = self.mlp(cur)\n        x = x + self.drop_path(cur)\n\n        return x\n\n\nclass ParallelBlock(nn.Module):\n    \"\"\" Parallel block class. \"\"\"\n    def __init__(self, dims, num_heads, mlp_ratios=[], qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,\n                 drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm,\n                 shared_cpes=None, shared_crpes=None):\n        super().__init__()\n\n        # Conv-Attention.\n        self.cpes = shared_cpes\n\n        self.norm12 = norm_layer(dims[1])\n        self.norm13 = norm_layer(dims[2])\n        self.norm14 = norm_layer(dims[3])\n        self.factoratt_crpe2 = FactorAtt_ConvRelPosEnc(\n            dims[1], num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop, \n            shared_crpe=shared_crpes[1]\n        )\n        self.factoratt_crpe3 = FactorAtt_ConvRelPosEnc(\n            dims[2], num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop, \n            shared_crpe=shared_crpes[2]\n        )\n        self.factoratt_crpe4 = FactorAtt_ConvRelPosEnc(\n            dims[3], num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop, \n            shared_crpe=shared_crpes[3]\n        )\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n        # MLP.\n        self.norm22 = norm_layer(dims[1])\n        self.norm23 = norm_layer(dims[2])\n        self.norm24 = norm_layer(dims[3])\n        assert dims[1] == dims[2] == dims[3]                              # In parallel block, we assume dimensions are the same and share the linear transformation.\n        assert mlp_ratios[1] == mlp_ratios[2] == mlp_ratios[3]\n        mlp_hidden_dim = int(dims[1] * mlp_ratios[1])\n        self.mlp2 = self.mlp3 = self.mlp4 = Mlp(in_features=dims[1], hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n\n    def upsample(self, x, output_size, size):\n        \"\"\" Feature map up-sampling. \"\"\"\n        return self.interpolate(x, output_size=output_size, size=size)\n\n    def downsample(self, x, output_size, size):\n        \"\"\" Feature map down-sampling. \"\"\"\n        return self.interpolate(x, output_size=output_size, size=size)\n\n    def interpolate(self, x, output_size, size):\n        \"\"\" Feature map interpolation. \"\"\"\n        B, N, C = x.shape\n        H, W = size\n        assert N == 1 + H * W\n\n        cls_token  = x[:, :1, :]\n        img_tokens = x[:, 1:, :]\n        \n        img_tokens = img_tokens.transpose(1, 2).reshape(B, C, H, W)\n        img_tokens = F.interpolate(img_tokens, size=output_size, mode='bilinear')  # FIXME: May have alignment issue.\n        img_tokens = img_tokens.reshape(B, C, -1).transpose(1, 2)\n        \n        out = torch.cat((cls_token, img_tokens), dim=1)\n\n        return out\n\n    def forward(self, x1, x2, x3, x4, sizes):\n        _, (H2, W2), (H3, W3), (H4, W4) = sizes\n        \n        # Conv-Attention.\n        x2 = self.cpes[1](x2, size=(H2, W2))  # Note: x1 is ignored.\n        x3 = self.cpes[2](x3, size=(H3, W3))\n        x4 = self.cpes[3](x4, size=(H4, W4))\n        \n        cur2 = self.norm12(x2)\n        cur3 = self.norm13(x3)\n        cur4 = self.norm14(x4)\n        cur2 = self.factoratt_crpe2(cur2, size=(H2,W2))\n        cur3 = self.factoratt_crpe3(cur3, size=(H3,W3))\n        cur4 = self.factoratt_crpe4(cur4, size=(H4,W4))\n        upsample3_2 = self.upsample(cur3, output_size=(H2,W2), size=(H3,W3))\n        upsample4_3 = self.upsample(cur4, output_size=(H3,W3), size=(H4,W4))\n        upsample4_2 = self.upsample(cur4, output_size=(H2,W2), size=(H4,W4))\n        downsample2_3 = self.downsample(cur2, output_size=(H3,W3), size=(H2,W2))\n        downsample3_4 = self.downsample(cur3, output_size=(H4,W4), size=(H3,W3))\n        downsample2_4 = self.downsample(cur2, output_size=(H4,W4), size=(H2,W2))\n        cur2 = cur2  + upsample3_2   + upsample4_2\n        cur3 = cur3  + upsample4_3   + downsample2_3\n        cur4 = cur4  + downsample3_4 + downsample2_4\n        x2 = x2 + self.drop_path(cur2) \n        x3 = x3 + self.drop_path(cur3) \n        x4 = x4 + self.drop_path(cur4) \n\n        # MLP. \n        cur2 = self.norm22(x2)\n        cur3 = self.norm23(x3)\n        cur4 = self.norm24(x4)\n        cur2 = self.mlp2(cur2)\n        cur3 = self.mlp3(cur3)\n        cur4 = self.mlp4(cur4)\n        x2 = x2 + self.drop_path(cur2)\n        x3 = x3 + self.drop_path(cur3)\n        x4 = x4 + self.drop_path(cur4) \n\n        return x1, x2, x3, x4\n\n\nclass PatchEmbed(nn.Module):\n    \"\"\" Image to Patch Embedding \"\"\"\n    def __init__(self, patch_size=16, in_chans=3, embed_dim=768):\n        super().__init__()\n        patch_size = to_2tuple(patch_size)\n\n        self.patch_size = patch_size\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n        self.norm = nn.LayerNorm(embed_dim)\n\n    def forward(self, x):\n        _, _, H, W = x.shape\n        out_H, out_W = H // self.patch_size[0], W // self.patch_size[1]\n\n        x = self.proj(x).flatten(2).transpose(1, 2)\n        out = self.norm(x)\n        \n        return out, (out_H, out_W)\n\n\nclass CoaT(nn.Module):\n    \"\"\" CoaT class. \"\"\"\n    def __init__(self, patch_size=16, in_chans=3, num_classes=1000, embed_dims=[0, 0, 0, 0], \n                 serial_depths=[0, 0, 0, 0], parallel_depth=0,\n                 num_heads=0, mlp_ratios=[0, 0, 0, 0], qkv_bias=True, qk_scale=None, drop_rate=0., attn_drop_rate=0.,\n                 drop_path_rate=0., norm_layer=partial(nn.LayerNorm, eps=1e-6),\n                 return_interm_layers=False, out_features=None, crpe_window={3:2, 5:3, 7:3},\n                 **kwargs):\n        super().__init__()\n        self.return_interm_layers = return_interm_layers\n        self.out_features = out_features\n        self.num_classes = num_classes\n\n        # Patch embeddings.\n        self.patch_embed1 = PatchEmbed(patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dims[0])\n        self.patch_embed2 = PatchEmbed(patch_size=2, in_chans=embed_dims[0], embed_dim=embed_dims[1])\n        self.patch_embed3 = PatchEmbed(patch_size=2, in_chans=embed_dims[1], embed_dim=embed_dims[2])\n        self.patch_embed4 = PatchEmbed(patch_size=2, in_chans=embed_dims[2], embed_dim=embed_dims[3])\n\n        # Class tokens.\n        self.cls_token1 = nn.Parameter(torch.zeros(1, 1, embed_dims[0]))\n        self.cls_token2 = nn.Parameter(torch.zeros(1, 1, embed_dims[1]))\n        self.cls_token3 = nn.Parameter(torch.zeros(1, 1, embed_dims[2]))\n        self.cls_token4 = nn.Parameter(torch.zeros(1, 1, embed_dims[3]))\n\n        # Convolutional position encodings.\n        self.cpe1 = ConvPosEnc(dim=embed_dims[0], k=3)\n        self.cpe2 = ConvPosEnc(dim=embed_dims[1], k=3)\n        self.cpe3 = ConvPosEnc(dim=embed_dims[2], k=3)\n        self.cpe4 = ConvPosEnc(dim=embed_dims[3], k=3)\n\n        # Convolutional relative position encodings.\n        self.crpe1 = ConvRelPosEnc(Ch=embed_dims[0] // num_heads, h=num_heads, window=crpe_window)\n        self.crpe2 = ConvRelPosEnc(Ch=embed_dims[1] // num_heads, h=num_heads, window=crpe_window)\n        self.crpe3 = ConvRelPosEnc(Ch=embed_dims[2] // num_heads, h=num_heads, window=crpe_window)\n        self.crpe4 = ConvRelPosEnc(Ch=embed_dims[3] // num_heads, h=num_heads, window=crpe_window)\n\n        # Enable stochastic depth.\n        dpr = drop_path_rate\n        \n        # Serial blocks 1.\n        self.serial_blocks1 = nn.ModuleList([\n            SerialBlock(\n                dim=embed_dims[0], num_heads=num_heads, mlp_ratio=mlp_ratios[0], qkv_bias=qkv_bias, qk_scale=qk_scale,\n                drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer, \n                shared_cpe=self.cpe1, shared_crpe=self.crpe1\n            )\n            for _ in range(serial_depths[0])]\n        )\n\n        # Serial blocks 2.\n        self.serial_blocks2 = nn.ModuleList([\n            SerialBlock(\n                dim=embed_dims[1], num_heads=num_heads, mlp_ratio=mlp_ratios[1], qkv_bias=qkv_bias, qk_scale=qk_scale,\n                drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer, \n                shared_cpe=self.cpe2, shared_crpe=self.crpe2\n            )\n            for _ in range(serial_depths[1])]\n        )\n\n        # Serial blocks 3.\n        self.serial_blocks3 = nn.ModuleList([\n            SerialBlock(\n                dim=embed_dims[2], num_heads=num_heads, mlp_ratio=mlp_ratios[2], qkv_bias=qkv_bias, qk_scale=qk_scale,\n                drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer, \n                shared_cpe=self.cpe3, shared_crpe=self.crpe3\n            )\n            for _ in range(serial_depths[2])]\n        )\n\n        # Serial blocks 4.\n        self.serial_blocks4 = nn.ModuleList([\n            SerialBlock(\n                dim=embed_dims[3], num_heads=num_heads, mlp_ratio=mlp_ratios[3], qkv_bias=qkv_bias, qk_scale=qk_scale,\n                drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer, \n                shared_cpe=self.cpe4, shared_crpe=self.crpe4\n            )\n            for _ in range(serial_depths[3])]\n        )\n\n        # Parallel blocks.\n        self.parallel_depth = parallel_depth\n        if self.parallel_depth > 0:\n            self.parallel_blocks = nn.ModuleList([\n                ParallelBlock(\n                    dims=embed_dims, num_heads=num_heads, mlp_ratios=mlp_ratios, qkv_bias=qkv_bias, qk_scale=qk_scale,\n                    drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer, \n                    shared_cpes=[self.cpe1, self.cpe2, self.cpe3, self.cpe4],\n                    shared_crpes=[self.crpe1, self.crpe2, self.crpe3, self.crpe4]\n                )\n                for _ in range(parallel_depth)]\n            )\n\n        # Classification head(s).\n        if not self.return_interm_layers:\n            self.norm1 = norm_layer(embed_dims[0])\n            self.norm2 = norm_layer(embed_dims[1])\n            self.norm3 = norm_layer(embed_dims[2])\n            self.norm4 = norm_layer(embed_dims[3])\n\n            if self.parallel_depth > 0:                                  # CoaT series: Aggregate features of last three scales for classification.\n                assert embed_dims[1] == embed_dims[2] == embed_dims[3]\n                self.aggregate = torch.nn.Conv1d(in_channels=3, out_channels=1, kernel_size=1)\n                self.head = nn.Linear(embed_dims[3], num_classes)\n            else:\n                self.head = nn.Linear(embed_dims[3], num_classes)        # CoaT-Lite series: Use feature of last scale for classification.\n\n        # Initialize weights.\n        trunc_normal_(self.cls_token1, std=.02)\n        trunc_normal_(self.cls_token2, std=.02)\n        trunc_normal_(self.cls_token3, std=.02)\n        trunc_normal_(self.cls_token4, std=.02)\n        self.apply(self._init_weights)\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Linear):\n            trunc_normal_(m.weight, std=.02)\n            if isinstance(m, nn.Linear) and m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.bias, 0)\n            nn.init.constant_(m.weight, 1.0)\n\n    @torch.jit.ignore\n    def no_weight_decay(self):\n        return {'cls_token1', 'cls_token2', 'cls_token3', 'cls_token4'}\n\n    def get_classifier(self):\n        return self.head\n\n    def reset_classifier(self, num_classes, global_pool=''):\n        self.num_classes = num_classes\n        self.head = nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()\n\n    def insert_cls(self, x, cls_token):\n        \"\"\" Insert CLS token. \"\"\"\n        cls_tokens = cls_token.expand(x.shape[0], -1, -1)\n        x = torch.cat((cls_tokens, x), dim=1)\n        return x\n\n    def remove_cls(self, x):\n        \"\"\" Remove CLS token. \"\"\"\n        return x[:, 1:, :]\n\n    def forward_features(self, x0):\n        #print('hi')\n        B = x0.shape[0]\n\n        # Serial blocks 1.\n        x1, (H1, W1) = self.patch_embed1(x0)\n        x1 = self.insert_cls(x1, self.cls_token1)\n        for blk in self.serial_blocks1:\n            x1 = blk(x1, size=(H1, W1))\n        x1_nocls = self.remove_cls(x1)\n        x1_nocls = x1_nocls.reshape(B, H1, W1, -1).permute(0, 3, 1, 2).contiguous()\n        \n        # Serial blocks 2.\n        x2, (H2, W2) = self.patch_embed2(x1_nocls)\n        x2 = self.insert_cls(x2, self.cls_token2)\n        for blk in self.serial_blocks2:\n            x2 = blk(x2, size=(H2, W2))\n        x2_nocls = self.remove_cls(x2)\n        x2_nocls = x2_nocls.reshape(B, H2, W2, -1).permute(0, 3, 1, 2).contiguous()\n\n        # Serial blocks 3.\n        x3, (H3, W3) = self.patch_embed3(x2_nocls)\n        x3 = self.insert_cls(x3, self.cls_token3)\n        for blk in self.serial_blocks3:\n            x3 = blk(x3, size=(H3, W3))\n        x3_nocls = self.remove_cls(x3)\n        x3_nocls = x3_nocls.reshape(B, H3, W3, -1).permute(0, 3, 1, 2).contiguous()\n\n        # Serial blocks 4.\n        x4, (H4, W4) = self.patch_embed4(x3_nocls)\n        x4 = self.insert_cls(x4, self.cls_token4)\n        for blk in self.serial_blocks4:\n            x4 = blk(x4, size=(H4, W4))\n        x4_nocls = self.remove_cls(x4)\n        x4_nocls = x4_nocls.reshape(B, H4, W4, -1).permute(0, 3, 1, 2).contiguous()\n\n        # Only serial blocks: Early return.\n        if self.parallel_depth == 0:\n            if self.return_interm_layers:   # Return intermediate features for down-stream tasks (e.g. Deformable DETR and Detectron2).\n                feat_out = {}   \n                if 'x1_nocls' in self.out_features:\n                    feat_out['x1_nocls'] = x1_nocls\n                if 'x2_nocls' in self.out_features:\n                    feat_out['x2_nocls'] = x2_nocls\n                if 'x3_nocls' in self.out_features:\n                    feat_out['x3_nocls'] = x3_nocls\n                if 'x4_nocls' in self.out_features:\n                    feat_out['x4_nocls'] = x4_nocls\n                return feat_out\n            else:                           # Return features for classification.\n                x4 = self.norm4(x4)\n                #x4_cls = x4[:, 0]\n                #print(x4.shape, x4_cls.shape)\n                return x4\n\n        # Parallel blocks.\n        for blk in self.parallel_blocks:\n            x1, x2, x3, x4 = blk(x1, x2, x3, x4, sizes=[(H1, W1), (H2, W2), (H3, W3), (H4, W4)])\n\n        if self.return_interm_layers:       # Return intermediate features for down-stream tasks (e.g. Deformable DETR and Detectron2).\n            feat_out = {}   \n            if 'x1_nocls' in self.out_features:\n                x1_nocls = self.remove_cls(x1)\n                x1_nocls = x1_nocls.reshape(B, H1, W1, -1).permute(0, 3, 1, 2).contiguous()\n                feat_out['x1_nocls'] = x1_nocls\n            if 'x2_nocls' in self.out_features:\n                x2_nocls = self.remove_cls(x2)\n                x2_nocls = x2_nocls.reshape(B, H2, W2, -1).permute(0, 3, 1, 2).contiguous()\n                feat_out['x2_nocls'] = x2_nocls\n            if 'x3_nocls' in self.out_features:\n                x3_nocls = self.remove_cls(x3)\n                x3_nocls = x3_nocls.reshape(B, H3, W3, -1).permute(0, 3, 1, 2).contiguous()\n                feat_out['x3_nocls'] = x3_nocls\n            if 'x4_nocls' in self.out_features:\n                x4_nocls = self.remove_cls(x4)\n                x4_nocls = x4_nocls.reshape(B, H4, W4, -1).permute(0, 3, 1, 2).contiguous()\n                feat_out['x4_nocls'] = x4_nocls\n            return feat_out\n        else:\n            x2 = self.norm2(x2)\n            x3 = self.norm3(x3)\n            x4 = self.norm4(x4)\n            x2_cls = x2[:, :1]              # Shape: [B, 1, C].\n            x3_cls = x3[:, :1]\n            x4_cls = x4[:, :1]\n            merged_cls = torch.cat((x2_cls, x3_cls, x4_cls), dim=1)       # Shape: [B, 3, C].\n            #print(merged_cls.shape)\n            merged_cls = self.aggregate(merged_cls).squeeze(dim=1)        # Shape: [B, C].\n            return merged_cls\n\n    def forward(self, x):\n        if self.return_interm_layers:       # Return intermediate features (for down-stream tasks).\n            return self.forward_features(x)\n        else:                               # Return features for classification.\n            x = self.forward_features(x) \n            x = self.head(x)\n            return x\n\n\n# CoaT.\n@register_model\ndef coat_tiny(**kwargs):\n    model = CoaT(patch_size=4, embed_dims=[152, 152, 152, 152], serial_depths=[2, 2, 2, 2], parallel_depth=6, num_heads=8, mlp_ratios=[4, 4, 4, 4], **kwargs)\n    model.default_cfg = _cfg_coat()\n    return model\n\n@register_model\ndef coat_mini(**kwargs):\n    model = CoaT(patch_size=4, embed_dims=[152, 216, 216, 216], serial_depths=[2, 2, 2, 2], parallel_depth=6, num_heads=8, mlp_ratios=[4, 4, 4, 4], **kwargs)\n    model.default_cfg = _cfg_coat()\n    return model\n\n@register_model\ndef coat_small(**kwargs):\n    model = CoaT(patch_size=4, embed_dims=[152, 320, 320, 320], serial_depths=[2, 2, 2, 2], parallel_depth=6, num_heads=8, mlp_ratios=[4, 4, 4, 4], **kwargs)\n    model.default_cfg = _cfg_coat()\n    return model\n\n# CoaT-Lite.\n@register_model\ndef coat_lite_tiny(**kwargs):\n    model = CoaT(patch_size=4, embed_dims=[64, 128, 256, 320], serial_depths=[2, 2, 2, 2], parallel_depth=0, num_heads=8, mlp_ratios=[8, 8, 4, 4], **kwargs)\n    model.default_cfg = _cfg_coat()\n    return model\n\n@register_model\ndef coat_lite_mini(**kwargs):\n    model = CoaT(patch_size=4, embed_dims=[64, 128, 320, 512], serial_depths=[2, 2, 2, 2], parallel_depth=0, num_heads=8, mlp_ratios=[8, 8, 4, 4], **kwargs)\n    model.default_cfg = _cfg_coat()\n    return model\n\n@register_model\ndef coat_lite_small(**kwargs):\n    model = CoaT(patch_size=4, embed_dims=[64, 128, 320, 512], serial_depths=[3, 4, 6, 3], parallel_depth=0, num_heads=8, mlp_ratios=[8, 8, 4, 4], **kwargs)\n    model.default_cfg = _cfg_coat()\n    return model\n\n@register_model\ndef coat_lite_medium(**kwargs):\n    model = CoaT(patch_size=4, embed_dims=[128, 256, 320, 512], serial_depths=[3, 6, 10, 8], parallel_depth=0, num_heads=8, mlp_ratios=[4, 4, 4, 4], **kwargs)\n    model.default_cfg = _cfg_coat()\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-09-08T06:46:08.738923Z","iopub.execute_input":"2022-09-08T06:46:08.742148Z","iopub.status.idle":"2022-09-08T06:46:08.818673Z","shell.execute_reply.started":"2022-09-08T06:46:08.742109Z","shell.execute_reply":"2022-09-08T06:46:08.817829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pyvips\nimport os\nimport gc\nimport zipfile\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport cv2\n\nif DEBUG:\n    test = pd.read_csv('../input/mayo-clinic-strip-ai/train.csv').iloc[:280]\n    test_folder = '../input/mayo-clinic-strip-ai/train'\nelse:\n    test = pd.read_csv('../input/mayo-clinic-strip-ai/test.csv')\n    test_folder = '../input/mayo-clinic-strip-ai/test'\n    \nsubmission = pd.read_csv('../input/mayo-clinic-strip-ai/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-09-08T06:50:48.578848Z","iopub.execute_input":"2022-09-08T06:50:48.579897Z","iopub.status.idle":"2022-09-08T06:50:48.598197Z","shell.execute_reply.started":"2022-09-08T06:50:48.57984Z","shell.execute_reply":"2022-09-08T06:50:48.597098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tile(img, sz=128, N=16):\n    shape = img.shape\n    pad0,pad1 = (sz - shape[0]%sz)%sz, (sz - shape[1]%sz)%sz\n    img = np.pad(img,[[pad0//2,pad0-pad0//2],[pad1//2,pad1-pad1//2],[0,0]],constant_values=255)\n    img = img.reshape(img.shape[0]//sz,sz,img.shape[1]//sz,sz,3)\n    img = img.transpose(0,2,1,3,4).reshape(-1,sz,sz,3)\n    if len(img) < N:\n        img = np.pad(img,[[0,N-len(img)],[0,0],[0,0],[0,0]],constant_values=255)\n    idxs = np.argsort(img.reshape(img.shape[0],-1).sum(-1))[:N] # pick up Top N dark tiles\n    \n    #idxs = np.argsort(img.std(axis=(1,2)).max(1))[-N:][::-1]\n    \n    img = img[idxs]\n    return img\n\ndef save_dataset(\n    df: pd.DataFrame, \n    N=16,\n    max_size=20000, \n    crop_size=1024, \n    image_dir='../input/mayo-clinic-strip-ai/train', \n    out_dir='./test',\n):\n    format_to_dtype = {\n       'uchar': np.uint8,\n       'char': np.int8,\n       'ushort': np.uint16,\n       'short': np.int16,\n       'uint': np.uint32,\n       'int': np.int32,\n       'float': np.float32,\n       'double': np.float64,\n       'complex': np.complex64,\n       'dpcomplex': np.complex128,\n    }\n    def vips2numpy(vi):\n        return np.ndarray(\n            buffer=vi.write_to_memory(),\n            dtype=format_to_dtype[vi.format],\n            shape=[vi.height, vi.width, vi.bands])\n    \n    if not os.path.isdir(out_dir):\n        os.makedirs(out_dir)\n        \n    tk0 = tqdm(enumerate(df[\"image_id\"].values), total=len(df))\n    for i, image_id in tk0:\n        print(f\"[{i+1}/{len(df)}] image_id: {image_id}\")\n        image = pyvips.Image.thumbnail(f'{image_dir}/{image_id}.tif', max_size)\n        image = vips2numpy(image)\n        width, height, c = image.shape\n        print(f\"Input width: {width} height: {height}\")\n        images = tile(image, sz=crop_size, N=N)\n        for idx, img in enumerate(images):\n            img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)\n            cv2.imwrite(f\"{out_dir}/{image_id}_{idx}.jpg\", img, [cv2.IMWRITE_JPEG_QUALITY, 100])\n            \n        del img, image, images; gc.collect()\n\n\nsave_dataset(\n    test,\n    N=16, \n    max_size=20000,\n    crop_size=1024, \n    image_dir=test_folder, \n    out_dir=f'./test'\n)","metadata":{"execution":{"iopub.status.busy":"2022-09-08T06:46:08.848441Z","iopub.execute_input":"2022-09-08T06:46:08.849069Z","iopub.status.idle":"2022-09-08T06:47:17.020898Z","shell.execute_reply.started":"2022-09-08T06:46:08.849026Z","shell.execute_reply":"2022-09-08T06:47:17.019808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\n\n_fd = 0\nimport os\nimport argparse\n\nos.environ['CUDA_VISIBLE_DEVICES'] = f\"{_fd}\"\nos.environ['TOKENIZERS_PARALLELISM'] = \"true\"\n\nimport sys\nsys.path.insert(0, '../input/timm-0-6-9/pytorch-image-models-master')\n\nclass CFG:\n    wandb=False\n    debug=False\n    apex=True\n    print_freq=50\n    monitor_freq=10000\n    num_workers=4\n    model='swin_large_patch4_window12_384'\n    model_reformat=False\n    image_size=384\n    seq_len=144\n    model_name='../input/mayo-1024-16-swin-l-v6' #'dbv3-ccc-train-relabel-05'\n    num_instance = 16\n    n_fold=5\n    batch_size=4\n    target_ohe_columns=['CE', 'LAA']\n    \nclass CFG2:\n    wandb=False\n    debug=False\n    apex=True\n    print_freq=50\n    monitor_freq=10000\n    num_workers=4\n    model='coat_lite_medium'\n    model_reformat=False\n    image_size=384\n    seq_len=144\n    model_name='../input/mayo-1024-16-coat-lite-m-v1' #'dbv3-ccc-train-relabel-05'\n    num_instance = 16\n    n_fold=5\n    batch_size=4\n    target_ohe_columns=['CE', 'LAA']\n    \nif CFG.debug:\n    CFG.epochs = 2\n    CFG.trn_fold = [0]\n\n\n# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\n\nOUTPUT_DIR = f'{CFG.model_name}/'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\n# ====================================================\n# Library\n# ====================================================\nimport os\nimport gc\nimport re\nimport ast\nimport sys\nimport copy\nimport json\nimport time\nimport math\nimport shutil\nimport string\nimport pickle\nimport random\nimport joblib\nimport itertools\nfrom pathlib import Path\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nfrom text_unidecode import unidecode\nfrom typing import Dict, List, Tuple\nimport codecs\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\npd.set_option('display.max_rows', 500)\npd.set_option('display.max_columns', 500)\npd.set_option('display.width', 1000)\nfrom tqdm.auto import tqdm; tqdm.pandas()\nfrom sklearn.metrics import f1_score\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nimport torch\nprint(f\"torch.__version__: {torch.__version__}\")\nimport torch.nn as nn\nfrom torch.nn import Parameter\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.utils.data import DataLoader, Dataset\n\nimport timm\nimport coat\nimport cv2\nfrom torch.utils.checkpoint import checkpoint # to support gradient checkpointing\nfrom torch.optim import Optimizer\nfrom collections import defaultdict\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import *\n\nfrom albumentations import (\n    HorizontalFlip, VerticalFlip, IAAPerspective, ShiftScaleRotate, CLAHE, RandomRotate90,\n    Transpose, ShiftScaleRotate, Blur, OpticalDistortion, GridDistortion, HueSaturationValue,\n    IAAAdditiveGaussianNoise, GaussNoise, MotionBlur, MedianBlur, IAAPiecewiseAffine, RandomResizedCrop,\n    IAASharpen, IAAEmboss, RandomBrightnessContrast, Flip, OneOf, Compose, Normalize, Cutout, CoarseDropout, ShiftScaleRotate, CenterCrop, Resize\n)\n\n# ====================================================\n# Data Loading\n# ====================================================\n\n\ndef parse_images(folder):\n    imgs = [os.path.join(folder, f) for f in os.listdir(folder)]\n\n    df = pd.DataFrame()\n    df['image_path'] = imgs\n    df['image_id'] = df['image_path'].apply(lambda x: '_'.join(x.split('/')[-1].replace('.jpg', '').split('_')[:2]))\n    df['instance_id'] = df['image_path'].apply(lambda x: int(x.split('_')[-1].replace('.jpg', '')))\n\n    df = df.sort_values(['image_id', 'instance_id']).reset_index(drop=True)\n\n    return df\n\ndef merge_image_info(image_df, info_df):\n    return image_df.merge(info_df, on='image_id', how='left').reset_index(drop=True)\n\n# ====================================================\n# Dataset\n# ====================================================\n\nclass TrainDataset(Dataset):\n    def __init__(self, cfg, csv, transform):\n        self.cfg = cfg\n        self.csv = csv.reset_index(drop=True)\n        self.transform = transform\n        self.image_ids = self.csv.image_id.unique()\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, index):\n        img_id = self.image_ids[index]\n        df = self.csv.loc[self.csv.image_id==img_id]\n        \n        images = []\n        for i, path in enumerate(df.image_path.values):\n            if i >= self.cfg.num_instance:\n                break\n\n            image = cv2.imread(path)\n            image = image[:, :, ::-1] #/255.\n            \n            res = self.transform(image=image)\n            image = res['image']\n\n            images += [image.unsqueeze(0)]\n\n        assert len(images) == self.cfg.num_instance\n\n        images = torch.cat(images, 0) # n, c, w, h\n        \n        labels = df.iloc[0][self.cfg.target_ohe_columns].values.astype(int)\n\n        return images, labels\n\n# ====================================================\n# Model\n# ====================================================\ndef mish(input):\n    \"\"\"\n    Applies the mish function element-wise:\n    mish(x) = x * tanh(softplus(x)) = x * tanh(ln(1 + exp(x)))\n    See additional documentation for mish class.\n    \"\"\"\n    return input * torch.tanh(F.softplus(input))\n\nfrom torch.nn.modules.loss import _WeightedLoss\n\n# reference: https://www.kaggle.com/c/siim-isic-melanoma-classification/discussion/173733\nclass SoftCrossEntropyLoss(_WeightedLoss):\n    def __init__(self, weight=None, reduction='mean'):\n        super().__init__(weight=weight, reduction=reduction)\n        self.weight = weight\n        self.reduction = reduction\n\n    def forward(self, inputs, targets):\n        #print(inputs, targets)\n        \n        lsm = F.log_softmax(inputs, -1)\n\n        if self.weight is not None:\n            lsm = lsm * self.weight.unsqueeze(0)\n\n        loss = -(targets * lsm).sum(-1)\n\n        if  self.reduction == 'sum':\n            loss = loss.sum()\n        elif  self.reduction == 'mean':\n            loss = loss.mean()\n        else:\n            pass\n\n        return loss\n\nclass MayoSoftCrossEntropyLoss(_WeightedLoss):\n    def __init__(self, weight=None, reduction='mean'):\n        super().__init__(weight=weight, reduction=reduction)\n        self.weight = weight # per sample class weights (need to multiplied by the same size)\n        self.reduction = reduction\n        #print(weight)\n\n    def forward(self, inputs, targets):\n        #print(inputs, targets)\n        \n        lsm = F.log_softmax(inputs, -1)\n        \n        #weight = torch.FloatTensor([(targets==0).sum(), (targets==1).sum()]).to(inputs.device)\n\n        weight = self.weight * inputs.shape[0]\n        weight = 1./weight\n\n        #print(self.weight, inputs.shape[0], weight)\n\n        if self.weight is not None:\n            lsm = lsm * weight.unsqueeze(0)\n\n        loss = -(targets * lsm).sum()/2.\n\n        return loss\n\nclass MSDPRegHead(nn.Module):\n    def __init__(self, in_feat, out_feat):\n        super(MSDPRegHead, self).__init__()\n        #self.norm = nn.LayerNorm(in_feat)\n        #self.fc1 = nn.Linear(in_feat, 1024)\n        #self.fc2 = nn.Linear(in_feat//2, out_feat)\n        \n        self.dropout1 = nn.Dropout(CFG.fc_dropout)\n        self.dropout2 = nn.Dropout(CFG.fc_dropout)\n        self.out = nn.Linear(512 , out_feat)\n        self.out_hidden = nn.Linear(in_feat , 512)\n        \n        #torch.nn.init.normal_(self.fc1.weight, std=0.02)\n        torch.nn.init.normal_(self.out.weight, std=0.02)\n        torch.nn.init.normal_(self.out_hidden.weight, std=0.02)\n        \n    def forward(self, x):\n        #x = self.norm(x)\n        \n        #x = self.fc1(x)\n        \n        x = torch.mean(\n            torch.stack(\n                [mish(self.out_hidden(self.dropout1(x))) for _ in range(5)],\n                dim=0,\n            ),\n            dim=0,\n        )\n        x = torch.mean(\n            torch.stack(\n                [self.out(self.dropout2(x)) for _ in range(5)],\n                dim=0,\n            ),\n            dim=0,\n        )\n        \n        return x\n\nclass Attention(nn.Module):\n    def __init__(self, feature_dim, step_dim, bias=True, **kwargs):\n        super(Attention, self).__init__(**kwargs)\n\n        self.supports_masking = True\n\n        self.bias = bias\n        self.feature_dim = feature_dim\n        self.step_dim = step_dim\n        self.features_dim = 0\n\n        weight = torch.zeros(feature_dim, 1)\n        nn.init.xavier_uniform_(weight)\n        self.weight = nn.Parameter(weight)\n\n        if bias:\n            self.b = nn.Parameter(torch.zeros(step_dim))\n\n    def forward(self, x, mask=None):\n        feature_dim = self.feature_dim\n        step_dim = x.shape[1] #self.step_dim\n\n        eij = torch.mm(\n            x.contiguous().view(-1, feature_dim),\n            self.weight\n        ).view(-1, step_dim)\n\n        if self.bias:\n            eij = eij + self.b\n\n        eij = torch.tanh(eij)\n        a = torch.exp(eij)\n\n        if mask is not None:\n            a = a * mask\n\n        a = a / torch.sum(a, 1, keepdim=True) + 1e-10\n\n        weighted_input = x * torch.unsqueeze(a, -1)\n        return torch.sum(weighted_input, 1)\n\nsigmoid = torch.nn.Sigmoid()\nclass Swish(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, i):\n        result = i * sigmoid(i)\n        ctx.save_for_backward(i)\n        return result\n    @staticmethod\n    def backward(ctx, grad_output):\n        i = ctx.saved_variables[0]\n        sigmoid_i = sigmoid(i)\n        return grad_output * (sigmoid_i * (1 + i * (1 - sigmoid_i)))\n\nclass Swish_module(nn.Module):\n    def forward(self, x):\n        return Swish.apply(x)\n\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, pretrained=False, dropout=0.1):\n        super().__init__()\n        \n        if cfg.model == 'coat_lite_medium':\n            self.model = timm.create_model('coat_lite_medium', pretrained=False)#, checkpoint_path='./coat_lite_medium_384x384_f9129688.pth')\n        elif cfg.model == 'vit256':\n            self.model = get_vit256('vit256_small_dino.pth')\n            #self.model = get_vit256('fake')\n        else:\n            self.model = timm.create_model(cfg.model, pretrained=pretrained)\n        self.seq_len = cfg.seq_len\n        self.model_reformat = cfg.model_reformat\n        \n        if cfg.model == 'vit256':\n            feat_dim = 384\n        else:\n            feat_dim = self.model.get_classifier().in_features\n\n        self.norm = nn.LayerNorm(feat_dim)\n\n        self.att = Attention(feat_dim, cfg.num_instance, bias=False)\n\n        #self.max_pool = nn.AdaptiveMaxPool1d(1)\n        self.fc = nn.Sequential(\n            nn.Linear(feat_dim, len(cfg.target_ohe_columns))\n        )\n\n        if cfg.model == 'vit256':\n            pass\n        else:\n            self.model.reset_classifier(num_classes=0, global_pool=\"avg\")\n        \n\n    def forward(self, images):\n        b, n, c, w, h = images.shape\n        x = images.view(b*n, c, w, h) # (bxn, c, w, h)\n        \n        x = self.model.forward_features(x) # (bxn, 144, f)\n        \n        if self.model_reformat:\n            x = x.view(x.shape[0], x.shape[1], -1).permute(0, 2, 1) # (b, c, w, h)->(b, c, wxh) -> (b, wxh, c)\n            x = torch.mean(x, 1)[:,None,:]\n        #print(x.shape); assert False\n        \n        x = x.view(b, -1, x.shape[2]) # (b, nx144, f)\n\n        x = self.norm(x)\n        x = self.att(x)\n\n        output = self.fc(x)\n        return output\n    \n# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef get_valid_transforms(cfg):\n    return A.Compose([\n            #A.CenterCrop(img_size,img_size, p=1.),\n            A.Resize(cfg.image_size, cfg.image_size, interpolation=cv2.INTER_LANCZOS4),      \n            #A.LongestMaxSize(max_size=self.cfg['image_size'], interpolation=cv2.INTER_LANCZOS4, always_apply=True, p=1),\n            #A.PadIfNeeded(min_height=self.cfg['image_size'], min_width=self.cfg['image_size']),\n            #A.PadIfNeeded(min_height=1280, min_width=1280),\n            #A.Resize(self.cfg['image_size'], self.cfg['image_size'], interpolation=cv2.INTER_LANCZOS4),      \n            \n            A.Normalize(mean=[0], std=[1], max_pixel_value=255.0, p=1.0),\n            ToTensorV2(p=1.0),\n        ], p=1.0)\n       \n@torch.no_grad()\ndef get_predictions(dl, model, device):\n    \n    model.eval()\n    preds = []\n    for step, (images, _) in enumerate(tqdm(dl)):\n        images = images.to(device).half()\n        \n        y_preds = model(images)\n        y_preds = F.softmax(y_preds, dim=1).to('cpu').numpy() # (b, 23, 5)\n        \n        preds.append(y_preds)\n        \n    predictions = np.concatenate(preds)\n    return predictions\n\nif __name__ == '__main__':\n    \n    test = merge_image_info(parse_images('./test/'), test)\n    pred_cols = ['CE', 'LAA']\n    for c in pred_cols:\n        test[c] = 0.0\n        \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n    all_preds = []\n    all_weights = [2.5, 1. ]\n    cfgs = [CFG, CFG2]\n    for cfg in cfgs:\n        \n        ds = TrainDataset(cfg, test, get_valid_transforms(cfg))\n        dl = DataLoader(ds, batch_size=cfg.batch_size, shuffle=False,\n                        num_workers=2, pin_memory=True, drop_last=False)\n        \n        model = CustomModel(cfg, pretrained=False, dropout=0.1).to(device)\n        model.half()\n        fold_preds = []\n        for fold in range(cfg.n_fold):\n            ckpt = f\"{cfg.model_name}/{cfg.model}_fold{fold}_best.pth\"\n        \n            model.load_state_dict(torch.load(ckpt, map_location=torch.device('cuda'))['model'])\n            fold_preds += [get_predictions(dl, model, device)]\n        fold_preds = np.mean(fold_preds, 0)\n        all_preds += [fold_preds]\n    all_preds = np.average(all_preds, axis=0, weights=all_weights)\n    \n    test.loc[test.instance_id==0, pred_cols] = all_preds\n    submission = test.loc[test.instance_id==0].groupby('patient_id')[pred_cols].mean().reset_index()\n    submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-09-08T06:50:53.364281Z","iopub.execute_input":"2022-09-08T06:50:53.364645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}