{"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":"<a id=\"\"><font color='#425066'><h3>Install Timm</h3></font></a>\n\nWe will be using timm which we first need to install since its not available in Kaggle environment.\n","metadata":{}},{"cell_type":"code","source":"!pip -q install timm","metadata":{"execution":{"iopub.status.busy":"2022-08-02T05:18:48.366235Z","iopub.execute_input":"2022-08-02T05:18:48.366880Z","iopub.status.idle":"2022-08-02T05:18:57.434446Z","shell.execute_reply.started":"2022-08-02T05:18:48.366840Z","shell.execute_reply":"2022-08-02T05:18:57.433308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"\"><font color='#425066'><h3>Import Dependencies</h3></font></a>\n\nImport the neccessary libraries - timm, torch, torchvision and PIL","metadata":{}},{"cell_type":"code","source":"import timm\nimport torch\nimport torch.nn as nn\nfrom torchvision import transforms\nfrom PIL import Image","metadata":{"execution":{"iopub.status.busy":"2022-08-02T05:19:10.277625Z","iopub.execute_input":"2022-08-02T05:19:10.278312Z","iopub.status.idle":"2022-08-02T05:19:10.285375Z","shell.execute_reply.started":"2022-08-02T05:19:10.278262Z","shell.execute_reply":"2022-08-02T05:19:10.284426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"\"><font color='#425066'><h3>Model</h3></font></a>\n\n- Here, we are using ConvNext Base pre-trained on Imagenet-22k and further finetuned on [130k Images (128x128) - Universal Image Embeddings](https://www.kaggle.com/datasets/rhtsingh/google-universal-image-embeddings-128x128) dataset.\n\n- We will be using image size of 384.\n\n- We will change to `nn.Identity()` as final layer which will allow us to only extract the embeddings and not logits.\n\n- `forward()` does the necessary transformations like resize and normalization.","metadata":{}},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(\n        self, model_name='convnext_base_in22k', \n        pretrained=True, num_classes=11, image_size=384,\n        checkpoint_path='src/runs/exp/weights/best.pt'\n    ):\n        super().__init__()\n        self.model = timm.create_model(\n            model_name,\n            pretrained=pretrained, \n            num_classes=num_classes, \n            checkpoint_path=checkpoint_path\n        )\n        self.model.head.fc = nn.Identity()\n        self.pool = nn.AdaptiveAvgPool1d(64)\n        self.image_size = image_size\n\n    def forward(self, image):\n        image = transforms.functional.resize(image, size=[self.image_size, self.image_size])\n        image = image / 255.0\n        image = transforms.functional.normalize(image, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])        \n        with torch.no_grad():\n            output = self.model(image)\n        output = self.pool(output)\n        output = torch.nn.functional.normalize(output)\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-08-02T05:19:18.035225Z","iopub.execute_input":"2022-08-02T05:19:18.035596Z","iopub.status.idle":"2022-08-02T05:19:18.044648Z","shell.execute_reply.started":"2022-08-02T05:19:18.035555Z","shell.execute_reply":"2022-08-02T05:19:18.043597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"\"><font color='#425066'><h3>Load from Checkpoint</h3></font></a>\n\nWe will load our model weights from finetuned checkpoint path","metadata":{}},{"cell_type":"code","source":"model = Model(\n    model_name='convnext_base_in22k', pretrained=False, \n    num_classes=11, image_size=384,\n    checkpoint_path='../input/convnextbaseuiefinetuned/runs/exp/weights/best.pt'\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-02T05:19:20.213042Z","iopub.execute_input":"2022-08-02T05:19:20.213674Z","iopub.status.idle":"2022-08-02T05:19:25.245262Z","shell.execute_reply.started":"2022-08-02T05:19:20.213637Z","shell.execute_reply":"2022-08-02T05:19:25.244156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"\"><font color='#425066'><h3>Extract Embeddings</h3></font></a>\n\nThe following code is to check if the model is extracting embeddings correctly","metadata":{}},{"cell_type":"code","source":"image = Image.open('../input/convnextbaseuiefinetuned/sample/image0000.jpeg').convert(\"RGB\")\nconvert_to_tensor = transforms.Compose([transforms.PILToTensor()])\ninput_tensor = convert_to_tensor(image)\ninput_batch = input_tensor.unsqueeze(0)\nembedding = torch.flatten(model(input_batch)[0]).cpu().data.numpy()\nprint(\"Embedding shape: \", embedding.shape)","metadata":{"execution":{"iopub.status.busy":"2022-08-02T05:19:27.719850Z","iopub.execute_input":"2022-08-02T05:19:27.720214Z","iopub.status.idle":"2022-08-02T05:19:28.792770Z","shell.execute_reply.started":"2022-08-02T05:19:27.720183Z","shell.execute_reply":"2022-08-02T05:19:28.791729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"\"><font color='#425066'><h3>PyTorch's Torchscript Format</h3></font></a>\n\nOne can learn more about Torchscript [here](https://pytorch.org/docs/stable/jit.html)","metadata":{}},{"cell_type":"code","source":"model_scripted = torch.jit.script(model)\nmodel_scripted.save('saved_model.pt')","metadata":{"execution":{"iopub.status.busy":"2022-08-02T05:19:33.858433Z","iopub.execute_input":"2022-08-02T05:19:33.858791Z","iopub.status.idle":"2022-08-02T05:19:35.437550Z","shell.execute_reply.started":"2022-08-02T05:19:33.858760Z","shell.execute_reply":"2022-08-02T05:19:35.436524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"\"><font color='#425066'><h3>Submission</h3></font></a>\n\nWe create a `submission.zip` file containing the model.","metadata":{}},{"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')\n    \n!rm saved_model.pt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-02T05:19:39.351490Z","iopub.execute_input":"2022-08-02T05:19:39.351869Z","iopub.status.idle":"2022-08-02T05:19:41.305993Z","shell.execute_reply.started":"2022-08-02T05:19:39.351815Z","shell.execute_reply":"2022-08-02T05:19:41.304709Z"},"trusted":true},"execution_count":null,"outputs":[]}]}