{"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":"# Setup","metadata":{"id":"DAY5rHgTm7e8"}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')","metadata":{"papermill":{"duration":22.050076,"end_time":"2021-04-17T11:04:51.928845","exception":false,"start_time":"2021-04-17T11:04:29.878769","status":"completed"},"tags":[],"id":"alleged-legislation","executionInfo":{"status":"ok","timestamp":1619310548121,"user_tz":-120,"elapsed":16271,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"outputId":"c6541e5f-ffb4-4609-d6c6-39784e6a07b1","execution":{"iopub.status.busy":"2021-05-26T20:17:47.027374Z","iopub.execute_input":"2021-05-26T20:17:47.027811Z","iopub.status.idle":"2021-05-26T20:17:47.039866Z","shell.execute_reply.started":"2021-05-26T20:17:47.027711Z","shell.execute_reply":"2021-05-26T20:17:47.038836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{"id":"cZoSOL9Qm-Yr"}},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport random\nimport os\nimport math","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.030593,"end_time":"2021-04-17T11:04:51.983376","exception":false,"start_time":"2021-04-17T11:04:51.952783","status":"completed"},"tags":[],"id":"expired-matter","executionInfo":{"status":"ok","timestamp":1619310548121,"user_tz":-120,"elapsed":14459,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:47.042113Z","iopub.execute_input":"2021-05-26T20:17:47.042793Z","iopub.status.idle":"2021-05-26T20:17:47.057712Z","shell.execute_reply.started":"2021-05-26T20:17:47.042749Z","shell.execute_reply":"2021-05-26T20:17:47.056259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score\nfrom sklearn.utils import class_weight\nfrom PIL import Image as pil_image\nfrom tqdm import tqdm\nimport scipy\n\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport plotly.express as px\nimport plotly.graph_objects as go","metadata":{"papermill":{"duration":4.287352,"end_time":"2021-04-17T11:04:56.353165","exception":false,"start_time":"2021-04-17T11:04:52.065813","status":"completed"},"tags":[],"id":"extreme-problem","executionInfo":{"status":"ok","timestamp":1619310550014,"user_tz":-120,"elapsed":16003,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:47.061570Z","iopub.execute_input":"2021-05-26T20:17:47.061950Z","iopub.status.idle":"2021-05-26T20:17:51.145449Z","shell.execute_reply.started":"2021-05-26T20:17:47.061911Z","shell.execute_reply":"2021-05-26T20:17:51.144291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\n\nimport timm\nfrom timm.optim import Lookahead, RAdam","metadata":{"papermill":{"duration":1.641769,"end_time":"2021-04-17T11:04:58.018871","exception":false,"start_time":"2021-04-17T11:04:56.377102","status":"completed"},"tags":[],"id":"angry-domain","executionInfo":{"status":"ok","timestamp":1619310554099,"user_tz":-120,"elapsed":19672,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:51.148422Z","iopub.execute_input":"2021-05-26T20:17:51.149266Z","iopub.status.idle":"2021-05-26T20:17:54.162105Z","shell.execute_reply.started":"2021-05-26T20:17:51.149220Z","shell.execute_reply":"2021-05-26T20:17:54.160981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Global","metadata":{"id":"0B00pe7mnBTj"}},{"cell_type":"code","source":"print(os.listdir(\"/kaggle/input/\"))\nprint(os.listdir(\"/kaggle/input/hotelid-trained-models\"))","metadata":{"execution":{"iopub.status.busy":"2021-05-26T20:17:54.163851Z","iopub.execute_input":"2021-05-26T20:17:54.164314Z","iopub.status.idle":"2021-05-26T20:17:54.183360Z","shell.execute_reply.started":"2021-05-26T20:17:54.164250Z","shell.execute_reply":"2021-05-26T20:17:54.181629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 42\nPROJECT_FOLDER = \"/kaggle/input/hotel-id-2021-fgvc8/\"\nTRAIN_DATA_FOLDER = \"/kaggle/input/hotelid-images-512x512-padded/\"\nTEST_DATA_FOLDER = PROJECT_FOLDER + \"test_images/\"","metadata":{"papermill":{"duration":0.030445,"end_time":"2021-04-17T11:04:58.130825","exception":false,"start_time":"2021-04-17T11:04:58.10038","status":"completed"},"tags":[],"id":"contained-brief","executionInfo":{"status":"ok","timestamp":1619310979015,"user_tz":-120,"elapsed":589,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:54.185328Z","iopub.execute_input":"2021-05-26T20:17:54.185801Z","iopub.status.idle":"2021-05-26T20:17:54.192038Z","shell.execute_reply.started":"2021-05-26T20:17:54.185756Z","shell.execute_reply":"2021-05-26T20:17:54.190296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(os.listdir(PROJECT_FOLDER))","metadata":{"id":"PZvmFng7ctO3","executionInfo":{"status":"ok","timestamp":1619310979515,"user_tz":-120,"elapsed":879,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"outputId":"dce0cc91-8e70-4acc-a0b8-6763ffffd5ca","execution":{"iopub.status.busy":"2021-05-26T20:17:54.194072Z","iopub.execute_input":"2021-05-26T20:17:54.194618Z","iopub.status.idle":"2021-05-26T20:17:54.207112Z","shell.execute_reply.started":"2021-05-26T20:17:54.194571Z","shell.execute_reply":"2021-05-26T20:17:54.205690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions - seed and metric calculator","metadata":{"id":"9p7EE95ZnNpK"}},{"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":{"papermill":{"duration":0.031291,"end_time":"2021-04-17T11:04:58.424933","exception":false,"start_time":"2021-04-17T11:04:58.393642","status":"completed"},"tags":[],"id":"eastern-content","executionInfo":{"status":"ok","timestamp":1619310981653,"user_tz":-120,"elapsed":600,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:54.211808Z","iopub.execute_input":"2021-05-26T20:17:54.212480Z","iopub.status.idle":"2021-05-26T20:17:54.221247Z","shell.execute_reply.started":"2021-05-26T20:17:54.212430Z","shell.execute_reply":"2021-05-26T20:17:54.220051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset and transformations","metadata":{"id":"xaJKvvuKnW4k"}},{"cell_type":"code","source":"import albumentations as A\nimport albumentations.pytorch as APT\nimport cv2 \n\nIMG_SIZE = 512\n\nbase_transform = A.Compose([\n    A.ToFloat(),\n    APT.transforms.ToTensor(),\n])\n\n\ntest_tta_transforms = {\n    \"base\": A.Compose([A.ToFloat(), APT.transforms.ToTensor(),]),\n    \"h_flip\": A.Compose([A.ToFloat(), A.HorizontalFlip(p=1), APT.transforms.ToTensor(),]),\n    \"v_flip\": A.Compose([A.ToFloat(), A.VerticalFlip(p=1), APT.transforms.ToTensor(),]),\n    \"rotate+90\": A.Compose([A.ToFloat(), A.Rotate(limit=90, p=1), APT.transforms.ToTensor(),]),\n    \"rotate-90\": A.Compose([A.ToFloat(), A.Rotate(limit=-90, p=1), APT.transforms.ToTensor(),]),\n#     \"rand_bright\": A.Compose([A.ToFloat(), A.RandomBrightness(p=1), APT.transforms.ToTensor(),]),\n}","metadata":{"papermill":{"duration":0.033385,"end_time":"2021-04-17T11:04:58.538926","exception":false,"start_time":"2021-04-17T11:04:58.505541","status":"completed"},"tags":[],"id":"revolutionary-membership","executionInfo":{"status":"ok","timestamp":1619310984075,"user_tz":-120,"elapsed":1519,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:54.224658Z","iopub.execute_input":"2021-05-26T20:17:54.225039Z","iopub.status.idle":"2021-05-26T20:17:55.182600Z","shell.execute_reply.started":"2021-05-26T20:17:54.225008Z","shell.execute_reply":"2021-05-26T20:17:55.181364Z"},"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\"]\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        \n        return {\n            \"image\" : transformed[\"image\"],\n        }","metadata":{"papermill":{"duration":0.032811,"end_time":"2021-04-17T11:04:58.595928","exception":false,"start_time":"2021-04-17T11:04:58.563117","status":"completed"},"tags":[],"id":"found-mouth","executionInfo":{"status":"ok","timestamp":1619310984077,"user_tz":-120,"elapsed":1058,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:55.185294Z","iopub.execute_input":"2021-05-26T20:17:55.185599Z","iopub.status.idle":"2021-05-26T20:17:55.197934Z","shell.execute_reply.started":"2021-05-26T20:17:55.185568Z","shell.execute_reply":"2021-05-26T20:17:55.196670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"NMDM4PwPnced"}},{"cell_type":"code","source":"# source: https://github.com/ronghuaiyang/arcface-pytorch/blob/master/models/metrics.py\nclass 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.fc':\n            self.backbone.head.fc = nn.Identity()\n        elif fc_name == 'fc':\n            self.backbone.fc = 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":{"id":"GuAfw_a4m3PK","executionInfo":{"status":"ok","timestamp":1619310987166,"user_tz":-120,"elapsed":578,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:55.199664Z","iopub.execute_input":"2021-05-26T20:17:55.200596Z","iopub.status.idle":"2021-05-26T20:17:55.227252Z","shell.execute_reply.started":"2021-05-26T20:17:55.200538Z","shell.execute_reply":"2021-05-26T20:17:55.226009Z"},"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.fc':\n            self.backbone.head.fc = nn.Identity()\n        elif fc_name == 'fc':\n            self.backbone.fc = 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":"2021-05-26T20:17:55.229170Z","iopub.execute_input":"2021-05-26T20:17:55.229691Z","iopub.status.idle":"2021-05-26T20:17:55.246351Z","shell.execute_reply.started":"2021-05-26T20:17:55.229643Z","shell.execute_reply":"2021-05-26T20:17:55.245176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model helper functions","metadata":{"id":"YMZYKhUSneMY"}},{"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":"2021-05-26T20:17:55.248354Z","iopub.execute_input":"2021-05-26T20:17:55.249131Z","iopub.status.idle":"2021-05-26T20:17:55.261464Z","shell.execute_reply.started":"2021-05-26T20:17:55.248984Z","shell.execute_reply":"2021-05-26T20:17:55.260186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def predict(loader, base_df, base_embeds, model, n_matches=5, bar_desc=\"Generating embeds\"):\n#     preds = []\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#             distances = cosine_similarity(output.detach().cpu().numpy(), base_embeds)\n            \n#             for j in range(0, output.shape[0]):\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\n# def find_closest_match_tta(args, test_df, tta_transforms, base_loader, model, n_matches=5):\n#     base_embeds = get_embeds(base_loader, model, \"Generating embeds for train\")\n\n#     test_dataset = HotelImageDataset(test_df, tta_transforms[\"base\"], data_folder=TEST_DATA_FOLDER)\n#     test_loader = DataLoader(test_dataset, num_workers=args.num_workers, batch_size=args.batch_size, shuffle=False)\n#     preds = predict(test_loader, base_loader.dataset.data, base_embeds, model, n_matches, f\"Generating predictions\")\n        \n#     return preds","metadata":{"papermill":{"duration":0.032565,"end_time":"2021-04-17T11:04:58.652672","exception":false,"start_time":"2021-04-17T11:04:58.620107","status":"completed"},"tags":[],"id":"massive-makeup","executionInfo":{"status":"ok","timestamp":1619310991310,"user_tz":-120,"elapsed":452,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:55.264939Z","iopub.execute_input":"2021-05-26T20:17:55.265403Z","iopub.status.idle":"2021-05-26T20:17:55.276301Z","shell.execute_reply.started":"2021-05-26T20:17:55.265348Z","shell.execute_reply":"2021-05-26T20:17:55.274846Z"},"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":"2021-05-26T20:17:55.278281Z","iopub.execute_input":"2021-05-26T20:17:55.278778Z","iopub.status.idle":"2021-05-26T20:17:55.293631Z","shell.execute_reply.started":"2021-05-26T20:17:55.278733Z","shell.execute_reply":"2021-05-26T20:17:55.292268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare data","metadata":{"id":"AwShW1wXniD6"}},{"cell_type":"markdown","source":"### for validation","metadata":{}},{"cell_type":"code","source":"# data_df = pd.read_csv(PROJECT_FOLDER + \"train.csv\", parse_dates=[\"timestamp\"])\n# data_df = data_df[data_df[\"hotel_id\"].isin(data_df[\"hotel_id\"].unique()[-500:])]\n# test_df = data_df.groupby(\"hotel_id\").sample(1, random_state=SEED)\n# data_df = data_df[~data_df[\"image\"].isin(test_df[\"image\"])]\n\n# TEST_DATA_FOLDER = TRAIN_DATA_FOLDER\n\n# print(f\"Base: {len(data_df)}, test: {len(test_df)}\")","metadata":{"execution":{"iopub.status.busy":"2021-05-26T20:17:55.295414Z","iopub.execute_input":"2021-05-26T20:17:55.295948Z","iopub.status.idle":"2021-05-26T20:17:56.060795Z","shell.execute_reply.started":"2021-05-26T20:17:55.295904Z","shell.execute_reply":"2021-05-26T20:17:56.059636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### for submission","metadata":{}},{"cell_type":"code","source":"data_df = pd.read_csv(PROJECT_FOLDER + \"train.csv\", parse_dates=[\"timestamp\"])\nsample_submission_df = pd.read_csv(PROJECT_FOLDER + \"sample_submission.csv\")\ntest_df = pd.DataFrame(data={\"image\": os.listdir(TEST_DATA_FOLDER), \"hotel_id\": \"\"}).sort_values(by=\"image\")","metadata":{"papermill":{"duration":2.790179,"end_time":"2021-04-17T11:05:01.702988","exception":false,"start_time":"2021-04-17T11:04:58.912809","status":"completed"},"tags":[],"id":"discrete-right","executionInfo":{"status":"ok","timestamp":1619311036476,"user_tz":-120,"elapsed":3742,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"outputId":"c21ed589-3139-4919-b5d5-07bcf6f1df15","execution":{"iopub.status.busy":"2021-05-26T20:17:56.062413Z","iopub.execute_input":"2021-05-26T20:17:56.062912Z","iopub.status.idle":"2021-05-26T20:17:56.067542Z","shell.execute_reply.started":"2021-05-26T20:17:56.062863Z","shell.execute_reply":"2021-05-26T20:17:56.066199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train and evaluate","metadata":{"id":"5JPdD2bpnniP"}},{"cell_type":"code","source":"def get_model(model_type, backbone_name, embed_size, checkpoint_path, args):\n    if model_type == 'arcmargin':\n        model = HotelIdModel(7770, embed_size, backbone_name)\n    else:\n        model = EmbeddingNet(7770, 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":"2021-05-26T20:17:56.069216Z","iopub.execute_input":"2021-05-26T20:17:56.070024Z","iopub.status.idle":"2021-05-26T20:17:56.085405Z","shell.execute_reply.started":"2021-05-26T20:17:56.069964Z","shell.execute_reply":"2021-05-26T20:17:56.084271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class args:\n    batch_size = 32\n    num_workers = 4\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":{"papermill":{"duration":0.59707,"end_time":"2021-04-17T11:05:02.330381","exception":false,"start_time":"2021-04-17T11:05:01.733311","status":"completed"},"tags":[],"id":"appointed-machinery","executionInfo":{"status":"ok","timestamp":1619311064188,"user_tz":-120,"elapsed":450,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:17:56.088860Z","iopub.execute_input":"2021-05-26T20:17:56.089232Z","iopub.status.idle":"2021-05-26T20:17:56.160111Z","shell.execute_reply.started":"2021-05-26T20:17:56.089189Z","shell.execute_reply":"2021-05-26T20:17:56.158835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_array = [get_model(\"arcmargin\", \n                         \"efficientnet_b1\", 4096,\n                         \"../input/hotelid-trained-models/checkpoint-arcmargin-model-efficientnet_b1-512x512-4096embeds-7770hotels.pt\", \n                         args),\n               \n               get_model(\"cosface\", \n                         \"ecaresnet50d_pruned\", 4096,\n                         \"../input/hotelid-trained-models/checkpoint-cosface-model-ecaresnet50d_pruned-512x512-4096embeds-7770hotels.pt\", \n                         args),\n               \n               get_model(\"classification\", \n                         \"eca_nfnet_l0\", 4096,\n                         \"../input/hotelid-trained-models/checkpoint-classification-model-eca_nfnet_l0-512x512-4096embeds-7770hotels.pt\", \n                         args),\n               \n               get_model(\"arcmargin\", \n                         \"eca_nfnet_l0\", 1024,\n                         \"../input/hotelid-trained-models/checkpoint-arcmargin-model-eca_nfnet_l0-512x512-1024embeds-7770hotels.pt\", \n                         args),\n              ]","metadata":{"execution":{"iopub.status.busy":"2021-05-26T20:17:56.161990Z","iopub.execute_input":"2021-05-26T20:17:56.162427Z","iopub.status.idle":"2021-05-26T20:19:19.225561Z","shell.execute_reply.started":"2021-05-26T20:17:56.162353Z","shell.execute_reply":"2021-05-26T20:19:19.224337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### submission","metadata":{}},{"cell_type":"code","source":"%%time\n\nif len(test_df) > 3:\n    preds = find_closest_match(args, test_loader, base_loader, model_array, n_matches=5)\n    test_df[\"hotel_id\"] = [str(list(l)).strip(\"[]\").replace(\",\", \"\") for l in preds]\n\ntest_df.to_csv(\"submission.csv\", index=False)\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2021-05-26T20:19:19.227419Z","iopub.execute_input":"2021-05-26T20:19:19.227827Z","iopub.status.idle":"2021-05-26T20:19:19.236000Z","shell.execute_reply.started":"2021-05-26T20:19:19.227781Z","shell.execute_reply":"2021-05-26T20:19:19.232513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### validation","metadata":{}},{"cell_type":"code","source":"# preds = find_closest_match(args, test_loader, base_loader, model_array, n_matches=5)\n\n# test_df[\"hotel_id_pred\"] = [str(list(l)).strip(\"[]\").replace(\",\", \"\") for l in preds]\n\n# y = np.repeat([test_df[\"hotel_id\"]], repeats=5, axis=0).T\n# preds = np.array(preds)\n\n# acc_top_1 = (preds[:, 0] == test_df[\"hotel_id\"]).mean()\n# acc_top_5 = (preds == y).any(axis=1).mean()\n\n# print(f\"Accuracy: {acc_top_1:0.4f}, top 5 accuracy: {acc_top_5:0.4f}\")","metadata":{"papermill":{"duration":10.500513,"end_time":"2021-04-17T14:14:58.022931","exception":false,"start_time":"2021-04-17T14:14:47.522418","status":"completed"},"tags":[],"id":"outer-company","executionInfo":{"status":"ok","timestamp":1619295476782,"user_tz":-120,"elapsed":931536,"user":{"displayName":"Jeom Jin-Ho","photoUrl":"","userId":"00155613517919499503"}},"execution":{"iopub.status.busy":"2021-05-26T20:20:40.118044Z","iopub.execute_input":"2021-05-26T20:20:40.118407Z","iopub.status.idle":"2021-05-26T20:23:59.112391Z","shell.execute_reply.started":"2021-05-26T20:20:40.118376Z","shell.execute_reply":"2021-05-26T20:23:59.111009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}