{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nfrom sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2022-04-19T21:51:36.007509Z","iopub.execute_input":"2022-04-19T21:51:36.007782Z","iopub.status.idle":"2022-04-19T21:51:37.001798Z","shell.execute_reply.started":"2022-04-19T21:51:36.007744Z","shell.execute_reply":"2022-04-19T21:51:37.001074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X = pd.read_csv('../input/happy-whale-and-dolphin/train.csv')\nX['path'] = '../input/happy-whale-and-dolphin/train_images/' + X['image']\nX['species'].replace({\n    'bottlenose_dolpin' : 'bottlenose_dolphin',\n    'kiler_whale' : 'killer_whale',\n    'beluga' : 'beluga_whale',\n    'globis' : 'short_finned_pilot_whale',\n    'pilot_whale' : 'short_finned_pilot_whale'\n},inplace =True)\n\nX['class'] = X['species'].apply(lambda x: x.split('_')[-1])\nX.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-19T21:51:37.003214Z","iopub.execute_input":"2022-04-19T21:51:37.003451Z","iopub.status.idle":"2022-04-19T21:51:37.15172Z","shell.execute_reply.started":"2022-04-19T21:51:37.003419Z","shell.execute_reply":"2022-04-19T21:51:37.151033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train_test_split in the training dataset","metadata":{}},{"cell_type":"code","source":"y = X['individual_id']\n\nX_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.33, random_state=42)\n\ntrain_df = X_train\n\n# train_df = X_test","metadata":{"execution":{"iopub.status.busy":"2022-04-19T21:51:37.152733Z","iopub.execute_input":"2022-04-19T21:51:37.152982Z","iopub.status.idle":"2022-04-19T21:51:37.176003Z","shell.execute_reply.started":"2022-04-19T21:51:37.152947Z","shell.execute_reply":"2022-04-19T21:51:37.175331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot = sns.countplot(x = train_df['class'], color = '#2596be')\nsns.despine()\nplot.set_title('Class Distribution\\n', font = 'serif', x = 0.1, y=1, fontsize = 16);\nplot.set_ylabel(\"Count\", x = 0.02, font = 'serif', fontsize = 12)\nplot.set_xlabel(\"Specie\", fontsize = 12, font = 'serif')\n\nfor p in plot.patches:\n    plot.annotate(format(p.get_height(), '.0f'), (p.get_x() + p.get_width() / 2, p.get_height()), \n       ha = 'center', va = 'center', xytext = (0, -20),font = 'serif', textcoords = 'offset points', size = 15)\n    \nplt.figure(figsize=(5,5))\nclass_cnt = train_df.groupby(['class']).size().reset_index(name = 'counts')\ncolors = sns.color_palette('Paired')[0:9]\nplt.pie(class_cnt['counts'], labels=class_cnt['class'], colors=colors, autopct='%1.1f%%')\nplt.legend(loc='upper left')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2022-04-19T21:51:37.177906Z","iopub.execute_input":"2022-04-19T21:51:37.178226Z","iopub.status.idle":"2022-04-19T21:51:37.506635Z","shell.execute_reply.started":"2022-04-19T21:51:37.178188Z","shell.execute_reply":"2022-04-19T21:51:37.505811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(6,8))\nsns.countplot(data=train_df, y='individual_id',  palette='crest', dodge=False,\n             order=train_df['individual_id'].value_counts().index[:1000:40])\nplt.title(\"Count of some individual Id\")\nplt.show()\nsum(train_df['individual_id'].value_counts() == 1)\norder=train_df['individual_id'].value_counts().index[::200]\nX.shape","metadata":{"execution":{"iopub.status.busy":"2022-04-19T21:51:37.510823Z","iopub.execute_input":"2022-04-19T21:51:37.51117Z","iopub.status.idle":"2022-04-19T21:51:37.92022Z","shell.execute_reply.started":"2022-04-19T21:51:37.511129Z","shell.execute_reply":"2022-04-19T21:51:37.919433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🐋🐬 PyTorch Lightning ArcFace Focal Loss\n\nLet's train [`timm`](https://github.com/rwightman/pytorch-image-models) models \nwith [PyTorch Lightning](https://www.pytorchlightning.ai/)!\n\n**Sources:**\n- [[Pytorch] ArcFace + GeM Pooling Starter](https://www.kaggle.com/debarshichanda/pytorch-arcface-gem-pooling-starter)\n- [FAISS Pytorch Inference](https://www.kaggle.com/debarshichanda/faiss-pytorch-inference)\n- [whale2-cropped-dataset](https://www.kaggle.com/phalanx/whale2-cropped-dataset)\n\n**Public LB scores**\n- V02: 0.725(`image_size=512`, `\"tf_efficientnet_b5\"`, `batch_size=16`, `loss=focal loss`) \n- V03: 0.722(`image_size=512`, `\"tf_efficientnet_b5\"`, `batch_size=16`, `loss=cross entropy`) \n\n**Reference**\n- https://www.kaggle.com/code/clemchris/pytorch-backfin-convnext-arcface\n\n# *Due to the uneven data distribution, the effect of focal loss is stronger than that of cross entropy loss.*","metadata":{"papermill":{"duration":0.023074,"end_time":"2022-03-20T06:54:27.735725","exception":false,"start_time":"2022-03-20T06:54:27.712651","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Installations (timm + FAISS)","metadata":{"papermill":{"duration":0.021893,"end_time":"2022-03-20T06:54:27.781233","exception":false,"start_time":"2022-03-20T06:54:27.75934","status":"completed"},"tags":[]}},{"cell_type":"code","source":"pip install timm faiss-gpu","metadata":{"papermill":{"duration":14.765868,"end_time":"2022-03-20T06:54:42.56991","exception":false,"start_time":"2022-03-20T06:54:27.804042","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:37.921651Z","iopub.execute_input":"2022-04-19T21:51:37.9219Z","iopub.status.idle":"2022-04-19T21:51:52.623049Z","shell.execute_reply.started":"2022-04-19T21:51:37.921866Z","shell.execute_reply":"2022-04-19T21:51:52.622194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{"papermill":{"duration":0.057471,"end_time":"2022-03-20T06:54:42.688799","exception":false,"start_time":"2022-03-20T06:54:42.631328","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import math\nfrom typing import Callable\nfrom typing import Dict\nfrom typing import Optional\nfrom typing import Tuple\nfrom pathlib import Path\n\nimport faiss\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom PIL import Image\nfrom pytorch_lightning.callbacks.model_checkpoint import ModelCheckpoint\nfrom timm.data.transforms_factory import create_transform\nfrom timm.optim import create_optimizer_v2\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset\nfrom tqdm.notebook import tqdm\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.preprocessing import normalize\nfrom sklearn.preprocessing import LabelEncoder","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":4.73183,"end_time":"2022-03-20T06:54:47.479283","exception":false,"start_time":"2022-03-20T06:54:42.747453","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:52.626097Z","iopub.execute_input":"2022-04-19T21:51:52.626569Z","iopub.status.idle":"2022-04-19T21:51:56.354422Z","shell.execute_reply.started":"2022-04-19T21:51:52.626536Z","shell.execute_reply":"2022-04-19T21:51:56.35358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Focal loss","metadata":{}},{"cell_type":"code","source":"from torch.autograd import Variable\nclass FocalLoss(nn.Module):\n    r\"\"\"\n        This criterion is a implemenation of Focal Loss, which is proposed in \n        Focal Loss for Dense Object Detection.\n\n            Loss(x, class) = - \\alpha (1-softmax(x)[class])^gamma \\log(softmax(x)[class])\n\n        The losses are averaged across observations for each minibatch.\n\n        Args:\n            alpha(1D Tensor, Variable) : the scalar factor for this criterion\n            gamma(float, double) : gamma > 0; reduces the relative loss for well-classiﬁed examples (p > .5), \n                                   putting more focus on hard, misclassiﬁed examples\n            size_average(bool): By default, the losses are averaged over observations for each minibatch.\n                                However, if the field size_average is set to False, the losses are\n                                instead summed for each minibatch.\n\n\n    \"\"\"\n    def __init__(self, class_num=15587, alpha=None, gamma=2, size_average=True):\n        super(FocalLoss, self).__init__()\n        if alpha is None:\n            self.alpha = Variable(torch.ones(class_num, 1))\n        else:\n            if isinstance(alpha, Variable):\n                self.alpha = alpha\n            else:\n                self.alpha = Variable(alpha)\n        self.gamma = gamma\n        self.class_num = class_num\n        self.size_average = size_average\n\n    def forward(self, inputs, targets):\n        N = inputs.size(0)\n        C = inputs.size(1)\n        P = F.softmax(inputs,dim=1)\n\n        class_mask = inputs.data.new(N, C).fill_(0)\n        class_mask = Variable(class_mask)\n        ids = targets.view(-1, 1)\n        class_mask.scatter_(1, ids.data, 1.)\n        #print(class_mask)\n\n\n        if inputs.is_cuda and not self.alpha.is_cuda:\n            self.alpha = self.alpha.cuda()\n        alpha = self.alpha[ids.data.view(-1)]\n\n        probs = (P*class_mask).sum(1).view(-1,1)\n\n        log_p = probs.log()\n        #print('probs size= {}'.format(probs.size()))\n        #print(probs)\n\n        batch_loss = -alpha*(torch.pow((1-probs), self.gamma))*log_p \n        #print('-----bacth_loss------')\n        #print(batch_loss)\n\n\n        if self.size_average:\n            loss = batch_loss.mean()\n        else:\n            loss = batch_loss.sum()\n        return loss","metadata":{"papermill":{"duration":0.070661,"end_time":"2022-03-20T06:54:47.608682","exception":false,"start_time":"2022-03-20T06:54:47.538021","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:56.35647Z","iopub.execute_input":"2022-04-19T21:51:56.356764Z","iopub.status.idle":"2022-04-19T21:51:56.37108Z","shell.execute_reply.started":"2022-04-19T21:51:56.35673Z","shell.execute_reply":"2022-04-19T21:51:56.370413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Paths & Settings","metadata":{"papermill":{"duration":0.057522,"end_time":"2022-03-20T06:54:47.722858","exception":false,"start_time":"2022-03-20T06:54:47.665336","status":"completed"},"tags":[]}},{"cell_type":"code","source":"INPUT_DIR = Path(\"..\") / \"input\"\nOUTPUT_DIR = Path(\"/\") / \"kaggle\" / \"working\"\n\nDATA_ROOT_DIR = INPUT_DIR / \"convert-backfintfrecords\" / \"happy-whale-and-dolphin-backfin\"\nTRAIN_DIR = DATA_ROOT_DIR / \"train_images\"\nTEST_DIR = DATA_ROOT_DIR / \"test_images\"\nTRAIN_CSV_PATH = DATA_ROOT_DIR / \"train.csv\"\nSAMPLE_SUBMISSION_CSV_PATH = DATA_ROOT_DIR / \"sample_submission.csv\"\nPUBLIC_SUBMISSION_CSV_PATH = INPUT_DIR / \"0-720-eff-b5-640-rotate\" / \"submission.csv\"\nIDS_WITHOUT_BACKFIN_PATH = INPUT_DIR / \"ids-without-backfin\" / \"ids_without_backfin.npy\"\n\nN_SPLITS = 5\n\nENCODER_CLASSES_PATH = OUTPUT_DIR / \"encoder_classes.npy\"\nTEST_CSV_PATH = OUTPUT_DIR / \"test.csv\"\nTRAIN_CSV_ENCODED_FOLDED_PATH = OUTPUT_DIR / \"train_encoded_folded.csv\"\nCHECKPOINTS_DIR = OUTPUT_DIR / \"checkpoints\"\nSUBMISSION_CSV_PATH = OUTPUT_DIR / \"submission.csv\"\n\nDEBUG = False","metadata":{"papermill":{"duration":0.065328,"end_time":"2022-03-20T06:54:47.844221","exception":false,"start_time":"2022-03-20T06:54:47.778893","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:56.372361Z","iopub.execute_input":"2022-04-19T21:51:56.37289Z","iopub.status.idle":"2022-04-19T21:51:56.386363Z","shell.execute_reply.started":"2022-04-19T21:51:56.372781Z","shell.execute_reply":"2022-04-19T21:51:56.385647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare DataFrames","metadata":{"papermill":{"duration":0.055587,"end_time":"2022-03-20T06:54:47.956036","exception":false,"start_time":"2022-03-20T06:54:47.900449","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_image_path(id: str, dir: Path) -> str:\n    return f\"{dir / id}\"","metadata":{"papermill":{"duration":0.066483,"end_time":"2022-03-20T06:54:48.079214","exception":false,"start_time":"2022-03-20T06:54:48.012731","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:56.390897Z","iopub.execute_input":"2022-04-19T21:51:56.391339Z","iopub.status.idle":"2022-04-19T21:51:56.399879Z","shell.execute_reply.started":"2022-04-19T21:51:56.391306Z","shell.execute_reply":"2022-04-19T21:51:56.398976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train DataFrame","metadata":{"papermill":{"duration":0.056289,"end_time":"2022-03-20T06:54:48.1961","exception":false,"start_time":"2022-03-20T06:54:48.139811","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV_PATH)\n\ntrain_df[\"image_path\"] = train_df[\"image\"].apply(get_image_path, dir=TRAIN_DIR)\n\nencoder = LabelEncoder()\ntrain_df[\"individual_id\"] = encoder.fit_transform(train_df[\"individual_id\"])\nnp.save(ENCODER_CLASSES_PATH, encoder.classes_)\n\nskf = StratifiedKFold(n_splits=N_SPLITS)\nfor fold, (_, val_) in enumerate(skf.split(X=train_df, y=train_df.individual_id)):\n    train_df.loc[val_, \"kfold\"] = fold\n    \ntrain_df.to_csv(TRAIN_CSV_ENCODED_FOLDED_PATH, index=False)\n    \ntrain_df.head()\n","metadata":{"papermill":{"duration":1.110036,"end_time":"2022-03-20T06:54:49.362372","exception":false,"start_time":"2022-03-20T06:54:48.252336","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:56.401152Z","iopub.execute_input":"2022-04-19T21:51:56.402014Z","iopub.status.idle":"2022-04-19T21:51:57.668638Z","shell.execute_reply.started":"2022-04-19T21:51:56.401972Z","shell.execute_reply":"2022-04-19T21:51:57.667975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test DataFrame","metadata":{"papermill":{"duration":0.057491,"end_time":"2022-03-20T06:54:49.47876","exception":false,"start_time":"2022-03-20T06:54:49.421269","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Use sample submission csv as template\ntest_df = pd.read_csv(SAMPLE_SUBMISSION_CSV_PATH)\ntest_df[\"image_path\"] = test_df[\"image\"].apply(get_image_path, dir=TEST_DIR)\n\ntest_df.drop(columns=[\"predictions\"], inplace=True)\n\n# Dummy id\ntest_df[\"individual_id\"] = 0\n\ntest_df.to_csv(TEST_CSV_PATH, index=False)\n\ntest_df.head()","metadata":{"papermill":{"duration":0.456441,"end_time":"2022-03-20T06:54:49.992142","exception":false,"start_time":"2022-03-20T06:54:49.535701","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:57.66987Z","iopub.execute_input":"2022-04-19T21:51:57.6701Z","iopub.status.idle":"2022-04-19T21:51:58.061615Z","shell.execute_reply.started":"2022-04-19T21:51:57.670068Z","shell.execute_reply":"2022-04-19T21:51:58.060951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"papermill":{"duration":0.058379,"end_time":"2022-03-20T06:54:50.108741","exception":false,"start_time":"2022-03-20T06:54:50.050362","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class HappyWhaleDataset(Dataset):\n    def __init__(self, df: pd.DataFrame, transform: Optional[Callable] = None):\n        self.df = df\n        self.transform = transform\n\n        self.image_names = self.df[\"image\"].values\n        self.image_paths = self.df[\"image_path\"].values\n        self.targets = self.df[\"individual_id\"].values\n\n    def __getitem__(self, index: int) -> Dict[str, torch.Tensor]:\n        image_name = self.image_names[index]\n\n        image_path = self.image_paths[index]\n\n        image = Image.open(image_path)\n        \n        if self.transform:\n            image = self.transform(image)\n\n        target = self.targets[index]\n        target = torch.tensor(target, dtype=torch.long)\n\n        return {\"image_name\": image_name, \"image\": image, \"target\": target}\n\n    def __len__(self) -> int:\n        return len(self.df)","metadata":{"papermill":{"duration":0.069144,"end_time":"2022-03-20T06:54:50.235457","exception":false,"start_time":"2022-03-20T06:54:50.166313","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:58.062753Z","iopub.execute_input":"2022-04-19T21:51:58.06318Z","iopub.status.idle":"2022-04-19T21:51:58.073629Z","shell.execute_reply.started":"2022-04-19T21:51:58.063141Z","shell.execute_reply":"2022-04-19T21:51:58.07293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning DataModule","metadata":{"papermill":{"duration":0.056953,"end_time":"2022-03-20T06:54:50.349928","exception":false,"start_time":"2022-03-20T06:54:50.292975","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class LitDataModule(pl.LightningDataModule):\n    def __init__(\n        self,\n        train_csv_encoded_folded: str,\n        test_csv: str,\n        val_fold: float,\n        image_size: int,\n        batch_size: int,\n        num_workers: int,\n    ):\n        super().__init__()\n\n        self.save_hyperparameters()\n\n        self.train_df = pd.read_csv(train_csv_encoded_folded)\n        self.test_df = pd.read_csv(test_csv)\n        \n        self.transform = create_transform(\n            input_size=(self.hparams.image_size, self.hparams.image_size),\n            crop_pct=1.0,\n        )\n        \n    def setup(self, stage: Optional[str] = None):\n        if stage == \"fit\" or stage is None:\n            # Split train df using fold\n            train_df = self.train_df[self.train_df.kfold != self.hparams.val_fold].reset_index(drop=True)\n            val_df = self.train_df[self.train_df.kfold == self.hparams.val_fold].reset_index(drop=True)\n\n            self.train_dataset = HappyWhaleDataset(train_df, transform=self.transform)\n            self.val_dataset = HappyWhaleDataset(val_df, transform=self.transform)\n\n        if stage == \"test\" or stage is None:\n            self.test_dataset = HappyWhaleDataset(self.test_df, transform=self.transform)\n\n    def train_dataloader(self) -> DataLoader:\n        return self._dataloader(self.train_dataset, train=True)\n\n    def val_dataloader(self) -> DataLoader:\n        return self._dataloader(self.val_dataset)\n\n    def test_dataloader(self) -> DataLoader:\n        return self._dataloader(self.test_dataset)\n\n    def _dataloader(self, dataset: HappyWhaleDataset, train: bool = False) -> DataLoader:\n        return DataLoader(\n            dataset,\n            batch_size=self.hparams.batch_size,\n            shuffle=train,\n            num_workers=self.hparams.num_workers,\n            pin_memory=True,\n            drop_last=train,\n        )","metadata":{"papermill":{"duration":0.072237,"end_time":"2022-03-20T06:54:50.479613","exception":false,"start_time":"2022-03-20T06:54:50.407376","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:58.075152Z","iopub.execute_input":"2022-04-19T21:51:58.075457Z","iopub.status.idle":"2022-04-19T21:51:58.08951Z","shell.execute_reply.started":"2022-04-19T21:51:58.075416Z","shell.execute_reply":"2022-04-19T21:51:58.088798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ArcMargin","metadata":{"papermill":{"duration":0.057828,"end_time":"2022-03-20T06:54:50.594798","exception":false,"start_time":"2022-03-20T06:54:50.53697","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# From https://github.com/lyakaap/Landmark2019-1st-and-3rd-Place-Solution/blob/master/src/modeling/metric_learning.py\n# Added type annotations, device, and 16bit support\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\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        s: float,\n        m: float,\n        easy_margin: bool,\n        ls_eps: float,\n    ):\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: torch.Tensor, label: torch.Tensor, device: str = \"cuda\") -> torch.Tensor:\n        # --------------------------- cos(theta) & phi(theta) ---------------------\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        # Enable 16 bit precision\n        cosine = cosine.to(torch.float32)\n\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=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":{"papermill":{"duration":0.073932,"end_time":"2022-03-20T06:54:50.725623","exception":false,"start_time":"2022-03-20T06:54:50.651691","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:58.091828Z","iopub.execute_input":"2022-04-19T21:51:58.092288Z","iopub.status.idle":"2022-04-19T21:51:58.106222Z","shell.execute_reply.started":"2022-04-19T21:51:58.092252Z","shell.execute_reply":"2022-04-19T21:51:58.105525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning Module","metadata":{"papermill":{"duration":0.081819,"end_time":"2022-03-20T06:54:50.865607","exception":false,"start_time":"2022-03-20T06:54:50.783788","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class LitModule(pl.LightningModule):\n    def __init__(\n        self,\n        model_name: str,\n        pretrained: bool,\n        drop_rate: float,\n        embedding_size: int,\n        num_classes: int,\n        arc_s: float,\n        arc_m: float,\n        arc_easy_margin: bool,\n        arc_ls_eps: float,\n        optimizer: str,\n        learning_rate: float,\n        weight_decay: float,\n        len_train_dl: int,\n        epochs:int\n    ):\n        super().__init__()\n\n        self.save_hyperparameters()\n\n        self.model = timm.create_model(model_name, pretrained=pretrained, drop_rate=drop_rate)\n        self.embedding = nn.Linear(self.model.get_classifier().in_features, embedding_size)\n        self.model.reset_classifier(num_classes=0, global_pool=\"avg\")\n\n        self.arc = ArcMarginProduct(\n            in_features=embedding_size,\n            out_features=num_classes,\n            s=arc_s,\n            m=arc_m,\n            easy_margin=arc_easy_margin,\n            ls_eps=arc_ls_eps,\n        )\n\n#         self.loss_fn = F.cross_entropy\n        self.loss_fn = FocalLoss()\n\n    def forward(self, images: torch.Tensor) -> torch.Tensor:\n        features = self.model(images)\n        embeddings = self.embedding(features)\n\n        return embeddings\n\n    def configure_optimizers(self):\n        optimizer = create_optimizer_v2(\n            self.parameters(),\n            opt=self.hparams.optimizer,\n            lr=self.hparams.learning_rate,\n            weight_decay=self.hparams.weight_decay,\n        )\n        \n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer,\n            self.hparams.learning_rate,\n            steps_per_epoch=self.hparams.len_train_dl,\n            epochs=self.hparams.epochs,\n        )\n        scheduler = {\"scheduler\": scheduler, \"interval\": \"step\"}\n\n        return [optimizer], [scheduler]\n\n    def training_step(self, batch: Dict[str, torch.Tensor], batch_idx: int) -> torch.Tensor:\n        return self._step(batch, \"train\")\n\n    def validation_step(self, batch: Dict[str, torch.Tensor], batch_idx: int) -> torch.Tensor:\n        return self._step(batch, \"val\")\n\n    def _step(self, batch: Dict[str, torch.Tensor], step: str) -> torch.Tensor:\n        images, targets = batch[\"image\"], batch[\"target\"]\n\n        embeddings = self(images)\n        outputs = self.arc(embeddings, targets, self.device)\n\n        loss = self.loss_fn(outputs, targets)\n        \n        self.log(f\"{step}_loss\", loss)\n\n        return loss","metadata":{"papermill":{"duration":0.13687,"end_time":"2022-03-20T06:54:51.13308","exception":false,"start_time":"2022-03-20T06:54:50.99621","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:58.107448Z","iopub.execute_input":"2022-04-19T21:51:58.107749Z","iopub.status.idle":"2022-04-19T21:51:58.122939Z","shell.execute_reply.started":"2022-04-19T21:51:58.107715Z","shell.execute_reply":"2022-04-19T21:51:58.122072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"papermill":{"duration":0.098053,"end_time":"2022-03-20T06:54:51.328839","exception":false,"start_time":"2022-03-20T06:54:51.230786","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train(\n    train_csv_encoded_folded: str = str(TRAIN_CSV_ENCODED_FOLDED_PATH),\n    test_csv: str = str(TEST_CSV_PATH),\n    val_fold: float = 0.0,\n    image_size: int = 256,\n    batch_size: int = 64,\n    num_workers: int = 2,\n    model_name: str = \"tf_efficientnet_b0\",\n    pretrained: bool = True,\n    drop_rate: float = 0.0,\n    embedding_size: int = 512,\n    num_classes: int = 15587,\n    arc_s: float = 30.0,\n    arc_m: float = 0.5,\n    arc_easy_margin: bool = False,\n    arc_ls_eps: float = 0.0,\n    optimizer: str = \"adam\",\n    learning_rate: float = 3e-4,\n    weight_decay: float = 1e-6,\n    checkpoints_dir: str = str(CHECKPOINTS_DIR),\n    accumulate_grad_batches: int = 1,\n    auto_lr_find: bool = False,\n    auto_scale_batch_size: bool = False,\n    fast_dev_run: bool = False,\n    gpus: int = 1,\n    max_epochs: int = 10,\n    precision: int = 16,\n    stochastic_weight_avg: bool = True,\n):\n    pl.seed_everything(42)\n\n    datamodule = LitDataModule(\n        train_csv_encoded_folded=train_csv_encoded_folded,\n        test_csv=test_csv,\n        val_fold=val_fold,\n        image_size=image_size,\n        batch_size=batch_size,\n        num_workers=num_workers,\n    )\n    \n    datamodule.setup()\n    len_train_dl = len(datamodule.train_dataloader())\n\n    module = LitModule(\n        model_name=model_name,\n        pretrained=pretrained,\n        drop_rate=drop_rate,\n        embedding_size=embedding_size,\n        num_classes=num_classes,\n        arc_s=arc_s,\n        arc_m=arc_m,\n        arc_easy_margin=arc_easy_margin,\n        arc_ls_eps=arc_ls_eps,\n        optimizer=optimizer,\n        learning_rate=learning_rate,\n        weight_decay=weight_decay,\n        len_train_dl=len_train_dl,\n        epochs=max_epochs\n    )\n    \n    model_checkpoint = ModelCheckpoint(\n        checkpoints_dir,\n        filename=f\"{model_name}_{image_size}\",\n        monitor=\"val_loss\",\n    )\n        \n    trainer = pl.Trainer(\n        accumulate_grad_batches=accumulate_grad_batches,\n        auto_lr_find=auto_lr_find,\n        auto_scale_batch_size=auto_scale_batch_size,\n        benchmark=True,\n        callbacks=[model_checkpoint],\n        deterministic=True,\n        fast_dev_run=fast_dev_run,\n        gpus=gpus,\n        max_epochs=2 if DEBUG else max_epochs,\n        precision=precision,\n        stochastic_weight_avg=stochastic_weight_avg,\n        limit_train_batches=0.1 if DEBUG else 1.0,\n        limit_val_batches=0.1 if DEBUG else 1.0,\n    )\n\n    trainer.tune(module, datamodule=datamodule)\n\n    trainer.fit(module, datamodule=datamodule)","metadata":{"papermill":{"duration":0.121308,"end_time":"2022-03-20T06:54:51.548926","exception":false,"start_time":"2022-03-20T06:54:51.427618","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-19T21:51:58.124404Z","iopub.execute_input":"2022-04-19T21:51:58.124839Z","iopub.status.idle":"2022-04-19T21:51:58.139716Z","shell.execute_reply.started":"2022-04-19T21:51:58.124804Z","shell.execute_reply":"2022-04-19T21:51:58.13896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name = \"tf_efficientnet_b5\"\nimage_size = 512\nbatch_size = 16\n\ntrain(model_name=model_name,\n      image_size=image_size,\n      batch_size=batch_size)","metadata":{"papermill":{"duration":23192.348927,"end_time":"2022-03-20T13:21:23.956569","exception":false,"start_time":"2022-03-20T06:54:51.607642","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-04-20T01:29:13.515177Z","iopub.execute_input":"2022-04-20T01:29:13.516046Z","iopub.status.idle":"2022-04-20T01:29:13.604666Z","shell.execute_reply.started":"2022-04-20T01:29:13.515917Z","shell.execute_reply":"2022-04-20T01:29:13.603723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{"papermill":{"duration":0.063566,"end_time":"2022-03-20T13:21:24.086693","exception":false,"start_time":"2022-03-20T13:21:24.023127","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def load_eval_module(checkpoint_path: str, device: torch.device) -> LitModule:\n    module = LitModule.load_from_checkpoint(checkpoint_path)\n    module.to(device)\n    module.eval()\n\n    return module\n\ndef load_dataloaders(\n    train_csv_encoded_folded: str,\n    test_csv: str,\n    val_fold: float,\n    image_size: int,\n    batch_size: int,\n    num_workers: int,\n) -> Tuple[DataLoader, DataLoader, DataLoader]:\n\n    datamodule = LitDataModule(\n        train_csv_encoded_folded=train_csv_encoded_folded,\n        test_csv=test_csv,\n        val_fold=val_fold,\n        image_size=image_size,\n        batch_size=batch_size,\n        num_workers=num_workers,\n    )\n\n    datamodule.setup()\n\n    train_dl = datamodule.train_dataloader()\n    val_dl = datamodule.val_dataloader()\n    test_dl = datamodule.test_dataloader()\n\n    return train_dl, val_dl, test_dl\n\n\ndef load_encoder() -> LabelEncoder:\n    encoder = LabelEncoder()\n    encoder.classes_ = np.load(ENCODER_CLASSES_PATH, allow_pickle=True)\n\n    return encoder\n\n\n@torch.inference_mode()\ndef get_embeddings(\n    module: pl.LightningModule, dataloader: DataLoader, encoder: LabelEncoder, stage: str\n) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:\n\n    all_image_names = []\n    all_embeddings = []\n    all_targets = []\n\n    for batch in tqdm(dataloader, desc=f\"Creating {stage} embeddings\"):\n        image_names = batch[\"image_name\"]\n        images = batch[\"image\"].to(module.device)\n        targets = batch[\"target\"].to(module.device)\n\n        embeddings = module(images)\n\n        all_image_names.append(image_names)\n        all_embeddings.append(embeddings.cpu().numpy())\n        all_targets.append(targets.cpu().numpy())\n        \n        if DEBUG:\n            break\n\n    all_image_names = np.concatenate(all_image_names)\n    all_embeddings = np.vstack(all_embeddings)\n    all_targets = np.concatenate(all_targets)\n\n    all_embeddings = normalize(all_embeddings, axis=1, norm=\"l2\")\n    all_targets = encoder.inverse_transform(all_targets)\n\n    return all_image_names, all_embeddings, all_targets\n\n\ndef create_and_search_index(embedding_size: int, train_embeddings: np.ndarray, val_embeddings: np.ndarray, k: int):\n    index = faiss.IndexFlatIP(embedding_size)\n    index.add(train_embeddings)\n    D, I = index.search(val_embeddings, k=k)  # noqa: E741\n\n    return D, I\n\n\ndef create_val_targets_df(\n    train_targets: np.ndarray, val_image_names: np.ndarray, val_targets: np.ndarray\n) -> pd.DataFrame:\n\n    allowed_targets = np.unique(train_targets)\n    val_targets_df = pd.DataFrame(np.stack([val_image_names, val_targets], axis=1), columns=[\"image\", \"target\"])\n    val_targets_df.loc[~val_targets_df.target.isin(allowed_targets), \"target\"] = \"new_individual\"\n\n    return val_targets_df\n\n\ndef create_distances_df(\n    image_names: np.ndarray, targets: np.ndarray, D: np.ndarray, I: np.ndarray, stage: str  # noqa: E741\n) -> pd.DataFrame:\n\n    distances_df = []\n    for i, image_name in tqdm(enumerate(image_names), desc=f\"Creating {stage}_df\"):\n        target = targets[I[i]]\n        distances = D[i]\n        subset_preds = pd.DataFrame(np.stack([target, distances], axis=1), columns=[\"target\", \"distances\"])\n        subset_preds[\"image\"] = image_name\n        distances_df.append(subset_preds)\n\n    distances_df = pd.concat(distances_df).reset_index(drop=True)\n    distances_df = distances_df.groupby([\"image\", \"target\"]).distances.max().reset_index()\n    distances_df = distances_df.sort_values(\"distances\", ascending=False).reset_index(drop=True)\n\n    return distances_df\n\n\ndef get_best_threshold(val_targets_df: pd.DataFrame, valid_df: pd.DataFrame) -> Tuple[float, float]:\n    best_th = 0\n    best_cv = 0\n    for th in [0.1 * x for x in range(11)]:\n        all_preds = get_predictions(valid_df, threshold=th)\n\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\n        cv = val_targets_df[th].mean()\n\n        print(f\"th={th} cv={cv}\")\n\n        if cv > best_cv:\n            best_th = th\n            best_cv = cv\n\n    print(f\"best_th={best_th}\")\n    print(f\"best_cv={best_cv}\")\n\n    # Adjustment: Since Public lb has nearly 10% 'new_individual' (Be Careful for private LB)\n    val_targets_df[\"is_new_individual\"] = val_targets_df.target == \"new_individual\"\n    val_scores = val_targets_df.groupby(\"is_new_individual\").mean().T\n    val_scores[\"adjusted_cv\"] = val_scores[True] * 0.1 + val_scores[False] * 0.9\n    best_th = val_scores[\"adjusted_cv\"].idxmax()\n    print(f\"best_th_adjusted={best_th}\")\n\n    return best_th, best_cv\n\n\ndef get_predictions(df: pd.DataFrame, threshold: float = 0.2):\n    sample_list = [\"938b7e931166\", \"5bf17305f073\", \"7593d2aee842\", \"7362d7a01d00\", \"956562ff2888\"]\n\n    predictions = {}\n    for i, row in tqdm(df.iterrows(), total=len(df), desc=f\"Creating predictions for threshold={threshold}\"):\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\n\n\n# TODO: add types\ndef 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\n\n\ndef create_predictions_df(test_df: pd.DataFrame, best_th: float) -> pd.DataFrame:\n    predictions = get_predictions(test_df, best_th)\n\n    predictions = pd.Series(predictions).reset_index()\n    predictions.columns = [\"image\", \"predictions\"]\n    predictions[\"predictions\"] = predictions[\"predictions\"].apply(lambda x: \" \".join(x))\n\n    return predictions","metadata":{"papermill":{"duration":0.37178,"end_time":"2022-03-20T13:21:24.522375","exception":false,"start_time":"2022-03-20T13:21:24.150595","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer(\n    checkpoint_path: str,\n    train_csv_encoded_folded: str = str(TRAIN_CSV_ENCODED_FOLDED_PATH),\n    test_csv: str = str(TEST_CSV_PATH),\n    val_fold: float = 0.0,\n    image_size: int = 256,\n    batch_size: int = 64,\n    num_workers: int = 2,\n    k: int = 50,\n):\n    module = load_eval_module(checkpoint_path, torch.device(\"cuda\"))\n\n    train_dl, val_dl, test_dl = load_dataloaders(\n        train_csv_encoded_folded=train_csv_encoded_folded,\n        test_csv=test_csv,\n        val_fold=val_fold,\n        image_size=image_size,\n        batch_size=batch_size,\n        num_workers=num_workers,\n    )\n\n    encoder = load_encoder()\n\n    train_image_names, train_embeddings, train_targets = get_embeddings(module, train_dl, encoder, stage=\"train\")\n    val_image_names, val_embeddings, val_targets = get_embeddings(module, val_dl, encoder, stage=\"val\")\n    test_image_names, test_embeddings, test_targets = get_embeddings(module, test_dl, encoder, stage=\"test\")\n\n    D, I = create_and_search_index(module.hparams.embedding_size, train_embeddings, val_embeddings, k)  # noqa: E741\n    print(\"Created index with train_embeddings\")\n\n    val_targets_df = create_val_targets_df(train_targets, val_image_names, val_targets)\n    print(f\"val_targets_df=\\n{val_targets_df.head()}\")\n\n    val_df = create_distances_df(val_image_names, train_targets, D, I, \"val\")\n    print(f\"val_df=\\n{val_df.head()}\")\n\n    best_th, best_cv = get_best_threshold(val_targets_df, val_df)\n    print(f\"val_targets_df=\\n{val_targets_df.describe()}\")\n\n    train_embeddings = np.concatenate([train_embeddings, val_embeddings])\n    train_targets = np.concatenate([train_targets, val_targets])\n    print(\"Updated train_embeddings and train_targets with val data\")\n\n    D, I = create_and_search_index(module.hparams.embedding_size, train_embeddings, test_embeddings, k)  # noqa: E741\n    print(\"Created index with train_embeddings\")\n\n    test_df = create_distances_df(test_image_names, train_targets, D, I, \"test\")\n    print(f\"test_df=\\n{test_df.head()}\")\n\n    predictions = create_predictions_df(test_df, best_th)\n    print(f\"predictions.head()={predictions.head()}\")\n    \n    # Fix missing predictions\n    # From https://www.kaggle.com/code/jpbremer/backfins-arcface-tpu-effnet/notebook\n    public_predictions = pd.read_csv(PUBLIC_SUBMISSION_CSV_PATH)\n    ids_without_backfin = np.load(IDS_WITHOUT_BACKFIN_PATH, allow_pickle=True)\n\n    ids2 = public_predictions[\"image\"][~public_predictions[\"image\"].isin(predictions[\"image\"])]\n\n    predictions = pd.concat(\n        [\n            predictions[~(predictions[\"image\"].isin(ids_without_backfin))],\n            public_predictions[public_predictions[\"image\"].isin(ids_without_backfin)],\n            public_predictions[public_predictions[\"image\"].isin(ids2)],\n        ]\n    )\n    predictions = predictions.drop_duplicates()\n\n    predictions.to_csv(SUBMISSION_CSV_PATH, index=False)\n    ","metadata":{"papermill":{"duration":0.08009,"end_time":"2022-03-20T13:21:24.667695","exception":false,"start_time":"2022-03-20T13:21:24.587605","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"infer(checkpoint_path=CHECKPOINTS_DIR / f\"{model_name}_{image_size}.ckpt\", image_size=image_size, batch_size=batch_size)","metadata":{"papermill":{"duration":1414.200016,"end_time":"2022-03-20T13:44:58.931841","exception":false,"start_time":"2022-03-20T13:21:24.731825","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]}]}