{"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":"# Intro\nInference notebook for [Hotel-ID starter - similarity - training](https://www.kaggle.com/code/michaln/hotel-id-starter-similarity-training)\n\nUsing model and embeddings from the training notebook to generate embeddings for test data and find similar images.","metadata":{"id":"DAY5rHgTm7e8","papermill":{"duration":0.024896,"end_time":"2022-03-24T14:00:54.588459","exception":false,"start_time":"2022-03-24T14:00:54.563563","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Setup","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/timm-062dev0/pytorch-image-models-master')","metadata":{"executionInfo":{"elapsed":16271,"status":"ok","timestamp":1619310548121,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"},"user_tz":-120},"id":"alleged-legislation","outputId":"c6541e5f-ffb4-4609-d6c6-39784e6a07b1","papermill":{"duration":0.036572,"end_time":"2022-03-24T14:00:54.649254","exception":false,"start_time":"2022-03-24T14:00:54.612682","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:04:28.584124Z","iopub.execute_input":"2022-05-25T23:04:28.58441Z","iopub.status.idle":"2022-05-25T23:04:28.615466Z","shell.execute_reply.started":"2022-05-25T23:04:28.58432Z","shell.execute_reply":"2022-05-25T23:04:28.614842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{"id":"cZoSOL9Qm-Yr","papermill":{"duration":0.023644,"end_time":"2022-03-24T14:00:54.696898","exception":false,"start_time":"2022-03-24T14:00:54.673254","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport random\nimport os\nimport math","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","executionInfo":{"elapsed":14459,"status":"ok","timestamp":1619310548121,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"},"user_tz":-120},"id":"expired-matter","papermill":{"duration":0.030271,"end_time":"2022-03-24T14:00:54.751131","exception":false,"start_time":"2022-03-24T14:00:54.72086","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:05:25.921307Z","iopub.execute_input":"2022-05-25T23:05:25.921747Z","iopub.status.idle":"2022-05-25T23:05:25.925367Z","shell.execute_reply.started":"2022-05-25T23:05:25.92171Z","shell.execute_reply":"2022-05-25T23:05:25.924654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image as pil_image\nfrom tqdm import tqdm","metadata":{"executionInfo":{"elapsed":16003,"status":"ok","timestamp":1619310550014,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"},"user_tz":-120},"id":"extreme-problem","papermill":{"duration":3.220402,"end_time":"2022-03-24T14:00:57.995239","exception":false,"start_time":"2022-03-24T14:00:54.774837","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:05:26.829676Z","iopub.execute_input":"2022-05-25T23:05:26.831424Z","iopub.status.idle":"2022-05-25T23:05:26.835734Z","shell.execute_reply.started":"2022-05-25T23:05:26.831376Z","shell.execute_reply":"2022-05-25T23:05:26.835074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\n\nimport timm\nfrom sklearn.metrics.pairwise import cosine_similarity","metadata":{"executionInfo":{"elapsed":19672,"status":"ok","timestamp":1619310554099,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"},"user_tz":-120},"id":"angry-domain","papermill":{"duration":2.727834,"end_time":"2022-03-24T14:01:00.766951","exception":false,"start_time":"2022-03-24T14:00:58.039117","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:05:27.703507Z","iopub.execute_input":"2022-05-25T23:05:27.704034Z","iopub.status.idle":"2022-05-25T23:05:31.867361Z","shell.execute_reply.started":"2022-05-25T23:05:27.703996Z","shell.execute_reply":"2022-05-25T23:05:31.866014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Global","metadata":{"id":"0B00pe7mnBTj","papermill":{"duration":0.023573,"end_time":"2022-03-24T14:01:00.814976","exception":false,"start_time":"2022-03-24T14:01:00.791403","status":"completed"},"tags":[]}},{"cell_type":"code","source":"print(os.listdir(\"/kaggle/input/\"))\nprint(os.listdir(\"/kaggle/input/checkpoint-arcmargin-1\"))","metadata":{"executionInfo":{"elapsed":589,"status":"ok","timestamp":1619310979015,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"},"user_tz":-120},"id":"contained-brief","papermill":{"duration":0.03175,"end_time":"2022-03-24T14:01:00.871686","exception":false,"start_time":"2022-03-24T14:01:00.839936","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:05:35.44524Z","iopub.execute_input":"2022-05-25T23:05:35.44595Z","iopub.status.idle":"2022-05-25T23:05:35.456067Z","shell.execute_reply.started":"2022-05-25T23:05:35.445911Z","shell.execute_reply":"2022-05-25T23:05:35.455027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 42\nPROJECT_FOLDER = \"../input/hotel-id-to-combat-human-trafficking-2022-fgvc9/\"\nDATA_FOLDER = \"../input/nopadding256/\"\nTRAIN_DATA_FOLDER = DATA_FOLDER + \"images/\"\nTEST_DATA_FOLDER = PROJECT_FOLDER + \"test_images/\"","metadata":{"executionInfo":{"elapsed":879,"status":"ok","timestamp":1619310979515,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"},"user_tz":-120},"id":"PZvmFng7ctO3","outputId":"dce0cc91-8e70-4acc-a0b8-6763ffffd5ca","papermill":{"duration":0.031651,"end_time":"2022-03-24T14:01:00.927239","exception":false,"start_time":"2022-03-24T14:01:00.895588","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:05:50.880901Z","iopub.execute_input":"2022-05-25T23:05:50.881153Z","iopub.status.idle":"2022-05-25T23:05:50.887758Z","shell.execute_reply.started":"2022-05-25T23:05:50.881124Z","shell.execute_reply":"2022-05-25T23:05:50.88706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(os.listdir(PROJECT_FOLDER))","metadata":{"executionInfo":{"elapsed":600,"status":"ok","timestamp":1619310981653,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"},"user_tz":-120},"id":"eastern-content","papermill":{"duration":0.032105,"end_time":"2022-03-24T14:01:01.031949","exception":false,"start_time":"2022-03-24T14:01:00.999844","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:05:51.919368Z","iopub.execute_input":"2022-05-25T23:05:51.91988Z","iopub.status.idle":"2022-05-25T23:05:51.925379Z","shell.execute_reply.started":"2022-05-25T23:05:51.919844Z","shell.execute_reply":"2022-05-25T23:05:51.92449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions - seed and metric calculator","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:05:56.981058Z","iopub.execute_input":"2022-05-25T23:05:56.981592Z","iopub.status.idle":"2022-05-25T23:05:56.986304Z","shell.execute_reply.started":"2022-05-25T23:05:56.981552Z","shell.execute_reply":"2022-05-25T23:05:56.985607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nimport albumentations.pytorch as APT\nimport cv2 \n\nIMG_SIZE = 256\n\nbase_transform = A.Compose([\n    A.ToFloat(),\n    APT.transforms.ToTensorV2(),\n])\n\n\ntest_tta_transforms = {\n    \"base\": A.Compose([A.ToFloat(), APT.transforms.ToTensorV2(),]),\n    \"h_flip\": A.Compose([A.ToFloat(), A.HorizontalFlip(p=1), APT.transforms.ToTensorV2(),]),\n    \"v_flip\": A.Compose([A.ToFloat(), A.VerticalFlip(p=1), APT.transforms.ToTensorV2(),]),\n    \"rotate+90\": A.Compose([A.ToFloat(), A.Rotate(limit=90, p=1), APT.transforms.ToTensorV2(),]),\n    \"rotate-90\": A.Compose([A.ToFloat(), A.Rotate(limit=-90, p=1), APT.transforms.ToTensorV2(),]),\n#     \"rand_bright\": A.Compose([A.ToFloat(), A.RandomBrightness(p=1), APT.transforms.ToTensor(),]),\n}","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:05:58.177868Z","iopub.execute_input":"2022-05-25T23:05:58.17858Z","iopub.status.idle":"2022-05-25T23:05:59.398501Z","shell.execute_reply.started":"2022-05-25T23:05:58.17854Z","shell.execute_reply":"2022-05-25T23:05:59.397748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pad_image(img):\n    w, h, c = np.shape(img)\n    if w > h:\n        pad = int((w - h) / 2)\n        img = cv2.copyMakeBorder(img, 0, 0, pad, pad, cv2.BORDER_CONSTANT, value=0)\n    else:\n        pad = int((h - w) / 2)\n        img = cv2.copyMakeBorder(img, pad, pad, 0, 0, cv2.BORDER_CONSTANT, value=0)\n        \n    return img\n\n\ndef open_and_preprocess_image(image_path):\n    img = cv2.imread(image_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img = pad_image(img)\n    return cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n\n\nclass HotelImageDataset:\n    def __init__(self, data, transform=None, data_folder=\"train_images/\"):\n        self.data = data\n        self.data_folder = data_folder\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        record = self.data.iloc[idx]\n        image_path = self.data_folder + record[\"image_id\"]\n        \n        if \"test\" in self.data_folder:\n            image = np.array(open_and_preprocess_image(image_path)).astype(np.uint8)\n        else:\n            image = np.array(pil_image.open(image_path)).astype(np.uint8)\n\n        if self.transform:\n            transformed = self.transform(image=image)\n            image = transformed[\"image\"]\n        \n        return {\n            \"image\" : image,\n        }","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:05:59.533558Z","iopub.execute_input":"2022-05-25T23:05:59.533972Z","iopub.status.idle":"2022-05-25T23:05:59.544498Z","shell.execute_reply.started":"2022-05-25T23:05:59.533941Z","shell.execute_reply":"2022-05-25T23:05:59.543846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"NMDM4PwPnced","papermill":{"duration":0.023902,"end_time":"2022-03-24T14:01:01.962307","exception":false,"start_time":"2022-03-24T14:01:01.938405","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class ArcMarginProduct(nn.Module):\n    r\"\"\"Implement of large margin arc distance: :\n        Args:\n            in_features: size of each input sample\n            out_features: size of each output sample\n            s: norm of input feature\n            m: margin\n            cos(theta + m)\n        \"\"\"\n    def __init__(self, in_features, out_features, s=30.0, m=0.50, easy_margin=False):\n        super(ArcMarginProduct, self).__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = s\n        self.m = m\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.easy_margin = easy_margin\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n\n    def forward(self, input, label):\n        # --------------------------- cos(theta) & phi(theta) ---------------------------\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        sine = torch.sqrt((1.0 - torch.pow(cosine, 2)).clamp(0, 1))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        if self.easy_margin:\n            phi = torch.where(cosine > 0, phi, cosine)\n        else:\n            phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        # --------------------------- convert label to one-hot ---------------------------\n        # one_hot = torch.zeros(cosine.size(), requires_grad=True, device='cuda')\n        one_hot = torch.zeros(cosine.size(), device='cuda')\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        # -------------torch.where(out_i = {x_i if condition_i else y_i) -------------\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)  # you can use torch.where if your torch.__version__ is 0.4\n        output *= self.s\n\n        return output\n\nclass HotelIdModel(nn.Module):\n    def __init__(self, out_features, embed_size=256, backbone_name=\"efficientnet_b3\"):\n        super(HotelIdModel, self).__init__()\n\n        self.embed_size = embed_size\n        self.backbone = timm.create_model(backbone_name, pretrained=False)\n        in_features = self.backbone.get_classifier().in_features\n\n        fc_name, _ = list(self.backbone.named_modules())[-1]\n        if fc_name == 'classifier':\n            self.backbone.classifier = nn.Identity()\n        elif fc_name == 'head.flatten':\n            self.backbone.head.fc = nn.Identity()\n        elif fc_name == 'fc':\n            self.backbone.fc = nn.Identity()\n        elif fc_name == 'head':\n            self.backbone.head = nn.Identity()\n        else:\n            raise Exception(\"unknown classifier layer: \" + fc_name)\n\n        self.arc_face = ArcMarginProduct(self.embed_size, out_features, s=30.0, m=0.50, easy_margin=False)\n\n        self.post = nn.Sequential(\n            nn.utils.weight_norm(nn.Linear(in_features, self.embed_size*2), dim=None),\n            nn.BatchNorm1d(self.embed_size*2),\n            nn.Dropout(0.2),\n            nn.utils.weight_norm(nn.Linear(self.embed_size*2, self.embed_size)),\n            nn.BatchNorm1d(self.embed_size),\n        )\n\n        print(f\"Model {backbone_name} ArcMarginProduct - Features: {in_features}, Embeds: {self.embed_size}\")\n        \n    def forward(self, input, targets = None):\n        x = self.backbone(input)\n        x = x.view(x.size(0), -1)\n        x = self.post(x)\n        \n        if targets is not None:\n            logits = self.arc_face(x, targets)\n            return logits\n        \n        return x","metadata":{"papermill":{"duration":0.032166,"end_time":"2022-03-24T14:01:02.018479","exception":false,"start_time":"2022-03-24T14:01:01.986313","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:06:38.09507Z","iopub.execute_input":"2022-05-25T23:06:38.095369Z","iopub.status.idle":"2022-05-25T23:06:38.113814Z","shell.execute_reply.started":"2022-05-25T23:06:38.095335Z","shell.execute_reply":"2022-05-25T23:06:38.11314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EmbeddingNet(nn.Module):\n    def __init__(self, n_classes=100, embed_size=64, backbone_name=\"efficientnet_b0\"):\n        super(EmbeddingNet, self).__init__()\n        \n        self.embed_size = embed_size   \n        self.backbone = timm.create_model(backbone_name, pretrained=False)\n        in_features = self.backbone.get_classifier().in_features\n\n        fc_name, _ = list(self.backbone.named_modules())[-1]\n        if fc_name == 'classifier':\n            self.backbone.classifier = nn.Identity()\n        elif fc_name == 'head.flatten':\n            self.backbone.head.fc = nn.Identity()\n        elif fc_name == 'fc':\n            self.backbone.fc = nn.Identity()\n        elif fc_name == 'head':\n            self.backbone.head = nn.Identity()\n        else:\n            raise Exception(\"unknown classifier layer: \" + fc_name)\n        \n        self.post = nn.Sequential(\n            nn.utils.weight_norm(nn.Linear(in_features, self.embed_size*2), dim=None),\n            nn.BatchNorm1d(self.embed_size*2),\n            nn.Dropout(0.2),\n            nn.utils.weight_norm(nn.Linear(self.embed_size*2, self.embed_size)),\n        )\n\n        self.classifier = nn.Sequential(\n            nn.BatchNorm1d(self.embed_size),\n            nn.Dropout(0.2),\n            nn.Linear(self.embed_size, n_classes),\n        )\n        \n        print(f\"Model {backbone_name} EmbeddingNet - Features: {in_features}, Embeds: {self.embed_size}\")\n        \n    def embed_and_classify(self, x):\n        x = self.forward(x)\n        return x, self.classifier(x)\n\n    def forward(self, x):\n        x = self.backbone(x)\n        x = x.view(x.size(0), -1)\n        x = self.post(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:06:51.993093Z","iopub.execute_input":"2022-05-25T23:06:51.993369Z","iopub.status.idle":"2022-05-25T23:06:52.006754Z","shell.execute_reply.started":"2022-05-25T23:06:51.993337Z","shell.execute_reply":"2022-05-25T23:06:52.005193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics.pairwise import cosine_similarity\n\ndef get_embeds(loader, model, bar_desc=\"Generating embeds\"):\n    outputs_all = []\n    \n    model.eval()\n    with torch.no_grad():\n        t = tqdm(loader, desc=bar_desc)\n        for i, sample in enumerate(t):\n            input = sample['image'].to(args.device)\n            output = model(input)\n            outputs_all.extend(output.detach().cpu().numpy())\n#             outputs_all.extend(output.detach().cpu().numpy().astype(np.float16))\n            \n            \n    return outputs_all","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:06:56.540921Z","iopub.execute_input":"2022-05-25T23:06:56.54118Z","iopub.status.idle":"2022-05-25T23:06:56.547796Z","shell.execute_reply.started":"2022-05-25T23:06:56.54115Z","shell.execute_reply":"2022-05-25T23:06:56.546821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_distances(input, base_embeds, model_array):\n    distances = None\n    for i, model in enumerate(model_array):\n        output = model(input)\n        output = output.detach().cpu().numpy()\n#         output = output.detach().cpu().numpy().astype(np.float16)\n        model_base_embeds = base_embeds[i]\n        output_distances = cosine_similarity(output, model_base_embeds)\n        \n        if distances is None:\n            distances = output_distances\n        else:\n            distances = distances * output_distances\n            \n    return distances\n    \n\ndef predict(loader, base_df, base_embeds, model_array, n_matches=5, bar_desc=\"Generating embeds\"):\n    preds = []\n    with torch.no_grad():\n        t = tqdm(loader, desc=bar_desc)\n        for i, sample in enumerate(t):\n            input = sample['image'].to(args.device)\n            distances = get_distances(input, base_embeds, model_array)\n            \n            for j in range(len(distances)):\n                tmp_df = base_df.copy()\n                tmp_df[\"distance\"] = distances[j]\n                tmp_df = tmp_df.sort_values(by=[\"distance\", \"hotel_id\"], ascending=False).reset_index(drop=True)\n                preds.extend([tmp_df[\"hotel_id\"].unique()[:n_matches]])\n\n    return preds\n\ndef find_closest_match(args, test_loader, base_loader, model_array, n_matches=5):\n    base_embeds = {}\n    for i, model in enumerate(model_array):\n        base_embeds[i] = get_embeds(base_loader, model, \"Generating embeds for train\")\n    \n    preds = predict(test_loader, base_loader.dataset.data, base_embeds, model_array, n_matches, f\"Generating predictions\")\n        \n    return preds","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:06:57.741995Z","iopub.execute_input":"2022-05-25T23:06:57.742851Z","iopub.status.idle":"2022-05-25T23:06:57.754833Z","shell.execute_reply.started":"2022-05-25T23:06:57.742797Z","shell.execute_reply":"2022-05-25T23:06:57.753419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_df = pd.read_csv(DATA_FOLDER + \"train_no_padding_256.csv\")\n# encode hotel ids\ndata_df[\"hotel_id_code\"] = data_df[\"hotel_id\"].astype('category').cat.codes.values.astype(np.int64)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:06:59.467421Z","iopub.execute_input":"2022-05-25T23:06:59.468118Z","iopub.status.idle":"2022-05-25T23:06:59.523817Z","shell.execute_reply.started":"2022-05-25T23:06:59.468074Z","shell.execute_reply":"2022-05-25T23:06:59.523088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_df","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:07:00.509903Z","iopub.execute_input":"2022-05-25T23:07:00.510473Z","iopub.status.idle":"2022-05-25T23:07:00.527804Z","shell.execute_reply.started":"2022-05-25T23:07:00.510432Z","shell.execute_reply":"2022-05-25T23:07:00.527137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission_df = pd.read_csv(PROJECT_FOLDER + \"sample_submission.csv\")\ntest_df = pd.DataFrame(data={\"image_id\": os.listdir(TEST_DATA_FOLDER), \"hotel_id\": \"\"}).sort_values(by=\"image_id\")","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:07:02.255864Z","iopub.execute_input":"2022-05-25T23:07:02.256712Z","iopub.status.idle":"2022-05-25T23:07:02.274075Z","shell.execute_reply.started":"2022-05-25T23:07:02.256673Z","shell.execute_reply":"2022-05-25T23:07:02.273393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:07:03.545225Z","iopub.execute_input":"2022-05-25T23:07:03.545764Z","iopub.status.idle":"2022-05-25T23:07:03.553665Z","shell.execute_reply.started":"2022-05-25T23:07:03.545728Z","shell.execute_reply":"2022-05-25T23:07:03.552865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train and evaluate","metadata":{}},{"cell_type":"code","source":"def get_model(model_type, backbone_name, embed_size, checkpoint_path, args):\n    if model_type == 'arcmargin':\n        model = HotelIdModel(3116, embed_size, backbone_name)\n    else:\n        model = EmbeddingNet(3116, embed_size, backbone_name)\n        \n    checkpoint = torch.load(checkpoint_path)\n    model.load_state_dict(checkpoint[\"model\"])\n    model = model.to(args.device)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:07:13.263816Z","iopub.execute_input":"2022-05-25T23:07:13.264072Z","iopub.status.idle":"2022-05-25T23:07:13.269313Z","shell.execute_reply.started":"2022-05-25T23:07:13.264042Z","shell.execute_reply":"2022-05-25T23:07:13.268486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class args:\n    batch_size = 64\n    num_workers = 2\n    n_classes = data_df[\"hotel_id\"].nunique()\n    device = ('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    \nseed_everything(seed=SEED)\n\nbase_dataset = HotelImageDataset(data_df, base_transform, data_folder=TRAIN_DATA_FOLDER)\nbase_loader = DataLoader(base_dataset, num_workers=args.num_workers, batch_size=args.batch_size, shuffle=False)\n\ntest_dataset = HotelImageDataset(test_df, test_tta_transforms[\"base\"], data_folder=TEST_DATA_FOLDER)\ntest_loader = DataLoader(test_dataset, num_workers=args.num_workers, batch_size=args.batch_size, shuffle=False)","metadata":{"executionInfo":{"elapsed":450,"status":"ok","timestamp":1619311064188,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"},"user_tz":-120},"id":"appointed-machinery","papermill":{"duration":0.069839,"end_time":"2022-03-24T14:01:02.65177","exception":false,"start_time":"2022-03-24T14:01:02.581931","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:07:24.922665Z","iopub.execute_input":"2022-05-25T23:07:24.922925Z","iopub.status.idle":"2022-05-25T23:07:24.990068Z","shell.execute_reply.started":"2022-05-25T23:07:24.922897Z","shell.execute_reply":"2022-05-25T23:07:24.989338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_array = [get_model(\"arcmargin\", \n                         \"eca_nfnet_l0\", 2048,\n                         \"../input/checkpoint-arcmargin-1/checkpoint-arcmargin-model-eca_nfnet_l0-256x256-2048embeds-3116hotels.pt\", \n                         args),\n\n               get_model(\"arcmargin\", \n                         \"efficientnet_b1\", 4096,\n                         \"../input/just-for-test/checkpoint-arcmargin-model-efficientnet_b1-256x256-4096embeds-3116hotels.pt\", \n                         args),\n               \n               get_model(\"arcmargin\", \n                         \"resnest101e\", 2048,\n                         \"../input/checkpoint-arcmargin-2/checkpoint-arcmargin-model-resnest101e-256x256-2048embeds-3116hotels.pt\", \n                         args),\n\n               get_model(\"arcmargin\", \n                         \"swinv2_base_window16_256\", 2048,\n                         \"../input/checkpoint-arcmargin-4/checkpoint-arcmargin-model-swinv2_base_window16_256-256x256-2048embeds-3116hotels.pt\", \n                         args),\n               \n               get_model(\"classification\", \n                         \"eca_nfnet_l0\", 2048,\n                         \"../input/checkpoint-classification-1/checkpoint-classification-model-eca_nfnet_l0-256x256-2048embeds-3116hotels.pt\", \n                         args),\n               \n               get_model(\"cosface\", \n                         \"eca_nfnet_l0\", 4096,\n                         \"../input/checkpoint-cosface/checkpoint-cosface-model-eca_nfnet_l0-256x256-4096embeds-3116hotels.pt\", \n                         args),\n\n              ]","metadata":{"execution":{"iopub.status.busy":"2022-05-25T23:12:57.362159Z","iopub.execute_input":"2022-05-25T23:12:57.362686Z","iopub.status.idle":"2022-05-25T23:13:53.645369Z","shell.execute_reply.started":"2022-05-25T23:12:57.36265Z","shell.execute_reply":"2022-05-25T23:13:53.644619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\npreds = find_closest_match(args, test_loader, base_loader, model_array, n_matches=5)\ntest_df[\"hotel_id\"] = [str(list(l)).strip(\"[]\").replace(\",\", \"\") for l in preds]","metadata":{"papermill":{"duration":5.553999,"end_time":"2022-03-24T14:01:08.229948","exception":false,"start_time":"2022-03-24T14:01:02.675949","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:14:04.208811Z","iopub.execute_input":"2022-05-25T23:14:04.209063Z","iopub.status.idle":"2022-05-25T23:38:40.021558Z","shell.execute_reply.started":"2022-05-25T23:14:04.209035Z","shell.execute_reply":"2022-05-25T23:38:40.02074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{"papermill":{"duration":0.024227,"end_time":"2022-03-24T14:01:08.278753","exception":false,"start_time":"2022-03-24T14:01:08.254526","status":"completed"},"tags":[]}},{"cell_type":"code","source":"test_df","metadata":{"papermill":{"duration":0.047156,"end_time":"2022-03-24T14:01:08.414557","exception":false,"start_time":"2022-03-24T14:01:08.367401","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-25T23:39:24.898807Z","iopub.execute_input":"2022-05-25T23:39:24.899094Z","iopub.status.idle":"2022-05-25T23:39:24.90834Z","shell.execute_reply.started":"2022-05-25T23:39:24.899058Z","shell.execute_reply":"2022-05-25T23:39:24.907356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.to_csv(\"submission.csv\", index=False)","metadata":{},"execution_count":null,"outputs":[]}]}