{"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":"# About this notebook\n\n## Version 1\n- PyTorch Efficientnet_b0 starter code\n- 5 folds\n- 30 epochs\n- batch size 64 no accumulation\n- Custom learning scheduler\n- With data augmentation\n\n## Version 2\n- PyTorch Efficientnet_b0 starter code\n- 5 folds\n- 30 epochs\n- batch size 64 no accumulation\n- Custom learning scheduler\n- With data augmentation\n\n## Version 3\n- PyTorch Efficientnet_b0 starter code\n- 5 folds\n- 30 epochs\n- batch size 64 no accumulation\n- Custom learning scheduler\n- With data augmentation\n- CrossEntropyLoss\n- meta features 'Age' 'variety'\n\n## Version 4\n- PyTorch Efficientnet_b0 starter code\n- 5 folds\n- 40 epochs\n- **ArcFace**\n- batch size 64 no accumulation\n- CosineAnnealingLR scheduler\n- With data augmentation\n- CrossEntropyLoss\n- without meta features 'Age' 'variety'\n\n# Improvements maybe\n- Use ArcFace or add triplet loss with cross entropy for accuracy improvement\n- Use meta featues 'Age' 'variety' to improove the performance of the models\n- Use focal Loss (already implemented)\n\n# acknowledgement\n- Y.NAKAMA great [notebook](https://www.kaggle.com/yasufuminakama/herbarium-2020-pytorch-resnet18-train/notebook)\n\nIf this notebook is helpful, feel free to upvote :)","metadata":{}},{"cell_type":"code","source":"!pip install -q --upgrade wandb\n!pip install timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-29T01:03:26.127547Z","iopub.execute_input":"2022-07-29T01:03:26.128103Z","iopub.status.idle":"2022-07-29T01:03:49.376515Z","shell.execute_reply.started":"2022-07-29T01:03:26.127981Z","shell.execute_reply":"2022-07-29T01:03:49.375365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport cv2 as cv\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\n\npd.options.display.max_columns = 300","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:49.378780Z","iopub.execute_input":"2022-07-29T01:03:49.379641Z","iopub.status.idle":"2022-07-29T01:03:50.105670Z","shell.execute_reply.started":"2022-07-29T01:03:49.379595Z","shell.execute_reply":"2022-07-29T01:03:50.104729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Quick EDA","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv('../input/paddy-disease-classification/train.csv')\n\nsubmission = pd.read_csv('../input/paddy-disease-classification/sample_submission.csv')\ntrain_dir = '../input/paddy-disease-classification/train_images/'\ntest_dir = '../input/paddy-disease-classification/test_images'\n\n#train['path_jpeg'] = train['label'].apply(lambda x: os.path.join('../inµput/paddy-disease-classification/train_images',f'{x}'))\n#test['path_jpeg'] = test['dcm_name'].apply(lambda x: os.path.join('../input/jpeg-melanoma-256x256/test',f'{x}.jpg'))","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:50.107124Z","iopub.execute_input":"2022-07-29T01:03:50.107472Z","iopub.status.idle":"2022-07-29T01:03:50.141459Z","shell.execute_reply.started":"2022-07-29T01:03:50.107436Z","shell.execute_reply":"2022-07-29T01:03:50.140515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn import preprocessing\n\nle_label = preprocessing.LabelEncoder()\nle_label.fit(train_df['label'])\ntrain_df['label'] = le_label.transform(train_df['label'])","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:50.144194Z","iopub.execute_input":"2022-07-29T01:03:50.144588Z","iopub.status.idle":"2022-07-29T01:03:50.208637Z","shell.execute_reply.started":"2022-07-29T01:03:50.144550Z","shell.execute_reply":"2022-07-29T01:03:50.207780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"le_variety = preprocessing.LabelEncoder()\nle_variety.fit(train_df['variety'])\ntrain_df['variety'] = le_variety.transform(train_df['variety'])","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:50.209899Z","iopub.execute_input":"2022-07-29T01:03:50.210322Z","iopub.status.idle":"2022-07-29T01:03:50.218841Z","shell.execute_reply.started":"2022-07-29T01:03:50.210286Z","shell.execute_reply":"2022-07-29T01:03:50.217775Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:50.235586Z","iopub.execute_input":"2022-07-29T01:03:50.235827Z","iopub.status.idle":"2022-07-29T01:03:50.241735Z","shell.execute_reply.started":"2022-07-29T01:03:50.235805Z","shell.execute_reply":"2022-07-29T01:03:50.240792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    apex=False\n    debug=False\n    print_freq=100\n    size=256\n    num_workers=2\n    scheduler='CosineAnnealingLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts','OneCycleLR']\n    epochs=32\n    # CosineAnnealingLR params\n    cosanneal_params={\n        'T_max':10,\n        'eta_min':1e-4*0.5,\n        'last_epoch':-1\n    }\n    #ReduceLROnPlateau params\n    reduce_params={\n        'mode':'min',\n        'factor':0.1,\n        'patience':6,\n        'eps':1e-6,\n        'verbose':True\n    }\n    #ReduceLROnPlateau params for auc\n    reduce_params_for_auc={\n        'mode':'max',\n        'factor':0.1,\n        'patience':3,\n        'eps':1e-6,\n        'verbose':True\n    }\n    # CosineAnnealingWarmRestarts params\n    cosanneal_res_params={\n        'T_0':3,\n        'eta_min':1e-6,\n        'T_mult':1,\n        'last_epoch':-1\n    }\n    # OneCycleLR params\n    onecycle_params={\n        'pct_start':0.1,\n        'div_factor':1e1,\n        'max_lr':1e-3,\n        'steps_per_epoch':3, \n        'epochs':3\n    }\n    #batch_size=64\n    momentum=0.9\n    lr=1e-3\n    weight_decay=1e-4\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    nfolds=5\n    trn_folds=[0, 1, 2, 3, 4]\n    model_name='efficientnet_b0'     #'vit_base_patch32_224_in21k' 'tf_efficientnetv2_b0' 'resnext50_32x4d' 'resnet50d' 'efficientnet_b0'\n    preds_col = ['bacterial_leaf_blight', 'bacterial_leaf_streak',\n               'bacterial_panicle_blight', 'blast', 'brown_spot', 'dead_heart',\n               'downy_mildew', 'hispa', 'normal', 'tungro']\n    train=True\n    early_stop=True\n    target_size=len(preds_col)\n    scale=30.0\n    margin=0.50\n    easy_margin=False\n    ls_eps=0.0\n    fc_dim=512\n    early_stopping_steps=5\n    grad_cam=False\n    seed=42\n    image_size = 256\n    batch_size = 64\n    scale=30.0\n    margin=0.50\n    easy_margin=False\n    ls_eps=0.0\n    efficientnet_b0 = '../input/paddy-disease-starter-pytorch-arcface/'\n    tta=15\n    nfolds=5\n    trn_fold=[0, 1, 2, 3, 4]","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:50.243326Z","iopub.execute_input":"2022-07-29T01:03:50.243891Z","iopub.status.idle":"2022-07-29T01:03:50.257837Z","shell.execute_reply.started":"2022-07-29T01:03:50.243854Z","shell.execute_reply":"2022-07-29T01:03:50.256678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\nimport os\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nfrom torch.cuda import amp\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn import metrics\nfrom sklearn.metrics import roc_auc_score, roc_curve, f1_score, accuracy_score, log_loss\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\nfrom PIL import ImageFile\n# sometimes, you will have images without an ending bit\n# this takes care of those kind of (corrupt) images\nImageFile.LOAD_TRUNCATED_IMAGES = True\nimport albumentations as A \nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.optim.optimizer import Optimizer\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau, OneCycleLR\n\nimport albumentations as A\nfrom torchvision import transforms ,datasets\nfrom torchvision.utils import make_grid\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nfrom albumentations import (ToFloat, Normalize, VerticalFlip, HorizontalFlip, Compose, Resize,\n                            RandomBrightnessContrast, HueSaturationValue, Blur, GaussNoise,\n                            Rotate, RandomResizedCrop, Cutout, ShiftScaleRotate)\n\n\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nVERSION = 3","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:50.259460Z","iopub.execute_input":"2022-07-29T01:03:50.259857Z","iopub.status.idle":"2022-07-29T01:03:52.640962Z","shell.execute_reply.started":"2022-07-29T01:03:50.259821Z","shell.execute_reply":"2022-07-29T01:03:52.639801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# W&B","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandb_api = user_secrets.get_secret(\"wandb_key\")\n\nimport wandb\nwandb.login(key=wandb_api)\n\ndef class2dict(f):\n    return dict((name, getattr(f, name)) for name in dir(f) if not name.startswith('__'))\n\nrun = wandb.init(project=\"Paddy Doctor Competition\", \n                 name=f\"{CFG.model_name} batch size\",\n                 config=class2dict(CFG),\n                 group=CFG.model_name,\n                 job_type=\"train\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-29T01:03:52.646356Z","iopub.execute_input":"2022-07-29T01:03:52.647110Z","iopub.status.idle":"2022-07-29T01:03:56.329736Z","shell.execute_reply.started":"2022-07-29T01:03:52.647078Z","shell.execute_reply":"2022-07-29T01:03:56.328655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    score = log_loss(y_true, y_pred)\n    return score\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef set_seed(seed = 1234):\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    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    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed()","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:56.334211Z","iopub.execute_input":"2022-07-29T01:03:56.336922Z","iopub.status.idle":"2022-07-29T01:03:56.352072Z","shell.execute_reply.started":"2022-07-29T01:03:56.336857Z","shell.execute_reply":"2022-07-29T01:03:56.351141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV schem","metadata":{"_kg_hide-output":true}},{"cell_type":"code","source":"%%time\ntrain_df[\"fold\"] = -1\ntrain_df = train_df.sample(frac=1).reset_index(drop=True)\ny = train_df.label.values\nkf = StratifiedKFold(n_splits=5)\nfor f, (t_, v_) in enumerate(kf.split(X=train_df, y=y)):\n    train_df.loc[v_, 'fold'] = f","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:03:56.353741Z","iopub.execute_input":"2022-07-29T01:03:56.354474Z","iopub.status.idle":"2022-07-29T01:03:56.378641Z","shell.execute_reply.started":"2022-07-29T01:03:56.354436Z","shell.execute_reply":"2022-07-29T01:03:56.377571Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class TestDataset:\n    def __init__(self,dataframe, augmentations=None):\n\n        self.dataframe = dataframe\n        self.augmentations = augmentations\n    \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self, item):\n        image_path = os.path.join(test_dir, self.dataframe.iloc[item]['image_id'])\n        image = Image.open(image_path)\n        image = np.array(image)\n        \n        if self.augmentations is not None:\n            augmented = self.augmentations(image=image)\n            image = augmented[\"image\"]\n\n            \n        return image, torch.tensor(1)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:10:27.109535Z","iopub.execute_input":"2022-07-29T01:10:27.109919Z","iopub.status.idle":"2022-07-29T01:10:27.124018Z","shell.execute_reply.started":"2022-07-29T01:10:27.109886Z","shell.execute_reply":"2022-07-29T01:10:27.122722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"mean = (0.485, 0.456, 0.406)\nstd = (0.229, 0.224, 0.225)\n\ntrain_aug = A.Compose(\n    [\n            A.Resize(height=256, width=256),\n            A.ShiftScaleRotate(rotate_limit=90, scale_limit = [0.8, 1.2]),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.HueSaturationValue(sat_shift_limit=[0.7, 1.3], hue_shift_limit=[-0.1, 0.1]),\n\n            #A.RandomBrightnessContrast(brightness_limit=[0.7, 1.3],contrast_limit= [0.7, 1.3]),\n            A.Normalize(),\n            A.pytorch.transforms.ToTensorV2()\n    ]\n)\ntta_aug = A.Compose([\n            A.Resize(height=256, width=256),\n            A.ShiftScaleRotate(rotate_limit=90, scale_limit = [0.8, 1.2]),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.HueSaturationValue(sat_shift_limit=[0.7, 1.3], hue_shift_limit=[-0.1, 0.1]),\n            A.CenterCrop(224, 224),\n            #A.RandomBrightnessContrast(brightness_limit=[0.7, 1.3],contrast_limit= [0.7, 1.3]),\n            A.Normalize(),\n            A.pytorch.transforms.ToTensorV2()\n])\n\nvalid_aug = A.Compose(\n    [       A.Resize(height=256, width=256),   \n            A.Normalize(),\n            A.pytorch.transforms.ToTensorV2()\n    ]\n)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-29T01:10:28.272946Z","iopub.execute_input":"2022-07-29T01:10:28.273472Z","iopub.status.idle":"2022-07-29T01:10:28.289895Z","shell.execute_reply.started":"2022-07-29T01:10:28.273437Z","shell.execute_reply":"2022-07-29T01:10:28.288703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data):\n    \n    if data == 'tta':\n        return A.Compose(\n        [\n           A.Resize(CFG.size, CFG.size),\n           A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            A.Flip(p=0.5),\n            \n            #A.Cutout(p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Rotate(limit=180, p=0.5),\n            A.ShiftScaleRotate(\n                shift_limit = 0.1, scale_limit=0.1, rotate_limit=45, p=0.5\n            ),\n           \n            ToTensorV2(p=1.0),\n        ]\n    )\n\n    elif data == 'valid':\n        return A.Compose([\n            A.Resize(CFG.size, CFG.size),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:10:29.949518Z","iopub.execute_input":"2022-07-29T01:10:29.949858Z","iopub.status.idle":"2022-07-29T01:10:29.960152Z","shell.execute_reply.started":"2022-07-29T01:10:29.949831Z","shell.execute_reply":"2022-07-29T01:10:29.958882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = TestDataset(dataframe=submission,augmentations=get_transforms(data='tta'))\nfor i in range(5):\n    plt.figure(figsize=(4, 4))\n    image,_ = test_data[i]\n    plt.imshow(image.permute(2,1,0))\n    plt.show() ","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:11:15.213290Z","iopub.execute_input":"2022-07-29T01:11:15.213742Z","iopub.status.idle":"2022-07-29T01:11:16.682490Z","shell.execute_reply.started":"2022-07-29T01:11:15.213702Z","shell.execute_reply":"2022-07-29T01:11:16.681553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# src: https://amaarora.github.io/2020/08/30/gempool.html\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        # Applies 2D average-pooling operation in kH * kW regions by step size\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__ + '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:11:19.046863Z","iopub.execute_input":"2022-07-29T01:11:19.047201Z","iopub.status.idle":"2022-07-29T01:11:19.060688Z","shell.execute_reply.started":"2022-07-29T01:11:19.047171Z","shell.execute_reply":"2022-07-29T01:11:19.058413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# src: https://github.com/lyakaap/Landmark2019-1st-and-3rd-Place-Solution/blob/master/src/modeling/metric_learning.py\n\nclass ArcMarginProduct(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, \n                 m=0.50, easy_margin=False, ls_eps=0.0):\n        '''\n        in_features: dimension of the input\n        out_features: dimension of the last layer (in our case the classification)\n        s: norm of input feature\n        m: margin\n        ls_eps: label smoothing'''\n        \n        super(ArcMarginProduct, self).__init__()\n        self.in_features, self.out_features = in_features, out_features\n        self.s = s\n        self.m = m\n        self.ls_eps = ls_eps\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        # Fills the input `Tensor` with values according to the method described in\n        # `Understanding the difficulty of training deep feedforward neural networks`\n        # Glorot, X. & Bengio, Y. (2010)\n        # using a uniform distribution.\n        nn.init.xavier_uniform_(self.weight)\n\n        self.easy_margin = easy_margin\n        self.cos_m, self.sin_m = math.cos(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()).to(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-07-29T01:11:22.596854Z","iopub.execute_input":"2022-07-29T01:11:22.597220Z","iopub.status.idle":"2022-07-29T01:11:22.613222Z","shell.execute_reply.started":"2022-07-29T01:11:22.597189Z","shell.execute_reply":"2022-07-29T01:11:22.611922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HappyWhaleModel(nn.Module):\n    def __init__(self, modelName='efficientnet_b0', numClasses=CFG.target_size, noNeurons=250, embeddingSize=128):\n        super(HappyWhaleModel, self).__init__()\n        # Retrieve pretrained weights\n        self.backbone = timm.create_model(modelName, pretrained=True)\n        # Save the number features from the backbone\n        ### different models have different numbers e.g. EffnetB3 has 1536\n        backbone_features = self.backbone.classifier.in_features\n        self.backbone.classifier = nn.Identity() \n        self.backbone.global_pool = nn.Identity() \n        self.gem = GeM()\n        # Embedding layer (what we actually need)\n        self.embedding = nn.Sequential(nn.Linear(backbone_features, noNeurons),\n                                       nn.BatchNorm1d(noNeurons),\n                                       nn.ReLU(),\n                                       nn.Dropout(p=0.2),\n                                       \n                                       nn.Linear(noNeurons, embeddingSize),\n                                       nn.BatchNorm1d(embeddingSize),\n                                       nn.ReLU(),\n                                       nn.Dropout(p=0.2))\n        self.arcface = ArcMarginProduct(in_features=embeddingSize, \n                                        out_features=numClasses,\n                                        s=30.0, m=0.50, easy_margin=False, ls_eps=0.0)\n        \n        \n    def forward(self, image, target=None, prints=False):\n        '''If there is a target it means that the model is training on the dataset.\n        If there is no target, that means the model is predicting on the test dataset.\n        In this case we would skip the ArcFace layer and return only the image embeddings.\n        '''\n        \n        features = self.backbone(image)\n        # flatten transforms from e.g.: [3, 1536, 1, 1] to [3, 1536]\n        gem_pool = self.gem(features).flatten(1)\n        embedding = self.embedding(gem_pool)\n        if target != None:\n            out = self.arcface(embedding, target)\n        \n        if prints:\n            print(clr.S+\"0. IN:\", \"image shape:\"+clr.E, image.shape, \"target:\", target)\n            print(clr.S+\"1. Backbone Output:\"+clr.E, features.shape)\n            print(clr.S+\"2. GeM Pool Output:\"+clr.E, gem_pool.shape)\n            print(clr.S+\"3. Embedding Output:\"+clr.E, embedding.shape)\n            if target != None:\n                print(clr.S+\"4. ArcFace Output:\"+clr.E, out.shape)\n        \n        if target != None:\n            return out, embedding\n        else:\n            return embedding","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:11:30.226845Z","iopub.execute_input":"2022-07-29T01:11:30.227195Z","iopub.status.idle":"2022-07-29T01:11:30.244187Z","shell.execute_reply.started":"2022-07-29T01:11:30.227164Z","shell.execute_reply":"2022-07-29T01:11:30.242810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# inference\n# ====================================================\ndef inference(model, states, test_loader, device):\n    model.to(device)\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    for i, (images, labels) in tk0:\n        images = images.to(device).float()\n        labels = labels.to(device).float()\n        avg_confs = []\n        for state in states:\n            model.load_state_dict(state['model'])\n            model.eval()\n            with torch.no_grad():\n                y_preds, _ = model(images, labels)\n            avg_confs.append(y_preds.softmax(1).detach().to(\"cpu\").numpy())\n        avg_confs = np.mean(avg_confs, axis=0)\n        probs.append(avg_confs)\n    probs = np.concatenate(probs)\n    return probs","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:31:32.387711Z","iopub.execute_input":"2022-07-29T01:31:32.388118Z","iopub.status.idle":"2022-07-29T01:31:32.402482Z","shell.execute_reply.started":"2022-07-29T01:31:32.388081Z","shell.execute_reply":"2022-07-29T01:31:32.401523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{}},{"cell_type":"code","source":"efficientnet_b0 = HappyWhaleModel()\n\nstates = [torch.load(CFG.efficientnet_b0+f'{CFG.model_name}_fold{fold}_best_Accuracy.pth', map_location=device) for fold in CFG.trn_fold]\n\ntta_efficientnet_b0_predictions = np.zeros((submission.shape[0], CFG.target_size))\n\nif CFG.tta:\n    \n    for i in range(CFG.tta):\n        test_dataset = TestDataset(submission, augmentations=get_transforms(data='tta'))\n        test_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, \n                         num_workers=CFG.num_workers, pin_memory=True)\n        tta_efficientnet_b0_predictions += inference(efficientnet_b0, states, test_loader, device) / CFG.tta\n        \nelse:\n    \n    test_dataset = TestDataset(submission, augmentations=get_transforms(data='valid'))\n    test_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, \n                         num_workers=CFG.num_workers, pin_memory=True)\n    efficientnet_b0_predictions = inference(efficientnet_b0, states, test_loader, device)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:31:36.815875Z","iopub.execute_input":"2022-07-29T01:31:36.816215Z","iopub.status.idle":"2022-07-29T01:42:58.809387Z","shell.execute_reply.started":"2022-07-29T01:31:36.816184Z","shell.execute_reply":"2022-07-29T01:42:58.808391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.tta:\n    \n    submission['label'] = le_label.inverse_transform(tta_efficientnet_b0_predictions.argmax(1))\n    submission.to_csv(f'tta{CFG.tta}_{CFG.model_name}_version{VERSION}.csv', index=False)\n    \nelse:\n    submission['label'] = le_label.inverse_transform(tta_efficientnet_b0_predictions.argmax(1))\n    submission.to_csv(f'{CFG.model_name}_version{VERSION}.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:45:44.231702Z","iopub.execute_input":"2022-07-29T01:45:44.232095Z","iopub.status.idle":"2022-07-29T01:45:44.253343Z","shell.execute_reply.started":"2022-07-29T01:45:44.232064Z","shell.execute_reply":"2022-07-29T01:45:44.252150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv('./tta15_efficientnet_b0_version3.csv')\nsub","metadata":{"execution":{"iopub.status.busy":"2022-07-29T01:45:45.496683Z","iopub.execute_input":"2022-07-29T01:45:45.497725Z","iopub.status.idle":"2022-07-29T01:45:45.521019Z","shell.execute_reply.started":"2022-07-29T01:45:45.497674Z","shell.execute_reply":"2022-07-29T01:45:45.520028Z"},"trusted":true},"execution_count":null,"outputs":[]}]}