{"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;\">[Pytorch] ArcFace Starter</h2>\n<br>","metadata":{}},{"cell_type":"markdown","source":"<br>\n\nThis is the training notebook adapted for the inference done with [https://www.kaggle.com/vladvdv/pytorch-inference-notebok-arcface-gem-pooling](https://www.kaggle.com/vladvdv/pytorch-inference-notebok-arcface-gem-pooling) . It is based on the work of Debarshi Chanda https://www.kaggle.com/debarshichanda/pytorch-arcface-gem-pooling-starter  \n    \nModifications for version 8:\n* changed arhitecture to B4\n* included ArcFace into optimizer (small bug, thank for observing in the comments)\n* add functionality for half precision training (just for implementing on local machine, having trouble installing dependencies for nvidia apex to kaggle environemets)\n\nModifications for version 13:\n* changed arhitecture to B7 with noisy student weights\n* use fins training data\n* increase number of embeddings from 512 to 2048","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 --upgrade wandb","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-04-10T16:07:22.391716Z","iopub.execute_input":"2022-04-10T16:07:22.394735Z","iopub.status.idle":"2022-04-10T16:07:43.838340Z","shell.execute_reply.started":"2022-04-10T16:07:22.394611Z","shell.execute_reply":"2022-04-10T16:07:43.837481Z"},"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.optim as optim\nimport torch.nn.functional as F\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\n\n# Utils\nimport joblib\nfrom tqdm import tqdm\nfrom collections import defaultdict\n\n# Sklearn Imports\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold\n\n# For Image Models\nimport timm\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\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-04-10T16:07:43.841738Z","iopub.execute_input":"2022-04-10T16:07:43.841996Z","iopub.status.idle":"2022-04-10T16:07:48.180067Z","shell.execute_reply.started":"2022-04-10T16:07:43.841962Z","shell.execute_reply":"2022-04-10T16:07:48.179252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<img src=\"https://i.imgur.com/gb6B4ig.png\" width=\"400\" alt=\"Weights & Biases\" />\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\"> Weights & Biases (W&B) is a set of machine learning tools that helps you build better models faster. <strong>Kaggle competitions require fast-paced model development and evaluation</strong>. There are a lot of components: exploring the training data, training different models, combining trained models in different combinations (ensembling), and so on.</span>\n\n> <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">⏳ Lots of components = Lots of places to go wrong = Lots of time spent debugging</span>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">W&B can be useful for Kaggle competition with it's lightweight and interoperable tools:</span>\n\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Quickly track experiments,<br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Version and iterate on datasets, <br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Evaluate model performance,<br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Reproduce models,<br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Visualize results and spot regressions,<br></span>\n* <span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">Share findings with colleagues.</span>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.2em; font-weight: 300;\">To learn more about Weights and Biases check out this <strong><a href=\"https://www.kaggle.com/ayuraj/experiment-tracking-with-weights-and-biases\">kernel</a></strong>.</span>","metadata":{}},{"cell_type":"code","source":"import wandb\n\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret(\"wandb_api\")\n    wandb.login(key=api_key)\n    anony = None\nexcept:\n    anony = \"must\"\n    print('If you want to use your W&B account, go to Add-ons -> Secrets and provide your W&B access token. Use the Label name as wandb_api. \\nGet your W&B access token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:48.181826Z","iopub.execute_input":"2022-04-10T16:07:48.182108Z","iopub.status.idle":"2022-04-10T16:07:48.431217Z","shell.execute_reply.started":"2022-04-10T16:07:48.182063Z","shell.execute_reply":"2022-04-10T16:07:48.430489Z"},"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\">Training Configuration ⚙️</h1></span>","metadata":{}},{"cell_type":"code","source":"CONFIG = {\"seed\": 2022,\n          \"epochs\": 20,\n          \"img_size\": 512,\n          \"model_name\": \"tf_efficientnet_b7_ns\",\n          \"num_classes\": 15587,\n          \"train_batch_size\": 4,\n          \"valid_batch_size\": 4,\n          \"learning_rate\": 0.0001,\n          \"scheduler\": 'OneCycleLR',\n          \"min_lr\": 1e-6,\n          \"T_max\": 500,\n          \"weight_decay\": 1e-6,\n          \"n_fold\": 100, #use almost all data for training after determining the right arhitecture and proper parameters\n          \"n_accumulate\": 1,\n          \"device\": torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"),\n          \"test_mode\":True, # enable for testing pipeline, changes epochs to 2 and uses just 100 training samples\n          \"enable_amp_half_precision\": False, # Try it in your local machine (the code is made for working with !pip install apex, not the pytorch native apex)\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-04-10T16:07:48.433435Z","iopub.execute_input":"2022-04-10T16:07:48.433923Z","iopub.status.idle":"2022-04-10T16:07:48.480895Z","shell.execute_reply.started":"2022-04-10T16:07:48.433883Z","shell.execute_reply":"2022-04-10T16:07:48.480133Z"},"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-04-10T16:07:48.482708Z","iopub.execute_input":"2022-04-10T16:07:48.483269Z","iopub.status.idle":"2022-04-10T16:07:48.496121Z","shell.execute_reply.started":"2022-04-10T16:07:48.483230Z","shell.execute_reply":"2022-04-10T16:07:48.495312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '../input/happy-whale-and-dolphin'\nTEST_DIR = '../input/convert-backfintfrecords/happy-whale-and-dolphin-backfin/test_images'\nTRAIN_DIR = '../input/convert-backfintfrecords/happy-whale-and-dolphin-backfin/train_images'","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:48.498450Z","iopub.execute_input":"2022-04-10T16:07:48.499424Z","iopub.status.idle":"2022-04-10T16:07:48.504333Z","shell.execute_reply.started":"2022-04-10T16:07:48.499384Z","shell.execute_reply":"2022-04-10T16:07:48.503649Z"},"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-04-10T16:07:48.505677Z","iopub.execute_input":"2022-04-10T16:07:48.506131Z","iopub.status.idle":"2022-04-10T16:07:48.513255Z","shell.execute_reply.started":"2022-04-10T16:07:48.506096Z","shell.execute_reply":"2022-04-10T16:07:48.512440Z"},"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(\"../input/finsmetadata/train.csv\")\ndf['file_path'] = df['image'].apply(get_train_file_path)\ndf.head()\n\nif CONFIG[\"test_mode\"]==True:\n    df=df[:100]\n    CONFIG[\"epochs\"] = 2\n    CONFIG[\"n_fold\"] = 2\nencoder = LabelEncoder()\ndf['individual_id'] = encoder.fit_transform(df['individual_id'])\n\nwith open(\"le.pkl\", \"wb\") as fp:\n    joblib.dump(encoder, fp)","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:48.514663Z","iopub.execute_input":"2022-04-10T16:07:48.514917Z","iopub.status.idle":"2022-04-10T16:07:48.608399Z","shell.execute_reply.started":"2022-04-10T16:07:48.514884Z","shell.execute_reply":"2022-04-10T16:07:48.607643Z"},"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\">Create Folds</h1></span>","metadata":{}},{"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-04-10T16:07:48.609697Z","iopub.execute_input":"2022-04-10T16:07:48.610366Z","iopub.status.idle":"2022-04-10T16:07:48.620977Z","shell.execute_reply.started":"2022-04-10T16:07:48.610324Z","shell.execute_reply":"2022-04-10T16:07:48.620209Z"},"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.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        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        }","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:48.624378Z","iopub.execute_input":"2022-04-10T16:07:48.624567Z","iopub.status.idle":"2022-04-10T16:07:48.630893Z","shell.execute_reply.started":"2022-04-10T16:07:48.624543Z","shell.execute_reply":"2022-04-10T16:07:48.630246Z"},"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.HorizontalFlip(p=0.5),\n        A.HueSaturationValue(always_apply=False, hue_shift_limit=(-20, 20), sat_shift_limit=(-30, 30), val_shift_limit=(-20, 20), p=0.5),\n#        A.ShiftScaleRotate(always_apply=False, shift_limit=(-0.10999999940395355, 0.10999999940395355), scale_limit=(-0.5899999737739563, -0.28999999165534973), rotate_limit=(-26, 24), interpolation=0, border_mode=0, value=(0, 0, 0), mask_value=None, p=0.5),\n#        A.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.3, always_apply=False, p=0.5),\n        A.Blur(blur_limit=7, p=0.5),\n#        A.Cutout(num_holes=8, max_h_size=32, max_w_size=32, fill_value=0, always_apply=False, p=0.5),\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-04-10T16:07:48.632226Z","iopub.execute_input":"2022-04-10T16:07:48.632689Z","iopub.status.idle":"2022-04-10T16:07:48.642776Z","shell.execute_reply.started":"2022-04-10T16:07:48.632656Z","shell.execute_reply":"2022-04-10T16:07:48.641876Z"},"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>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Code taken from <a href=\"https://amaarora.github.io/2020/08/30/gempool.html\">GeM Pooling Explained</a></span>\n\n![](https://i.imgur.com/thTgYWG.jpg)","metadata":{}},{"cell_type":"code","source":"class 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-04-10T16:07:48.644204Z","iopub.execute_input":"2022-04-10T16:07:48.644499Z","iopub.status.idle":"2022-04-10T16:07:48.655701Z","shell.execute_reply.started":"2022-04-10T16:07:48.644460Z","shell.execute_reply":"2022-04-10T16:07:48.654807Z"},"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>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Code taken from <a href=\"https://github.com/lyakaap/Landmark2019-1st-and-3rd-Place-Solution/blob/master/src/modeling/metric_learning.py\">Landmark2019-1st-and-3rd-Place-Solution</a></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        if(CONFIG['enable_amp_half_precision']==True):\n            cosine = cosine.to(torch.float32)\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-04-10T16:07:48.657281Z","iopub.execute_input":"2022-04-10T16:07:48.657835Z","iopub.status.idle":"2022-04-10T16:07:48.673181Z","shell.execute_reply.started":"2022-04-10T16:07:48.657795Z","shell.execute_reply":"2022-04-10T16:07:48.672408Z"},"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\">Create Model</h1></span>","metadata":{}},{"cell_type":"code","source":"class HappyWhaleModel(nn.Module):\n    def __init__(self, model_name, 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.model.global_pool = nn.Identity()\n        self.pooling = GeM()\n        self.drop = nn.Dropout(p=0.2, inplace=False)\n        self.fc = nn.Linear(in_features,2048)\n        self.arc = ArcMarginProduct(2048, \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    def forward(self, images, labels):\n        features = self.model(images)\n        pooled_features = self.pooling(features).flatten(1)\n        pooled_drop = self.drop(pooled_features)\n        emb = self.fc(pooled_drop)\n        output = self.arc(emb,labels)\n        return output,emb\n    \n    \nmodel = HappyWhaleModel(CONFIG['model_name'])\nmodel.to(CONFIG['device']); \noptimizer = optim.Adam(model.parameters(), lr=CONFIG['learning_rate'], \n                       weight_decay=CONFIG['weight_decay'])\nif (CONFIG['enable_amp_half_precision']==True):\n    opt_level = 'O1'\n    model, optimizer = amp.initialize(model, optimizer, opt_level=opt_level)\n","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:48.674617Z","iopub.execute_input":"2022-04-10T16:07:48.675078Z","iopub.status.idle":"2022-04-10T16:07:55.038207Z","shell.execute_reply.started":"2022-04-10T16:07:48.674990Z","shell.execute_reply":"2022-04-10T16:07:55.037465Z"},"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\">Loss Function</h1></span>","metadata":{}},{"cell_type":"code","source":"def criterion(outputs, labels):\n    return nn.CrossEntropyLoss()(outputs, labels)","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:55.039348Z","iopub.execute_input":"2022-04-10T16:07:55.039928Z","iopub.status.idle":"2022-04-10T16:07:55.045668Z","shell.execute_reply.started":"2022-04-10T16:07:55.039891Z","shell.execute_reply":"2022-04-10T16:07:55.043622Z"},"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\">Training Function</h1></span>","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    \n    dataset_size = 0\n    running_loss = 0.0\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        \n        batch_size = images.size(0)\n        \n        outputs, emb = model(images, labels)\n        loss = criterion(outputs, labels)\n        loss = loss / CONFIG['n_accumulate']\n        if (CONFIG['enable_amp_half_precision']==True):\n            with amp.scale_loss(loss, optimizer) as scaled_loss:\n                scaled_loss.backward()\n        else:\n            loss.backward()\n    \n        if (step + 1) % CONFIG['n_accumulate'] == 0:\n            optimizer.step()\n\n            # zero the parameter gradients\n            optimizer.zero_grad()\n\n            if scheduler is not None:\n                scheduler.step()\n                \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        bar.set_postfix(Epoch=epoch, Train_Loss=epoch_loss,\n                        LR=optimizer.param_groups[0]['lr'])\n    gc.collect()\n    \n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:55.046885Z","iopub.execute_input":"2022-04-10T16:07:55.050938Z","iopub.status.idle":"2022-04-10T16:07:55.071257Z","shell.execute_reply.started":"2022-04-10T16:07:55.048749Z","shell.execute_reply":"2022-04-10T16:07:55.069987Z"},"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\">Validation Function</h1></span>","metadata":{}},{"cell_type":"code","source":"@torch.inference_mode()\ndef valid_one_epoch(model, dataloader, device, epoch):\n    model.eval()\n    \n    dataset_size = 0\n    running_loss = 0.0\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        \n        batch_size = images.size(0)\n\n        outputs, emb = model(images, labels)\n        loss = criterion(outputs, labels)\n        \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        bar.set_postfix(Epoch=epoch, Valid_Loss=epoch_loss,\n                        LR=optimizer.param_groups[0]['lr'])   \n    \n    gc.collect()\n    \n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:55.078126Z","iopub.execute_input":"2022-04-10T16:07:55.080723Z","iopub.status.idle":"2022-04-10T16:07:55.096209Z","shell.execute_reply.started":"2022-04-10T16:07:55.080681Z","shell.execute_reply":"2022-04-10T16:07:55.095342Z"},"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\">Run Training</h1></span>","metadata":{}},{"cell_type":"code","source":"def run_training(model, optimizer, scheduler, device, num_epochs):\n    # To automatically log gradients\n    wandb.watch(model, log_freq=100)\n    \n    if torch.cuda.is_available():\n        print(\"[INFO] Using GPU: {}\\n\".format(torch.cuda.get_device_name()))\n    \n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_epoch_loss = np.inf\n    history = defaultdict(list)\n    \n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        train_epoch_loss = train_one_epoch(model, optimizer, scheduler, \n                                           dataloader=train_loader, \n                                           device=CONFIG['device'], epoch=epoch)\n        \n        val_epoch_loss = valid_one_epoch(model, valid_loader, device=CONFIG['device'], \n                                         epoch=epoch)\n    \n        history['Train Loss'].append(train_epoch_loss)\n        history['Valid Loss'].append(val_epoch_loss)\n        \n        # Log the metrics\n        wandb.log({\"Train Loss\": train_epoch_loss})\n        wandb.log({\"Valid Loss\": val_epoch_loss})\n        wandb.log({\"LR\": optimizer.param_groups[0]['lr']})\n        \n        # deep copy the model\n        if val_epoch_loss <= best_epoch_loss:\n            print(f\"{b_}Validation Loss Improved ({best_epoch_loss} ---> {val_epoch_loss})\")\n            best_epoch_loss = val_epoch_loss\n            run.summary[\"Best Loss\"] = best_epoch_loss\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = \"Loss{:.4f}_epoch{:.0f}.bin\".format(best_epoch_loss, epoch)\n            torch.save(model.state_dict(), PATH)\n            # Save a model file from the current directory\n            print(f\"Model Saved{sr_}\")\n            \n        print()\n    \n    end = time.time()\n    time_elapsed = end - start\n    print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n        time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n    print(\"Best Loss: {:.4f}\".format(best_epoch_loss))\n    \n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:55.100997Z","iopub.execute_input":"2022-04-10T16:07:55.108100Z","iopub.status.idle":"2022-04-10T16:07:55.134598Z","shell.execute_reply.started":"2022-04-10T16:07:55.108058Z","shell.execute_reply":"2022-04-10T16:07:55.133481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fetch_scheduler(optimizer):\n    if CONFIG['scheduler'] == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer,T_max=CONFIG['T_max'], \n                                                   eta_min=CONFIG['min_lr'])\n    elif CONFIG['scheduler'] == 'CosineAnnealingWarmRestarts':\n        scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer,T_0=CONFIG['T_0'], \n                                                             eta_min=CONFIG['min_lr'])\n                \n    elif CONFIG['scheduler'] == 'OneCycleLR':\n        scheduler = lr_scheduler.OneCycleLR(optimizer,max_lr =CONFIG['learning_rate'],total_steps = CONFIG['epochs'] * len(train_loader))\n        \n    elif CONFIG['scheduler'] == None:\n        return None\n        \n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:55.138500Z","iopub.execute_input":"2022-04-10T16:07:55.143901Z","iopub.status.idle":"2022-04-10T16:07:55.151719Z","shell.execute_reply.started":"2022-04-10T16:07:55.143863Z","shell.execute_reply":"2022-04-10T16:07:55.150981Z"},"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=True, pin_memory=True, drop_last=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-04-10T16:07:55.153625Z","iopub.execute_input":"2022-04-10T16:07:55.154145Z","iopub.status.idle":"2022-04-10T16:07:55.168567Z","shell.execute_reply.started":"2022-04-10T16:07:55.154107Z","shell.execute_reply":"2022-04-10T16:07:55.167837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Prepare Dataloaders</span>","metadata":{}},{"cell_type":"code","source":"train_loader, valid_loader = prepare_loaders(df, fold=0)","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:55.173824Z","iopub.execute_input":"2022-04-10T16:07:55.174325Z","iopub.status.idle":"2022-04-10T16:07:55.196665Z","shell.execute_reply.started":"2022-04-10T16:07:55.174289Z","shell.execute_reply":"2022-04-10T16:07:55.196004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Define LR Scheduler</span>","metadata":{}},{"cell_type":"code","source":"scheduler = fetch_scheduler(optimizer)","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:55.198036Z","iopub.execute_input":"2022-04-10T16:07:55.198664Z","iopub.status.idle":"2022-04-10T16:07:55.202954Z","shell.execute_reply.started":"2022-04-10T16:07:55.198626Z","shell.execute_reply":"2022-04-10T16:07:55.202286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\">Start Training</span>","metadata":{}},{"cell_type":"code","source":"run = wandb.init(project='HappyWhale', \n                 config=CONFIG,\n                 job_type='Train',\n                 tags=['arcface', 'gem-pooling', 'effnet-b0', '448'],\n                 anonymous='must')","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:07:55.209476Z","iopub.execute_input":"2022-04-10T16:07:55.209941Z","iopub.status.idle":"2022-04-10T16:08:03.915474Z","shell.execute_reply.started":"2022-04-10T16:07:55.209907Z","shell.execute_reply":"2022-04-10T16:08:03.914728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model, history = run_training(model, optimizer, scheduler,\n                              device=CONFIG['device'],\n                              num_epochs=CONFIG['epochs'])","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:08:03.916783Z","iopub.execute_input":"2022-04-10T16:08:03.917041Z","iopub.status.idle":"2022-04-10T16:08:39.725165Z","shell.execute_reply.started":"2022-04-10T16:08:03.917004Z","shell.execute_reply":"2022-04-10T16:08:39.724254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run.finish()","metadata":{"execution":{"iopub.status.busy":"2022-04-10T16:08:39.726610Z","iopub.execute_input":"2022-04-10T16:08:39.727059Z","iopub.status.idle":"2022-04-10T16:08:43.896462Z","shell.execute_reply.started":"2022-04-10T16:08:39.727018Z","shell.execute_reply":"2022-04-10T16:08:43.895662Z"},"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\">Visualizations</h1>\n\n<span style=\"color: #000508; font-family: Segoe UI; font-size: 1.5em; font-weight: 300;\"><a href=\"https://wandb.ai/dchanda/HappyWhale/runs/3j25um1k\">View the Complete Dashboard Here ⮕</a></span>","metadata":{}},{"cell_type":"markdown","source":"![](https://i.imgur.com/zD3rD0W.jpg)","metadata":{}},{"cell_type":"markdown","source":"![Upvote!](https://img.shields.io/badge/Upvote-If%20you%20like%20my%20work-07b3c8?style=for-the-badge&logo=kaggle)","metadata":{}}]}