{"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":"<br>\n<h2 style = \"font-size:60px; font-family:Garamond ; font-weight : normal; background-color: #f6f5f5 ; color : #fe346e; text-align: center; border-radius: 100px 100px;\">FAISS Pytorch Inference</h2>\n<br>","metadata":{}},{"cell_type":"markdown","source":"<h3>📌 Inference Pipeline is taken from this Notebook:</h3> <h4><a href='https://www.kaggle.com/ks2019/happywhale-arcface-baseline-tpu'>https://www.kaggle.com/ks2019/happywhale-arcface-baseline-tpu</a></h4>\n\n<h3>📌 Train Notebook:</h3> <h4><a href='https://www.kaggle.com/debarshichanda/pytorch-arcface-gem-pooling-starter'>https://www.kaggle.com/debarshichanda/pytorch-arcface-gem-pooling-starter</a></h4>","metadata":{}},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Install Required Libraries</h1></span>","metadata":{}},{"cell_type":"code","source":"!pip install timm\n!pip install faiss-gpu","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-02-23T04:59:32.021658Z","iopub.execute_input":"2022-02-23T04:59:32.021990Z","iopub.status.idle":"2022-02-23T04:59:51.912886Z","shell.execute_reply.started":"2022-02-23T04:59:32.021907Z","shell.execute_reply":"2022-02-23T04:59:51.912022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Import Required Libraries 📚</h1></span>","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport math\nimport copy\nimport time\nimport random\n\n# For data manipulation\nimport numpy as np\nimport pandas as pd\n\n# Pytorch Imports\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\n\n# Utils\nimport joblib\nfrom tqdm import tqdm\nfrom collections import defaultdict\n\n# Sklearn Imports\nfrom sklearn.preprocessing import LabelEncoder, normalize\nfrom sklearn.model_selection import StratifiedKFold\n\n# For Image Models\nimport timm\n\n# For Similarity Search\nimport faiss\n\n# Albumentations for augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\nb_ = Fore.BLUE\ny_ = Fore.YELLOW\nsr_ = Style.RESET_ALL\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:51.916622Z","iopub.execute_input":"2022-02-23T04:59:51.916847Z","iopub.status.idle":"2022-02-23T04:59:56.733979Z","shell.execute_reply.started":"2022-02-23T04:59:51.916818Z","shell.execute_reply":"2022-02-23T04:59:56.733200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Configuration ⚙️</h1></span>","metadata":{}},{"cell_type":"code","source":"CONFIG = {\"seed\": 2022,\n          \"img_size\": 448,\n          \"model_name\": \"tf_efficientnet_b0_ns\",\n          \"num_classes\": 15587,\n          \"embedding_size\": 512,\n          \"train_batch_size\": 64,\n          \"valid_batch_size\": 64,\n          \"n_fold\": 5,\n          \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n          # ArcFace Hyperparameters\n          \"s\": 30.0, \n          \"m\": 0.30,\n          \"ls_eps\": 0.0,\n          \"easy_margin\": False\n          }","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:56.735623Z","iopub.execute_input":"2022-02-23T04:59:56.735890Z","iopub.status.idle":"2022-02-23T04:59:56.789007Z","shell.execute_reply.started":"2022-02-23T04:59:56.735856Z","shell.execute_reply":"2022-02-23T04:59:56.788121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Set Seed for Reproducibility</h1></span>","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:56.790674Z","iopub.execute_input":"2022-02-23T04:59:56.790941Z","iopub.status.idle":"2022-02-23T04:59:56.802207Z","shell.execute_reply.started":"2022-02-23T04:59:56.790906Z","shell.execute_reply":"2022-02-23T04:59:56.801454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '../input/happy-whale-and-dolphin'\nTRAIN_DIR = '../input/happy-whale-and-dolphin/train_images'\nTEST_DIR = '../input/happy-whale-and-dolphin/test_images'","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:56.806069Z","iopub.execute_input":"2022-02-23T04:59:56.806330Z","iopub.status.idle":"2022-02-23T04:59:56.811462Z","shell.execute_reply.started":"2022-02-23T04:59:56.806296Z","shell.execute_reply":"2022-02-23T04:59:56.810662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_file_path(id):\n    return f\"{TRAIN_DIR}/{id}\"","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:56.813025Z","iopub.execute_input":"2022-02-23T04:59:56.813291Z","iopub.status.idle":"2022-02-23T04:59:56.822830Z","shell.execute_reply.started":"2022-02-23T04:59:56.813258Z","shell.execute_reply":"2022-02-23T04:59:56.822043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Read the Data 📖</h1>","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f\"{ROOT_DIR}/train.csv\")\ndf['file_path'] = df['image'].apply(get_train_file_path)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:56.824466Z","iopub.execute_input":"2022-02-23T04:59:56.824895Z","iopub.status.idle":"2022-02-23T04:59:56.938328Z","shell.execute_reply.started":"2022-02-23T04:59:56.824790Z","shell.execute_reply":"2022-02-23T04:59:56.937662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = LabelEncoder()\n\nwith open(\"../input/arcface-gap-embed/le.pkl\", \"rb\") as fp:\n    encoder = joblib.load(fp)\n    \ndf['individual_id'] = encoder.transform(df['individual_id'])","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:56.939586Z","iopub.execute_input":"2022-02-23T04:59:56.939835Z","iopub.status.idle":"2022-02-23T04:59:56.972073Z","shell.execute_reply.started":"2022-02-23T04:59:56.939790Z","shell.execute_reply":"2022-02-23T04:59:56.971448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=CONFIG['n_fold'])\n\nfor fold, ( _, val_) in enumerate(skf.split(X=df, y=df.individual_id)):\n      df.loc[val_ , \"kfold\"] = fold","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:56.973410Z","iopub.execute_input":"2022-02-23T04:59:56.973667Z","iopub.status.idle":"2022-02-23T04:59:57.431088Z","shell.execute_reply.started":"2022-02-23T04:59:56.973632Z","shell.execute_reply":"2022-02-23T04:59:57.430350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Dataset Class</h1></span>","metadata":{}},{"cell_type":"code","source":"class HappyWhaleDataset(Dataset):\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.ids = df['image'].values\n        self.file_names = df['file_path'].values\n        self.labels = df['individual_id'].values\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        idx = self.ids[index]\n        img_path = self.file_names[index]\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        label = self.labels[index]\n        \n        if self.transforms:\n            img = self.transforms(image=img)[\"image\"]\n            \n        return {\n            'image': img,\n            'label': torch.tensor(label, dtype=torch.long),\n            'id': idx\n        }","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:57.432529Z","iopub.execute_input":"2022-02-23T04:59:57.432791Z","iopub.status.idle":"2022-02-23T04:59:57.440622Z","shell.execute_reply.started":"2022-02-23T04:59:57.432756Z","shell.execute_reply":"2022-02-23T04:59:57.439976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Augmentations</h1></span>","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.),\n    \n    \"valid\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.)\n}","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:57.441934Z","iopub.execute_input":"2022-02-23T04:59:57.442399Z","iopub.status.idle":"2022-02-23T04:59:57.451537Z","shell.execute_reply.started":"2022-02-23T04:59:57.442363Z","shell.execute_reply":"2022-02-23T04:59:57.450873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">GeM Pooling</h1></span>","metadata":{}},{"cell_type":"code","source":"# NOT USED\n\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n        \n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n        \n    def __repr__(self):\n        return self.__class__.__name__ + \\\n                '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + \\\n                ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:57.453860Z","iopub.execute_input":"2022-02-23T04:59:57.454366Z","iopub.status.idle":"2022-02-23T04:59:57.465032Z","shell.execute_reply.started":"2022-02-23T04:59:57.454330Z","shell.execute_reply":"2022-02-23T04:59:57.464243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">ArcFace</h1></span>","metadata":{}},{"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, \n                 m=0.50, easy_margin=False, ls_eps=0.0):\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.ls_eps = ls_eps  # label smoothing\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))\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=CONFIG['device'])\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        if self.ls_eps > 0:\n            one_hot = (1 - self.ls_eps) * one_hot + self.ls_eps / self.out_features\n        # -------------torch.where(out_i = {x_i if condition_i else y_i) ------------\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:57.466470Z","iopub.execute_input":"2022-02-23T04:59:57.466735Z","iopub.status.idle":"2022-02-23T04:59:57.480498Z","shell.execute_reply.started":"2022-02-23T04:59:57.466699Z","shell.execute_reply":"2022-02-23T04:59:57.479808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Model</h1></span>","metadata":{}},{"cell_type":"code","source":"class HappyWhaleModel(nn.Module):\n    def __init__(self, model_name, embedding_size, pretrained=True):\n        super(HappyWhaleModel, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        in_features = self.model.classifier.in_features\n        self.model.classifier = nn.Identity()\n        self.embedding = nn.Linear(in_features, embedding_size)\n        self.fc = ArcMarginProduct(embedding_size, \n                                   CONFIG[\"num_classes\"],\n                                   s=CONFIG[\"s\"], \n                                   m=CONFIG[\"m\"], \n                                   easy_margin=CONFIG[\"ls_eps\"], \n                                   ls_eps=CONFIG[\"ls_eps\"])\n    \n    def forward(self, images, labels):\n        features = self.model(images)\n        embedding = self.embedding(features)\n        output = self.fc(embedding, labels)\n        return output\n    \n    def extract(self, images):\n        features = self.model(images)\n        embedding = self.embedding(features)\n        return embedding\n    \n\nmodel = HappyWhaleModel(CONFIG['model_name'], CONFIG['embedding_size'])\nmodel.load_state_dict(torch.load(\"../input/arcface-gap-embed/Loss14.0082_epoch10.bin\"))\nmodel.to(CONFIG['device']);","metadata":{"execution":{"iopub.status.busy":"2022-02-23T04:59:57.484926Z","iopub.execute_input":"2022-02-23T04:59:57.485470Z","iopub.status.idle":"2022-02-23T05:00:02.290255Z","shell.execute_reply.started":"2022-02-23T04:59:57.485440Z","shell.execute_reply":"2022-02-23T05:00:02.289438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.inference_mode()\ndef get_embeddings(model, dataloader, device):\n    model.eval()\n    \n    LABELS = []\n    EMBEDS = []\n    IDS = []\n    \n    bar = tqdm(enumerate(dataloader), total=len(dataloader))\n    for step, data in bar:        \n        images = data['image'].to(device, dtype=torch.float)\n        labels = data['label'].to(device, dtype=torch.long)\n        ids = data['id']\n\n        outputs = model.extract(images)\n        \n        LABELS.append(labels.cpu().numpy())\n        EMBEDS.append(outputs.cpu().numpy())\n        IDS.append(ids)\n    \n    EMBEDS = np.vstack(EMBEDS)\n    LABELS = np.concatenate(LABELS)\n    IDS = np.concatenate(IDS)\n    \n    return EMBEDS, LABELS, IDS","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:02.291452Z","iopub.execute_input":"2022-02-23T05:00:02.291718Z","iopub.status.idle":"2022-02-23T05:00:02.300857Z","shell.execute_reply.started":"2022-02-23T05:00:02.291683Z","shell.execute_reply":"2022-02-23T05:00:02.300062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_loaders(df, fold):\n    df_train = df[df.kfold != fold].reset_index(drop=True)\n    df_valid = df[df.kfold == fold].reset_index(drop=True)\n    \n    train_dataset = HappyWhaleDataset(df_train, transforms=data_transforms[\"train\"])\n    valid_dataset = HappyWhaleDataset(df_valid, transforms=data_transforms[\"valid\"])\n\n    train_loader = DataLoader(train_dataset, batch_size=CONFIG['train_batch_size'], \n                              num_workers=2, shuffle=False, pin_memory=True)\n    valid_loader = DataLoader(valid_dataset, batch_size=CONFIG['valid_batch_size'], \n                              num_workers=2, shuffle=False, pin_memory=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:02.302604Z","iopub.execute_input":"2022-02-23T05:00:02.302935Z","iopub.status.idle":"2022-02-23T05:00:02.313607Z","shell.execute_reply.started":"2022-02-23T05:00:02.302901Z","shell.execute_reply":"2022-02-23T05:00:02.312768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader, valid_loader = prepare_loaders(df, fold=0)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:02.314684Z","iopub.execute_input":"2022-02-23T05:00:02.315706Z","iopub.status.idle":"2022-02-23T05:00:02.334444Z","shell.execute_reply.started":"2022-02-23T05:00:02.315670Z","shell.execute_reply":"2022-02-23T05:00:02.333815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_embeds, train_labels, train_ids = get_embeddings(model, train_loader, CONFIG['device'])\nvalid_embeds, valid_labels, valid_ids = get_embeddings(model, valid_loader, CONFIG['device'])","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:02.335818Z","iopub.execute_input":"2022-02-23T05:00:02.336201Z","iopub.status.idle":"2022-02-23T05:00:54.864317Z","shell.execute_reply.started":"2022-02-23T05:00:02.336167Z","shell.execute_reply":"2022-02-23T05:00:54.863466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_embeds = normalize(train_embeds, axis=1, norm='l2')\nvalid_embeds = normalize(valid_embeds, axis=1, norm='l2')","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:54.866186Z","iopub.execute_input":"2022-02-23T05:00:54.866663Z","iopub.status.idle":"2022-02-23T05:00:54.874859Z","shell.execute_reply.started":"2022-02-23T05:00:54.866617Z","shell.execute_reply":"2022-02-23T05:00:54.874007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = encoder.inverse_transform(train_labels)\nvalid_labels = encoder.inverse_transform(valid_labels)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:54.876372Z","iopub.execute_input":"2022-02-23T05:00:54.876722Z","iopub.status.idle":"2022-02-23T05:00:54.884933Z","shell.execute_reply.started":"2022-02-23T05:00:54.876684Z","shell.execute_reply":"2022-02-23T05:00:54.884214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"index = faiss.IndexFlatIP(CONFIG['embedding_size'])\nindex.add(train_embeds)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:54.886491Z","iopub.execute_input":"2022-02-23T05:00:54.887050Z","iopub.status.idle":"2022-02-23T05:00:54.893516Z","shell.execute_reply.started":"2022-02-23T05:00:54.887011Z","shell.execute_reply":"2022-02-23T05:00:54.892777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"D, I = index.search(valid_embeds, k=50)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:54.895057Z","iopub.execute_input":"2022-02-23T05:00:54.895681Z","iopub.status.idle":"2022-02-23T05:00:54.904151Z","shell.execute_reply.started":"2022-02-23T05:00:54.895642Z","shell.execute_reply":"2022-02-23T05:00:54.903299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"allowed_targets = np.unique(train_labels)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:54.917104Z","iopub.execute_input":"2022-02-23T05:00:54.917380Z","iopub.status.idle":"2022-02-23T05:00:54.924794Z","shell.execute_reply.started":"2022-02-23T05:00:54.917341Z","shell.execute_reply":"2022-02-23T05:00:54.924068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_targets_df = pd.DataFrame(np.stack([valid_ids, valid_labels], axis=1), columns=['image','target'])\nval_targets_df.loc[~val_targets_df.target.isin(allowed_targets), 'target'] = 'new_individual'\nval_targets_df.target.value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:54.926041Z","iopub.execute_input":"2022-02-23T05:00:54.926828Z","iopub.status.idle":"2022-02-23T05:00:54.945273Z","shell.execute_reply.started":"2022-02-23T05:00:54.926792Z","shell.execute_reply":"2022-02-23T05:00:54.944304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df = []\nfor i, val_id in tqdm(enumerate(valid_ids)):\n    targets = train_labels[I[i]]\n    distances = D[i]\n    subset_preds = pd.DataFrame(np.stack([targets,distances],axis=1),columns=['target','distances'])\n    subset_preds['image'] = val_id\n    valid_df.append(subset_preds)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:54.946382Z","iopub.execute_input":"2022-02-23T05:00:54.947146Z","iopub.status.idle":"2022-02-23T05:00:54.963002Z","shell.execute_reply.started":"2022-02-23T05:00:54.947086Z","shell.execute_reply":"2022-02-23T05:00:54.962313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df = pd.concat(valid_df).reset_index(drop=True)\nvalid_df = valid_df.groupby(['image','target']).distances.max().reset_index()\nvalid_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:54.964210Z","iopub.execute_input":"2022-02-23T05:00:54.964668Z","iopub.status.idle":"2022-02-23T05:00:55.014194Z","shell.execute_reply.started":"2022-02-23T05:00:54.964631Z","shell.execute_reply":"2022-02-23T05:00:55.013481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df = valid_df.sort_values('distances', ascending=False).reset_index(drop=True)\nvalid_df.to_csv('val_neighbors.csv')","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.015517Z","iopub.execute_input":"2022-02-23T05:00:55.015761Z","iopub.status.idle":"2022-02-23T05:00:55.025946Z","shell.execute_reply.started":"2022-02-23T05:00:55.015727Z","shell.execute_reply":"2022-02-23T05:00:55.025168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_list = ['938b7e931166', '5bf17305f073', '7593d2aee842', '7362d7a01d00','956562ff2888']","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.027034Z","iopub.execute_input":"2022-02-23T05:00:55.028249Z","iopub.status.idle":"2022-02-23T05:00:55.032995Z","shell.execute_reply.started":"2022-02-23T05:00:55.028208Z","shell.execute_reply":"2022-02-23T05:00:55.031959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_predictions(test_df, threshold=0.2):\n    predictions = {}\n    for i, row in tqdm(test_df.iterrows()):\n        if row.image in predictions:\n            if len(predictions[row.image]) == 5:\n                continue\n            predictions[row.image].append(row.target)\n        elif row.distances > threshold:\n            predictions[row.image] = [row.target, 'new_individual']\n        else:\n            predictions[row.image] = ['new_individual', row.target]\n\n    for x in tqdm(predictions):\n        if len(predictions[x]) < 5:\n            remaining = [y for y in sample_list if y not in predictions]\n            predictions[x] = predictions[x] + remaining\n            predictions[x] = predictions[x][:5]\n        \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.034558Z","iopub.execute_input":"2022-02-23T05:00:55.035144Z","iopub.status.idle":"2022-02-23T05:00:55.044809Z","shell.execute_reply.started":"2022-02-23T05:00:55.035090Z","shell.execute_reply":"2022-02-23T05:00:55.044064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def map_per_image(label, predictions):\n    \"\"\"Computes the precision score of one image.\n\n    Parameters\n    ----------\n    label : string\n            The true label of the image\n    predictions : list\n            A list of predicted elements (order does matter, 5 predictions allowed per image)\n\n    Returns\n    -------\n    score : double\n    \"\"\"    \n    try:\n        return 1 / (predictions[:5].index(label) + 1)\n    except ValueError:\n        return 0.0","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.046296Z","iopub.execute_input":"2022-02-23T05:00:55.046666Z","iopub.status.idle":"2022-02-23T05:00:55.053961Z","shell.execute_reply.started":"2022-02-23T05:00:55.046629Z","shell.execute_reply":"2022-02-23T05:00:55.053214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Compute CV</h1></span>","metadata":{}},{"cell_type":"code","source":"best_th = 0\nbest_cv = 0\nfor th in [0.1*x for x in range(11)]:\n    all_preds = get_predictions(valid_df, threshold=th)\n    cv = 0\n    for i,row in val_targets_df.iterrows():\n        target = row.target\n        preds = all_preds[row.image]\n        val_targets_df.loc[i,th] = map_per_image(target, preds)\n    cv = val_targets_df[th].mean()\n    print(f\"CV at threshold {th}: {cv}\")\n    if cv > best_cv:\n        best_th = th\n        best_cv = cv","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.056312Z","iopub.execute_input":"2022-02-23T05:00:55.056517Z","iopub.status.idle":"2022-02-23T05:00:55.360911Z","shell.execute_reply.started":"2022-02-23T05:00:55.056494Z","shell.execute_reply":"2022-02-23T05:00:55.360179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Best threshold\", best_th)\nprint(\"Best cv\", best_cv)\nval_targets_df.describe()","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.362273Z","iopub.execute_input":"2022-02-23T05:00:55.362859Z","iopub.status.idle":"2022-02-23T05:00:55.412706Z","shell.execute_reply.started":"2022-02-23T05:00:55.362818Z","shell.execute_reply":"2022-02-23T05:00:55.411972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Adjustment: Since Public lb has nearly 10% 'new_individual' (Be Careful for private LB)\nval_targets_df['is_new_individual'] = val_targets_df.target=='new_individual'\nprint(val_targets_df.is_new_individual.value_counts().to_dict())\nval_scores = val_targets_df.groupby('is_new_individual').mean().T\nval_scores['adjusted_cv'] = val_scores[True]*0.1+val_scores[False]*0.9\nbest_threshold_adjusted = val_scores['adjusted_cv'].idxmax()\nprint(\"best_threshold\",best_threshold_adjusted)\nval_scores","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.413862Z","iopub.execute_input":"2022-02-23T05:00:55.414191Z","iopub.status.idle":"2022-02-23T05:00:55.682535Z","shell.execute_reply.started":"2022-02-23T05:00:55.414153Z","shell.execute_reply":"2022-02-23T05:00:55.680975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span><h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Inference</h1></span>","metadata":{}},{"cell_type":"code","source":"train_embeds = np.concatenate([train_embeds, valid_embeds])\ntrain_labels = np.concatenate([train_labels, valid_labels])\nprint(train_embeds.shape,train_labels.shape)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.683791Z","iopub.status.idle":"2022-02-23T05:00:55.684442Z","shell.execute_reply.started":"2022-02-23T05:00:55.684185Z","shell.execute_reply":"2022-02-23T05:00:55.684211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"index = faiss.IndexFlatIP(CONFIG['embedding_size'])\nindex.add(train_embeds)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.685711Z","iopub.status.idle":"2022-02-23T05:00:55.686338Z","shell.execute_reply.started":"2022-02-23T05:00:55.686082Z","shell.execute_reply":"2022-02-23T05:00:55.686106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pd.DataFrame()\ntest[\"image\"] = os.listdir(\"../input/happy-whale-and-dolphin/test_images\")\ntest[\"file_path\"] = test[\"image\"].apply(lambda x: f\"{TEST_DIR}/{x}\")\ntest[\"individual_id\"] = -1  #dummy value\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.687503Z","iopub.status.idle":"2022-02-23T05:00:55.688112Z","shell.execute_reply.started":"2022-02-23T05:00:55.687875Z","shell.execute_reply":"2022-02-23T05:00:55.687899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = HappyWhaleDataset(test, transforms=data_transforms[\"valid\"])\ntest_loader = DataLoader(test_dataset, batch_size=CONFIG['valid_batch_size'], \n                         num_workers=2, shuffle=False, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.689310Z","iopub.status.idle":"2022-02-23T05:00:55.689909Z","shell.execute_reply.started":"2022-02-23T05:00:55.689679Z","shell.execute_reply":"2022-02-23T05:00:55.689702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_embeds, _, test_ids = get_embeddings(model, test_loader, CONFIG['device'])\ntest_embeds = normalize(test_embeds, axis=1, norm='l2')","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.691059Z","iopub.status.idle":"2022-02-23T05:00:55.691685Z","shell.execute_reply.started":"2022-02-23T05:00:55.691453Z","shell.execute_reply":"2022-02-23T05:00:55.691476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"D, I = index.search(test_embeds, k=50)","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.692857Z","iopub.status.idle":"2022-02-23T05:00:55.693477Z","shell.execute_reply.started":"2022-02-23T05:00:55.693237Z","shell.execute_reply":"2022-02-23T05:00:55.693262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = []\nfor i, test_id in tqdm(enumerate(test_ids)):\n    targets = train_labels[I[i]]\n    distances = D[i]\n    subset_preds = pd.DataFrame(np.stack([targets, distances], axis=1), columns=['target','distances'])\n    subset_preds['image'] = test_id\n    test_df.append(subset_preds)\n    \ntest_df = pd.concat(test_df).reset_index(drop=True)\ntest_df = test_df.groupby(['image','target']).distances.max().reset_index()\ntest_df = test_df.sort_values('distances', ascending=False).reset_index(drop=True)\ntest_df.to_csv('test_neighbors.csv')","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.694676Z","iopub.status.idle":"2022-02-23T05:00:55.695311Z","shell.execute_reply.started":"2022-02-23T05:00:55.695055Z","shell.execute_reply":"2022-02-23T05:00:55.695080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = get_predictions(test_df, best_threshold_adjusted)\n\npredictions = pd.Series(predictions).reset_index()\npredictions.columns = ['image','predictions']\npredictions['predictions'] = predictions['predictions'].apply(lambda x: ' '.join(x))\npredictions.to_csv('submission.csv',index=False)\npredictions.head()","metadata":{"execution":{"iopub.status.busy":"2022-02-23T05:00:55.696469Z","iopub.status.idle":"2022-02-23T05:00:55.697068Z","shell.execute_reply.started":"2022-02-23T05:00:55.696838Z","shell.execute_reply":"2022-02-23T05:00:55.696862Z"},"trusted":true},"execution_count":null,"outputs":[]}]}