{"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/timmmaster') # Newer timm verson to support swinv2","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":"2022-05-29T08:12:46.577127Z","iopub.execute_input":"2022-05-29T08:12:46.577507Z","iopub.status.idle":"2022-05-29T08:12:46.586693Z","shell.execute_reply.started":"2022-05-29T08:12:46.577412Z","shell.execute_reply":"2022-05-29T08:12:46.585459Z"},"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":"2022-05-29T08:12:46.592877Z","iopub.execute_input":"2022-05-29T08:12:46.593143Z","iopub.status.idle":"2022-05-29T08:12:46.597438Z","shell.execute_reply.started":"2022-05-29T08:12:46.593119Z","shell.execute_reply":"2022-05-29T08:12:46.596633Z"},"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":"2022-05-29T08:12:46.599063Z","iopub.execute_input":"2022-05-29T08:12:46.599672Z","iopub.status.idle":"2022-05-29T08:12:49.716507Z","shell.execute_reply.started":"2022-05-29T08:12:46.599636Z","shell.execute_reply":"2022-05-29T08:12:49.715691Z"},"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":"2022-05-29T08:12:49.719807Z","iopub.execute_input":"2022-05-29T08:12:49.720061Z","iopub.status.idle":"2022-05-29T08:12:53.067351Z","shell.execute_reply.started":"2022-05-29T08:12:49.720037Z","shell.execute_reply":"2022-05-29T08:12:53.066224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Global","metadata":{"id":"0B00pe7mnBTj"}},{"cell_type":"code","source":"print(os.listdir(\"/kaggle/input/\"))","metadata":{"execution":{"iopub.status.busy":"2022-05-29T08:12:53.072179Z","iopub.execute_input":"2022-05-29T08:12:53.072559Z","iopub.status.idle":"2022-05-29T08:12:53.079655Z","shell.execute_reply.started":"2022-05-29T08:12:53.07251Z","shell.execute_reply":"2022-05-29T08:12:53.078412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 42\nIMG_SIZE = 512\nPROJECT_FOLDER = \"../input/hotelid-2022-train-images-512x512/\"\nTRAIN_DATA_FOLDER = \"../input/hotelid-2022-train-images-512x512/images/\"\nTEST_DATA_FOLDER = \"../input/hotel-id-to-combat-human-trafficking-2022-fgvc9/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":"2022-05-29T08:12:53.081171Z","iopub.execute_input":"2022-05-29T08:12:53.081724Z","iopub.status.idle":"2022-05-29T08:12:53.097164Z","shell.execute_reply.started":"2022-05-29T08:12:53.081594Z","shell.execute_reply":"2022-05-29T08:12:53.096305Z"},"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":"2022-05-29T08:12:53.099133Z","iopub.execute_input":"2022-05-29T08:12:53.099452Z","iopub.status.idle":"2022-05-29T08:12:53.116197Z","shell.execute_reply.started":"2022-05-29T08:12:53.099418Z","shell.execute_reply":"2022-05-29T08:12:53.114583Z"},"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\n\n\nbase_transform = A.Compose([\n    A.ToFloat(),\n    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":"2022-05-29T08:12:53.118098Z","iopub.execute_input":"2022-05-29T08:12:53.118571Z","iopub.status.idle":"2022-05-29T08:12:54.066917Z","shell.execute_reply.started":"2022-05-29T08:12:53.118535Z","shell.execute_reply":"2022-05-29T08:12:54.065786Z"},"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        \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":"2022-05-29T08:12:54.069357Z","iopub.execute_input":"2022-05-29T08:12:54.069692Z","iopub.status.idle":"2022-05-29T08:12:54.081444Z","shell.execute_reply.started":"2022-05-29T08:12:54.069655Z","shell.execute_reply":"2022-05-29T08:12:54.080442Z"},"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=args.device)\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 = self.backbone.default_cfg['classifier']\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        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":{"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":"2022-05-29T08:12:54.083467Z","iopub.execute_input":"2022-05-29T08:12:54.083842Z","iopub.status.idle":"2022-05-29T08:12:54.283173Z","shell.execute_reply.started":"2022-05-29T08:12:54.083806Z","shell.execute_reply":"2022-05-29T08:12:54.281863Z"},"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 = self.backbone.default_cfg['classifier']\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        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-29T08:12:54.284906Z","iopub.execute_input":"2022-05-29T08:12:54.285321Z","iopub.status.idle":"2022-05-29T08:12:54.299313Z","shell.execute_reply.started":"2022-05-29T08:12:54.285279Z","shell.execute_reply":"2022-05-29T08:12:54.298248Z"},"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":"2022-05-29T08:12:54.300975Z","iopub.execute_input":"2022-05-29T08:12:54.301463Z","iopub.status.idle":"2022-05-29T08:12:54.310846Z","shell.execute_reply.started":"2022-05-29T08:12:54.301428Z","shell.execute_reply":"2022-05-29T08:12:54.309968Z"},"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']/255.).to(args.device)\n            distances = get_distances(input/255., 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-29T08:12:54.312177Z","iopub.execute_input":"2022-05-29T08:12:54.312568Z","iopub.status.idle":"2022-05-29T08:12:54.325549Z","shell.execute_reply.started":"2022-05-29T08:12:54.312531Z","shell.execute_reply":"2022-05-29T08:12:54.324684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare data","metadata":{"id":"AwShW1wXniD6"}},{"cell_type":"markdown","source":"### for submission","metadata":{}},{"cell_type":"code","source":"data_df = pd.read_csv(PROJECT_FOLDER + \"train.csv\")\ntest_df = pd.DataFrame(data={\"image_id\": os.listdir(TEST_DATA_FOLDER), \"hotel_id\": \"\"}).sort_values(by=\"image_id\")\nprint(test_df)","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":"2022-05-29T08:12:54.326964Z","iopub.execute_input":"2022-05-29T08:12:54.327402Z","iopub.status.idle":"2022-05-29T08:12:54.4002Z","shell.execute_reply.started":"2022-05-29T08:12:54.327368Z","shell.execute_reply":"2022-05-29T08:12:54.399358Z"},"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(args.n_classes, embed_size, backbone_name)\n    else:\n        model = EmbeddingNet(args.n_classes, 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-29T08:12:54.401353Z","iopub.execute_input":"2022-05-29T08:12:54.401836Z","iopub.status.idle":"2022-05-29T08:12:54.408248Z","shell.execute_reply.started":"2022-05-29T08:12:54.401798Z","shell.execute_reply":"2022-05-29T08:12:54.406842Z"},"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, base_transform, 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":"2022-05-29T08:12:54.409792Z","iopub.execute_input":"2022-05-29T08:12:54.410194Z","iopub.status.idle":"2022-05-29T08:12:54.474184Z","shell.execute_reply.started":"2022-05-29T08:12:54.410161Z","shell.execute_reply":"2022-05-29T08:12:54.473387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_array = [\n                get_model(\"arcmargin\", \n                         \"efficientnet_b7\", 4096,\n                         \"../input/hotelarcmarginmodels/checkpoint-arcmargin-model-efficientnet_b1-512x512-4096embeds-3116hotels.pt\", \n                         args),\n               \n               get_model(\"cosface\", \n                         \"ecaresnet50d_pruned\", 4096,\n                         \"../input/hotelidcosfaceecaresnet50dtrained/checkpoint-cosface-model-ecaresnet50d_pruned-512x512-4096embeds-3116hotels.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/hotelarcmarginmodels/checkpoint-arcmargin-model-eca_nfnet_l0-512x512-1024embeds-3116hotels.pt\",\n                         args),\n              ]","metadata":{"execution":{"iopub.status.busy":"2022-05-29T08:12:54.475452Z","iopub.execute_input":"2022-05-29T08:12:54.475798Z","iopub.status.idle":"2022-05-29T08:13:21.742009Z","shell.execute_reply.started":"2022-05-29T08:12:54.475764Z","shell.execute_reply":"2022-05-29T08:13:21.741171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### submission","metadata":{}},{"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]\n\ntest_df.to_csv(\"submission.csv\", index=False)\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-05-29T08:13:21.743399Z","iopub.execute_input":"2022-05-29T08:13:21.743722Z","iopub.status.idle":"2022-05-29T08:13:41.376639Z","shell.execute_reply.started":"2022-05-29T08:13:21.743688Z","shell.execute_reply":"2022-05-29T08:13:41.375658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}