{"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":"markdown","source":"This notebook is based on [PyTorch - Ensemble Pretrained Baselines [Training]](https://www.kaggle.com/code/carloalbertobarbano/pytorch-ensemble-pretrained-baselines-training). Thank you.","metadata":{}},{"cell_type":"code","source":"import sys\nfrom functools import partial","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:26.123083Z","iopub.execute_input":"2022-07-13T11:31:26.123568Z","iopub.status.idle":"2022-07-13T11:31:26.157371Z","shell.execute_reply.started":"2022-07-13T11:31:26.12347Z","shell.execute_reply":"2022-07-13T11:31:26.15637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torchvision import models\nfrom torchvision import transforms","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-13T11:31:26.197665Z","iopub.execute_input":"2022-07-13T11:31:26.198408Z","iopub.status.idle":"2022-07-13T11:31:28.524274Z","shell.execute_reply.started":"2022-07-13T11:31:26.198369Z","shell.execute_reply":"2022-07-13T11:31:28.523129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# tmp to download dino source code and weight\ntmp = torch.hub.load('facebookresearch/dino:main', 'dino_vits16')","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:28.526298Z","iopub.execute_input":"2022-07-13T11:31:28.526892Z","iopub.status.idle":"2022-07-13T11:31:39.800965Z","shell.execute_reply.started":"2022-07-13T11:31:28.526855Z","shell.execute_reply":"2022-07-13T11:31:39.799976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sys.path.append('/root/.cache/torch/hub/facebookresearch_dino_main')","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:39.802586Z","iopub.execute_input":"2022-07-13T11:31:39.80303Z","iopub.status.idle":"2022-07-13T11:31:39.808883Z","shell.execute_reply.started":"2022-07-13T11:31:39.802978Z","shell.execute_reply":"2022-07-13T11:31:39.807571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# code from https://github.com/facebookresearch/dino/blob/main/vision_transformer.py\n# Little fix for jit\n\n\"\"\"\nMostly copy-paste from timm library.\nhttps://github.com/rwightman/pytorch-image-models/blob/master/timm/models/vision_transformer.py\n\"\"\"\nimport math\nfrom functools import partial\n\nimport torch\nimport torch.nn as nn\n\nfrom utils import trunc_normal_\n\n\ndef drop_path(x, drop_prob: float = 0., training: bool = False):\n    if drop_prob == 0. or not training:\n        return x\n    keep_prob = 1 - drop_prob\n    shape = (x.shape[0],) + (1,) * (x.ndim - 1)  # work with diff dim tensors, not just 2D ConvNets\n    random_tensor = keep_prob + torch.rand(shape, dtype=x.dtype, device=x.device)\n    random_tensor.floor_()  # binarize\n    output = x.div(keep_prob) * random_tensor\n    return output\n\n\nclass DropPath(nn.Module):\n    \"\"\"Drop paths (Stochastic Depth) per sample  (when applied in main path of residual blocks).\n    \"\"\"\n    def __init__(self, drop_prob=None):\n        super(DropPath, self).__init__()\n        self.drop_prob = drop_prob\n\n    def forward(self, x):\n        return drop_path(x, self.drop_prob, self.training)\n\n\nclass Mlp(nn.Module):\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 Attention(nn.Module):\n    def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0.):\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)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n    def forward(self, x):\n        B, N, C = x.shape\n        qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale\n        attn = attn.softmax(dim=-1)\n        attn = self.attn_drop(attn)\n\n        x = (attn @ v).transpose(1, 2).reshape(B, N, C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x, attn\n\n\nclass Block(nn.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        super().__init__()\n        self.norm1 = norm_layer(dim)\n        self.attn = Attention(\n            dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\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):\n        y, attn = self.attn(self.norm1(x))\n        x = x + self.drop_path(y)\n        x = x + self.drop_path(self.mlp(self.norm2(x)))\n        return x\n\n\nclass PatchEmbed(nn.Module):\n    \"\"\" Image to Patch Embedding\n    \"\"\"\n    def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):\n        super().__init__()\n        num_patches = (img_size // patch_size) * (img_size // patch_size)\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.num_patches = num_patches\n\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        x = self.proj(x).flatten(2).transpose(1, 2)\n        return x\n\n\nclass VisionTransformer(nn.Module):\n    \"\"\" Vision Transformer \"\"\"\n    def __init__(self, img_size=[224], patch_size=16, in_chans=3, num_classes=0, embed_dim=768, depth=12,\n                 num_heads=12, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop_rate=0., attn_drop_rate=0.,\n                 drop_path_rate=0., norm_layer=nn.LayerNorm, **kwargs):\n        super().__init__()\n        self.num_features = self.embed_dim = embed_dim\n\n        self.patch_embed = PatchEmbed(\n            img_size=img_size[0], patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim)\n        num_patches = self.patch_embed.num_patches\n\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))\n        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))\n        self.pos_drop = nn.Dropout(p=drop_rate)\n\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]  # stochastic depth decay rule\n        self.blocks = nn.ModuleList([\n            Block(\n                dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,\n                drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[i], norm_layer=norm_layer)\n            for i in range(depth)])\n        self.norm = norm_layer(embed_dim)\n\n        # Classifier head\n        self.head = nn.Linear(embed_dim, num_classes) if num_classes > 0 else nn.Identity()\n\n        trunc_normal_(self.pos_embed, std=.02)\n        trunc_normal_(self.cls_token, 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    def interpolate_pos_encoding(self, x, w: int, h: int):\n        npatch = x.shape[1] - 1\n        N = self.pos_embed.shape[1] - 1\n        if npatch == N and w == h:\n            return self.pos_embed\n        class_pos_embed = self.pos_embed[:, 0]\n        patch_pos_embed = self.pos_embed[:, 1:]\n        dim = x.shape[-1]\n        w0 = w // self.patch_embed.patch_size\n        h0 = h // self.patch_embed.patch_size\n        # we add a small number to avoid floating point error in the interpolation\n        # see discussion at https://github.com/facebookresearch/dino/issues/8\n        w0, h0 = w0 + 0.1, h0 + 0.1\n        patch_pos_embed = nn.functional.interpolate(\n            patch_pos_embed.reshape(1, int(math.sqrt(N)), int(math.sqrt(N)), dim).permute(0, 3, 1, 2),\n            scale_factor=[float(w0 / math.sqrt(N)), float(h0 / math.sqrt(N))],\n            mode='bicubic',\n        )\n        assert int(w0) == patch_pos_embed.shape[-2] and int(h0) == patch_pos_embed.shape[-1]\n        patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)\n        return torch.cat((class_pos_embed.unsqueeze(0), patch_pos_embed), dim=1)\n\n    def prepare_tokens(self, x):\n        B, nc, w, h = x.shape\n        x = self.patch_embed(x)  # patch linear embedding\n\n        # add the [CLS] token to the embed patch tokens\n        cls_tokens = self.cls_token.expand(B, -1, -1)\n        x = torch.cat((cls_tokens, x), dim=1)\n\n        # add positional encoding to each token\n        x = x + self.interpolate_pos_encoding(x, w, h)\n\n        return self.pos_drop(x)\n\n    def forward(self, x):\n        x = self.prepare_tokens(x)\n        for blk in self.blocks:\n            x = blk(x)\n        x = self.norm(x)\n        return x[:, 0]\n\n    def get_last_selfattention(self, x):\n        x = self.prepare_tokens(x)\n        for i, blk in enumerate(self.blocks):\n            if i < len(self.blocks) - 1:\n                x = blk(x)\n            else:\n                # return attention of the last block\n                return blk(x, return_attention=True)\n\n    def get_intermediate_layers(self, x, n=1):\n        x = self.prepare_tokens(x)\n        # we return the output tokens from the `n` last blocks\n        output = []\n        for i, blk in enumerate(self.blocks):\n            x = blk(x)\n            if len(self.blocks) - i <= n:\n                output.append(self.norm(x))\n        return output\n","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:39.812497Z","iopub.execute_input":"2022-07-13T11:31:39.813238Z","iopub.status.idle":"2022-07-13T11:31:39.870777Z","shell.execute_reply.started":"2022-07-13T11:31:39.81319Z","shell.execute_reply":"2022-07-13T11:31:39.869799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def vit_tiny(patch_size=16, **kwargs):\n    model = VisionTransformer(\n        patch_size=patch_size, embed_dim=192, depth=12, num_heads=3, mlp_ratio=4,\n        qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)\n    return model\n\n\ndef vit_small(patch_size=16, **kwargs):\n    model = VisionTransformer(\n        patch_size=patch_size, embed_dim=384, depth=12, num_heads=6, mlp_ratio=4,\n        qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)\n    return model\n\n\ndef vit_base(patch_size=16, **kwargs):\n    model = VisionTransformer(\n        patch_size=patch_size, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4,\n        qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:39.872294Z","iopub.execute_input":"2022-07-13T11:31:39.873099Z","iopub.status.idle":"2022-07-13T11:31:39.889474Z","shell.execute_reply.started":"2022-07-13T11:31:39.873052Z","shell.execute_reply":"2022-07-13T11:31:39.888248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DINO(nn.Module):\n    def __init__(self, encoder, resize, size):\n        super().__init__()\n        self.encoder = encoder\n        self.pool = nn.AdaptiveAvgPool1d(64)\n        self.resize = resize \n        self.size = size\n        \n    def forward(self, x):\n        x = transforms.functional.resize(x, [self.resize, self.resize])\n        x = transforms.functional.center_crop(x, [self.size, self.size])\n        x = x / 255.\n        x = transforms.functional.normalize(x, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        x = self.encoder(x)\n        x = self.pool(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:39.893114Z","iopub.execute_input":"2022-07-13T11:31:39.89357Z","iopub.status.idle":"2022-07-13T11:31:39.903467Z","shell.execute_reply.started":"2022-07-13T11:31:39.893535Z","shell.execute_reply":"2022-07-13T11:31:39.90254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = vit_small(patch_size=16)\nencoder.load_state_dict(\n    tmp.state_dict(),\n)\n\ndino = DINO(encoder, 256, 224)\ndino.eval()","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:39.90494Z","iopub.execute_input":"2022-07-13T11:31:39.905943Z","iopub.status.idle":"2022-07-13T11:31:41.449711Z","shell.execute_reply.started":"2022-07-13T11:31:39.905898Z","shell.execute_reply":"2022-07-13T11:31:41.448565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dino(torch.randn((1, 3, 256, 256))).shape","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:41.451086Z","iopub.execute_input":"2022-07-13T11:31:41.45145Z","iopub.status.idle":"2022-07-13T11:31:43.466118Z","shell.execute_reply.started":"2022-07-13T11:31:41.451416Z","shell.execute_reply":"2022-07-13T11:31:43.464991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"saved_model = torch.jit.script(dino)\nsaved_model.save('saved_model.pt')","metadata":{"execution":{"iopub.status.busy":"2022-07-13T11:31:43.467705Z","iopub.execute_input":"2022-07-13T11:31:43.468617Z","iopub.status.idle":"2022-07-13T11:31:44.907615Z","shell.execute_reply.started":"2022-07-13T11:31:43.46857Z","shell.execute_reply":"2022-07-13T11:31:44.906419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from zipfile import ZipFile\n\nwith ZipFile('submission.zip','w') as zip:           \n    zip.write('saved_model.pt', arcname='saved_model.pt') ","metadata":{},"execution_count":null,"outputs":[]}]}