{"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 torchinfo","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:30:25.028026Z","iopub.execute_input":"2022-08-05T23:30:25.028734Z","iopub.status.idle":"2022-08-05T23:30:35.715157Z","shell.execute_reply.started":"2022-08-05T23:30:25.028697Z","shell.execute_reply":"2022-08-05T23:30:35.713973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nimport torch\nimport torch.nn as nn\nimport torchvision\nfrom torchvision import models\nfrom torchvision import transforms  \nimport torch.nn.functional as F\nfrom sklearn.decomposition import PCA\nimport numpy as np\nfrom tqdm import tqdm ","metadata":{"execution":{"iopub.status.busy":"2022-08-05T22:01:01.185749Z","iopub.execute_input":"2022-08-05T22:01:01.186628Z","iopub.status.idle":"2022-08-05T22:01:08.550176Z","shell.execute_reply.started":"2022-08-05T22:01:01.186566Z","shell.execute_reply":"2022-08-05T22:01:08.549019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = timm.create_model(\"swin_large_patch4_window12_384_in22k\", pretrained=True, num_classes=0)    \nmodel.to('cuda')","metadata":{"execution":{"iopub.status.busy":"2022-08-05T22:01:08.551756Z","iopub.execute_input":"2022-08-05T22:01:08.553805Z","iopub.status.idle":"2022-08-05T22:01:37.969435Z","shell.execute_reply.started":"2022-08-05T22:01:08.553761Z","shell.execute_reply":"2022-08-05T22:01:37.968331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# use cifar10 for pca \ntrfs = transforms.Compose(\n    [\n        transforms.Resize([384, 384]),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n\n    ]\n) \ndataset = torchvision.datasets.CIFAR10(root='.', train=False, download=True, transform=trfs)\nloader = torch.utils.data.DataLoader(dataset, batch_size=32, shuffle=True)\n\nembeds = []\nlabels = []\n\nwith torch.no_grad():\n    for images, labels_ in tqdm(loader):\n        out = model(images.to('cuda'))\n        embeds.append(out)\n        labels.append(labels_)\n\nembeds = torch.cat(embeds)\nlabels = torch.cat(labels)\n\ntorch.save(embeds, 'embeds.pt')\ntorch.save(labels, 'labels.pt')","metadata":{"execution":{"iopub.status.busy":"2022-08-05T22:01:58.526968Z","iopub.execute_input":"2022-08-05T22:01:58.527357Z","iopub.status.idle":"2022-08-05T22:08:48.769665Z","shell.execute_reply.started":"2022-08-05T22:01:58.527298Z","shell.execute_reply":"2022-08-05T22:08:48.768573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Compute PCA on the train embeddings matrix\npca = PCA(n_components=64)\npca.fit(embeds.cpu())\npca_comp = np.asarray(pca.components_).astype(np.float32)\nnp.save('pca.np', pca_comp)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:37:49.319469Z","iopub.execute_input":"2022-08-05T23:37:49.320124Z","iopub.status.idle":"2022-08-05T23:37:51.077885Z","shell.execute_reply.started":"2022-08-05T23:37:49.320087Z","shell.execute_reply":"2022-08-05T23:37:51.076819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"embeds.shape[1]","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:37:54.407060Z","iopub.execute_input":"2022-08-05T23:37:54.407457Z","iopub.status.idle":"2022-08-05T23:37:54.414017Z","shell.execute_reply.started":"2022-08-05T23:37:54.407421Z","shell.execute_reply":"2022-08-05T23:37:54.413013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# linear test \n#l = nn.Linear(embeds.shape[1], 64, bias=False)\n#l.weight = torch.nn.Parameter(torch.tensor(pca_comp))","metadata":{"execution":{"iopub.status.busy":"2022-08-05T16:17:04.244003Z","iopub.execute_input":"2022-08-05T16:17:04.244992Z","iopub.status.idle":"2022-08-05T16:17:04.252561Z","shell.execute_reply.started":"2022-08-05T16:17:04.244947Z","shell.execute_reply":"2022-08-05T16:17:04.251533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyModel(nn.Module):\n  def __init__(self, model_name, target_size=[224, 224], emb_dim=1536, pca_weights=None, normalize=True):\n    super().__init__()\n    self.target_size = target_size\n    \n    self.encoder = timm.create_model(model_name, pretrained=True, num_classes=0)    \n    \n    if pca_weights is None: \n        self.final = nn.AdaptiveAvgPool1d(64)\n    else: \n        self.final =  nn.Linear(emb_dim, 64, bias=False)\n        self.final.weight = torch.nn.Parameter(torch.tensor(pca_weights))\n        \n    self.normalize = normalize\n\n  def forward(self, x):\n    x = transforms.functional.resize(x,size=self.target_size)\n    x = x / 255.0\n    x = transforms.functional.normalize(x, \n                                            mean=[0.485, 0.456, 0.406], \n                                            std=[0.229, 0.224, 0.225])\n    x = self.encoder(x)\n    x = self.final(x)\n    \n    if self.normalize:\n        x = torch.nn.functional.normalize(x)\n    return x","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:37:57.532430Z","iopub.execute_input":"2022-08-05T23:37:57.532814Z","iopub.status.idle":"2022-08-05T23:37:57.544351Z","shell.execute_reply.started":"2022-08-05T23:37:57.532781Z","shell.execute_reply":"2022-08-05T23:37:57.542488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = MyModel(\"swin_large_patch4_window12_384_in22k\", target_size=[384, 384], emb_dim=embeds.shape[1], pca_weights=pca_comp, normalize=True)\nm.eval();","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:38:11.834641Z","iopub.execute_input":"2022-08-05T23:38:11.835011Z","iopub.status.idle":"2022-08-05T23:38:16.395126Z","shell.execute_reply.started":"2022-08-05T23:38:11.834979Z","shell.execute_reply":"2022-08-05T23:38:16.394130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"saved_model = torch.jit.script(m)\nsaved_model.save('saved_model.pt')","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:38:16.397014Z","iopub.execute_input":"2022-08-05T23:38:16.397391Z","iopub.status.idle":"2022-08-05T23:38:19.173061Z","shell.execute_reply.started":"2022-08-05T23:38:16.397361Z","shell.execute_reply":"2022-08-05T23:38:19.171681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"check_model = torch.jit.load('saved_model.pt')\n\nprint(check_model)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:38:19.177067Z","iopub.execute_input":"2022-08-05T23:38:19.177764Z","iopub.status.idle":"2022-08-05T23:38:19.969913Z","shell.execute_reply.started":"2022-08-05T23:38:19.177732Z","shell.execute_reply":"2022-08-05T23:38:19.968673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls -lh","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:38:26.308145Z","iopub.execute_input":"2022-08-05T23:38:26.308566Z","iopub.status.idle":"2022-08-05T23:38:27.412008Z","shell.execute_reply.started":"2022-08-05T23:38:26.308526Z","shell.execute_reply":"2022-08-05T23:38:27.410809Z"},"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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-05T23:38:30.032545Z","iopub.execute_input":"2022-08-05T23:38:30.032972Z","iopub.status.idle":"2022-08-05T23:38:33.028676Z","shell.execute_reply.started":"2022-08-05T23:38:30.032912Z","shell.execute_reply":"2022-08-05T23:38:33.027233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls -lh","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:38:33.031285Z","iopub.execute_input":"2022-08-05T23:38:33.038401Z","iopub.status.idle":"2022-08-05T23:38:34.153725Z","shell.execute_reply.started":"2022-08-05T23:38:33.038350Z","shell.execute_reply":"2022-08-05T23:38:34.152433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchinfo import summary\n\ntest_input_size = (2, 3, 384, 384)  # the model should work with any input_size\nsummary(model, input_size=test_input_size)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"saved_model = torch.jit.load('saved_model.pt').cuda()\ninput = torch.ones(test_input_size, device='cuda')\n\nassert saved_model(input).shape == torch.Size([2, 64])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T23:39:30.747628Z","iopub.execute_input":"2022-08-05T23:39:30.748265Z","iopub.status.idle":"2022-08-05T23:39:32.636613Z","shell.execute_reply.started":"2022-08-05T23:39:30.748228Z","shell.execute_reply":"2022-08-05T23:39:32.635455Z"},"trusted":true},"execution_count":null,"outputs":[]}]}