{"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":"!pip install timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-29T16:14:13.339594Z","iopub.execute_input":"2022-07-29T16:14:13.340371Z","iopub.status.idle":"2022-07-29T16:14:27.466455Z","shell.execute_reply.started":"2022-07-29T16:14:13.340328Z","shell.execute_reply":"2022-07-29T16:14:27.465464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn\nimport torch\nfrom collections import OrderedDict\nimport numpy as np\nimport torch.nn.functional as F\nfrom PIL import Image\nimport os\nfrom torchvision.transforms import Compose, Resize, CenterCrop, ToTensor, Normalize\nfrom tqdm import tqdm\nimport sys\nsys.path.append(\"../input/timmmaster/\")\nimport timm\ntry:\n    from torchvision.transforms import InterpolationMode\n    BICUBIC = InterpolationMode.BICUBIC\nexcept ImportError:\n    BICUBIC = Image.BICUBIC\nfrom torch.nn import LayerNorm\n\n\nclass QuickGELU(nn.Module):\n    def forward(self, x: torch.Tensor):\n        return x * torch.sigmoid(1.702 * x)\n\n\nclass ResidualAttentionBlock(nn.Module):\n    def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):\n        super().__init__()\n\n        self.attn = nn.MultiheadAttention(d_model, n_head)\n        self.ln_1 = LayerNorm(d_model)\n        self.mlp = nn.Sequential(OrderedDict([\n            (\"c_fc\", nn.Linear(d_model, d_model * 4)),\n            (\"gelu\", QuickGELU()),\n            (\"c_proj\", nn.Linear(d_model * 4, d_model))\n        ]))\n        self.ln_2 = LayerNorm(d_model)\n        self.attn_mask = attn_mask\n\n    def attention(self, x: torch.Tensor):\n        self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None\n        return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]\n\n    def forward(self, x: torch.Tensor):\n        x = x + self.attention(self.ln_1(x))\n        x = x + self.mlp(self.ln_2(x))\n        return x\n\n\nclass Transformer(nn.Module):\n    def __init__(self, width: int, layers: int, heads: int, attn_mask: torch.Tensor = None):\n        super().__init__()\n        self.width = width\n        self.layers = layers\n        self.resblocks = nn.Sequential(*[ResidualAttentionBlock(width, heads, attn_mask) for _ in range(layers)])\n\n    def forward(self, x: torch.Tensor):\n        return self.resblocks(x)\n\n\nclass VisionTransformer(nn.Module):\n    def __init__(self, input_resolution: int, patch_size: int, width: int, layers: int, heads: int, output_dim: int):\n        super().__init__()\n        self.input_resolution = input_resolution\n        self.output_dim = output_dim\n        self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False)\n\n        scale = width ** -0.5\n        self.class_embedding = nn.Parameter(scale * torch.randn(width))\n        self.positional_embedding = nn.Parameter(scale * torch.randn((input_resolution // patch_size) ** 2 + 1, width))\n        self.ln_pre = LayerNorm(width)\n\n        self.transformer = Transformer(width, layers, heads)\n\n        self.ln_post = LayerNorm(width)\n        self.proj = nn.Parameter(scale * torch.randn(width, output_dim))\n        #self.text_embs = torch.load('../input/text-embs2/text_embs.pt')\n    def forward(self, x: torch.Tensor):\n        x = self.conv1(x)  # shape = [*, width, grid, grid]\n        x = x.reshape(x.shape[0], x.shape[1], -1)  # shape = [*, width, grid ** 2]\n        x = x.permute(0, 2, 1)  # shape = [*, grid ** 2, width]\n        x = torch.cat([self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), x], dim=1)  # shape = [*, grid ** 2 + 1, width]\n        x = x + self.positional_embedding.to(x.dtype)\n        x = self.ln_pre(x)\n\n        x = x.permute(1, 0, 2)  # NLD -> LND\n        x = self.transformer(x)\n        x = x.permute(1, 0, 2)  # LND -> NLD\n\n        x = self.ln_post(x[:, 0, :])\n\n        if self.proj is not None:\n            \n            \n            x = x @ self.proj\n            #m = torch.argmax(x @ self.text_embs.T.to(x), dim=-1)[0]\n        return x #+ self.text_embs[m].to(x)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T16:14:27.468698Z","iopub.execute_input":"2022-07-29T16:14:27.469032Z","iopub.status.idle":"2022-07-29T16:14:28.638678Z","shell.execute_reply.started":"2022-07-29T16:14:27.468999Z","shell.execute_reply":"2022-07-29T16:14:28.637596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms\n\nclass MyModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.feature_extractor = VisionTransformer(\n              input_resolution=336,\n              patch_size=14,\n              width=1024,\n              layers=24,\n              heads=16,\n              output_dim=768\n          )\n        \n        self.feature_extractor.load_state_dict(torch.load('../input/vitl14336/ViT-L14_visual336px.pt'))\n        self.feature_extractor2 = timm.create_model(\"convnext_xlarge_384_in22ft1k\", pretrained=True, num_classes=0)\n        self.feature_extractor4 = nn.AdaptiveAvgPool1d(64)\n        self.feature_extractor3 = nn.AdaptiveAvgPool1d(64)\n    def forward(self, x):\n        x1 = transforms.functional.resize(x,size=[336, 336])\n        x1 = x1 / 255.0\n        x1 = transforms.functional.normalize(x1, \n                                            mean=[0.48145466, 0.4578275, 0.40821073], \n                                            std=[0.26862954, 0.26130258, 0.27577711])\n        x1 = self.feature_extractor(x1)\n        x1 = self.feature_extractor3(x1)\n        x2 = transforms.functional.resize(x,size=[384, 384])\n        x2 = x2/255.0\n        x2 = transforms.functional.normalize(x2, \n                                                mean=[0.485, 0.456, 0.406], \n                                                std=[0.229, 0.224, 0.225])\n        x2 = self.feature_extractor4(self.feature_extractor2(x2))\n        return x1 + x2\nmodel = MyModel()\nmodel.eval()\nsaved_model = torch.jit.script(model)\nsaved_model.save('saved_model.pt')","metadata":{"execution":{"iopub.status.busy":"2022-07-29T16:14:28.640334Z","iopub.execute_input":"2022-07-29T16:14:28.640721Z","iopub.status.idle":"2022-07-29T16:15:44.037317Z","shell.execute_reply.started":"2022-07-29T16:14:28.640686Z","shell.execute_reply":"2022-07-29T16:15:44.034179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from zipfile import ZipFile\nwith ZipFile('submission.zip','w') as zip:           \n    zip.write('./saved_model.pt', arcname='saved_model.pt') ","metadata":{},"execution_count":null,"outputs":[]}]}