{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-25T21:40:35.983942Z","iopub.execute_input":"2022-07-25T21:40:35.984543Z","iopub.status.idle":"2022-07-25T21:40:36.009417Z","shell.execute_reply.started":"2022-07-25T21:40:35.984465Z","shell.execute_reply":"2022-07-25T21:40:36.008617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install transformers","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:40:36.010904Z","iopub.execute_input":"2022-07-25T21:40:36.011295Z","iopub.status.idle":"2022-07-25T21:40:48.130465Z","shell.execute_reply.started":"2022-07-25T21:40:36.011262Z","shell.execute_reply":"2022-07-25T21:40:48.129356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import ViTFeatureExtractor, ViTModel\nfrom PIL import Image\nimport requests\n\nfeature_extractor = ViTFeatureExtractor.from_pretrained('google/vit-base-patch16-224-in21k', torchscript=True)\nmodel = ViTModel.from_pretrained('google/vit-base-patch16-224-in21k', torchscript=True)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:41:32.470049Z","iopub.execute_input":"2022-07-25T21:41:32.471318Z","iopub.status.idle":"2022-07-25T21:41:59.859519Z","shell.execute_reply.started":"2022-07-25T21:41:32.471256Z","shell.execute_reply":"2022-07-25T21:41:59.858193Z"},"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\nfrom torchvision.transforms.functional import InterpolationMode","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:42:05.230494Z","iopub.execute_input":"2022-07-25T21:42:05.231047Z","iopub.status.idle":"2022-07-25T21:42:05.454032Z","shell.execute_reply.started":"2022-07-25T21:42:05.231015Z","shell.execute_reply":"2022-07-25T21:42:05.453138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:42:05.517689Z","iopub.execute_input":"2022-07-25T21:42:05.518017Z","iopub.status.idle":"2022-07-25T21:42:05.524296Z","shell.execute_reply.started":"2022-07-25T21:42:05.517989Z","shell.execute_reply":"2022-07-25T21:42:05.523037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:42:05.815661Z","iopub.execute_input":"2022-07-25T21:42:05.815981Z","iopub.status.idle":"2022-07-25T21:42:05.821396Z","shell.execute_reply.started":"2022-07-25T21:42:05.815956Z","shell.execute_reply":"2022-07-25T21:42:05.819690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:42:07.106207Z","iopub.execute_input":"2022-07-25T21:42:07.107310Z","iopub.status.idle":"2022-07-25T21:42:07.256426Z","shell.execute_reply.started":"2022-07-25T21:42:07.107237Z","shell.execute_reply":"2022-07-25T21:42:07.255430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VitModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.model = ViTModel.from_pretrained('google/vit-base-patch16-224-in21k', torchscript=True)\n        self.feature_extractor = ViTFeatureExtractor.from_pretrained('google/vit-base-patch16-224-in21k', torchscript=True)\n        self.pool = nn.AdaptiveAvgPool1d(64)\n        \n    def forward(self, x):\n        #x = transforms.functional.resize(x,size=[224, 224], interpolation=transforms.InterpolationMode.BICUBIC)\n#         x = x/255.0\n#         x = transforms.functional.normalize(x,\n#                                             mean=[0.48145466, 0.4578275, 0.40821073], \n#                                             std=[0.26862954, 0.26130258, 0.27577711])\n\n        x = torch.squeeze(x, 0)\n        inputs = feature_extractor(x, return_tensors=\"pt\")\n        inputs = inputs.to(torch.device('cuda:0'))\n        \n        x = self.model(inputs[\"pixel_values\"])\n        return self.pool(x[1])","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:46:11.949475Z","iopub.execute_input":"2022-07-25T21:46:11.949813Z","iopub.status.idle":"2022-07-25T21:46:11.958249Z","shell.execute_reply.started":"2022-07-25T21:46:11.949786Z","shell.execute_reply":"2022-07-25T21:46:11.957366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_image_model():\n    model = VitModel().to(torch.device('cuda:0'))\n    model = model.eval()\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:46:12.125085Z","iopub.execute_input":"2022-07-25T21:46:12.125784Z","iopub.status.idle":"2022-07-25T21:46:12.131534Z","shell.execute_reply.started":"2022-07-25T21:46:12.125745Z","shell.execute_reply":"2022-07-25T21:46:12.130486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = load_image_model()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:46:12.264628Z","iopub.execute_input":"2022-07-25T21:46:12.265757Z","iopub.status.idle":"2022-07-25T21:46:14.631913Z","shell.execute_reply.started":"2022-07-25T21:46:12.265710Z","shell.execute_reply":"2022-07-25T21:46:14.630924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = torch.randn(1, 3, 4545, 2324)\nprint(x.shape)\n# Let's print it\nmodel(x).shape","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:46:14.634059Z","iopub.execute_input":"2022-07-25T21:46:14.634448Z","iopub.status.idle":"2022-07-25T21:46:15.220016Z","shell.execute_reply.started":"2022-07-25T21:46:14.634412Z","shell.execute_reply":"2022-07-25T21:46:15.219010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"traced_model = torch.jit.trace(model, torch.randn(1, 3, 224, 224))\ntorch.jit.save(traced_model, 'saved_model.pt')","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:49:09.948251Z","iopub.execute_input":"2022-07-25T21:49:09.948848Z","iopub.status.idle":"2022-07-25T21:49:11.904173Z","shell.execute_reply.started":"2022-07-25T21:49:09.948809Z","shell.execute_reply":"2022-07-25T21:49:11.903163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from PIL import Image\n# import torch\n# from torchvision import transforms\n\n# # Model loading.\n# model = torch.jit.load(\"./saved_model.pt\")\n# model.eval()\n# embedding_fn = model\n\n# # Load image and extract its embedding.\n# input_image = Image.open(\"../input/lost-poets-nft/lostpoets/lostpoet_10004.jpg\").convert(\"RGB\")\n# convert_to_tensor = transforms.Compose([transforms.PILToTensor()])\n# input_tensor = convert_to_tensor(input_image)\n# input_batch = input_tensor.unsqueeze(0)\n# with torch.no_grad():\n#     embedding = torch.flatten(embedding_fn(input_batch)[0]).cpu().data.numpy()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:49:16.263770Z","iopub.execute_input":"2022-07-25T21:49:16.264110Z","iopub.status.idle":"2022-07-25T21:49:16.819957Z","shell.execute_reply.started":"2022-07-25T21:49:16.264082Z","shell.execute_reply":"2022-07-25T21:49:16.818980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(embedding.shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:49:18.772361Z","iopub.execute_input":"2022-07-25T21:49:18.772722Z","iopub.status.idle":"2022-07-25T21:49:18.778138Z","shell.execute_reply.started":"2022-07-25T21:49:18.772689Z","shell.execute_reply":"2022-07-25T21:49:18.777166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Check is done","metadata":{}},{"cell_type":"code","source":"# # Reload traced model and compare outputs to original one\n# loaded_model = torch.jit.load('saved_model.pt')\n# loaded_model.eval()\n# traced_outputs = loaded_model(x)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:46:18.548225Z","iopub.execute_input":"2022-07-24T08:46:18.548579Z","iopub.status.idle":"2022-07-24T08:46:18.553232Z","shell.execute_reply.started":"2022-07-24T08:46:18.548551Z","shell.execute_reply":"2022-07-24T08:46:18.551988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# saved_model = torch.jit.script(model)\n# saved_model.save('saved_model.pt')","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:46:18.742934Z","iopub.execute_input":"2022-07-24T08:46:18.744027Z","iopub.status.idle":"2022-07-24T08:46:18.748368Z","shell.execute_reply.started":"2022-07-24T08:46:18.743989Z","shell.execute_reply":"2022-07-24T08:46:18.747144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Final save","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')","metadata":{"execution":{"iopub.status.busy":"2022-07-25T21:49:35.889488Z","iopub.execute_input":"2022-07-25T21:49:35.889837Z","iopub.status.idle":"2022-07-25T21:49:36.763436Z","shell.execute_reply.started":"2022-07-25T21:49:35.889808Z","shell.execute_reply":"2022-07-25T21:49:36.762414Z"},"trusted":true},"execution_count":null,"outputs":[]}]}