{"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] BirdCLEF Starter</h2>\n<br>","metadata":{}},{"cell_type":"markdown","source":"![](https://storage.googleapis.com/kaggle-media/competitions/Birdsong/Lifeclef.logo.jpg)","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":{"execution":{"iopub.status.busy":"2022-03-13T09:39:52.246178Z","iopub.execute_input":"2022-03-13T09:39:52.246532Z","iopub.status.idle":"2022-03-13T09:40:12.902446Z","shell.execute_reply.started":"2022-03-13T09:39:52.246441Z","shell.execute_reply":"2022-03-13T09:40:12.901676Z"},"_kg_hide-output":true,"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\nfrom torch.cuda import amp\n\n# Audio \nimport torchaudio\nfrom torchaudio.transforms import MelSpectrogram, Resample\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\nfrom sklearn.metrics import f1_score\n\n# For Image Models\nimport timm\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-03-13T09:40:12.906095Z","iopub.execute_input":"2022-03-13T09:40:12.906318Z","iopub.status.idle":"2022-03-13T09:40:16.165933Z","shell.execute_reply.started":"2022-03-13T09:40:12.906289Z","shell.execute_reply":"2022-03-13T09:40:16.165095Z"},"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-03-13T09:40:16.167564Z","iopub.execute_input":"2022-03-13T09:40:16.167828Z","iopub.status.idle":"2022-03-13T09:40:17.448684Z","shell.execute_reply.started":"2022-03-13T09:40:16.167792Z","shell.execute_reply":"2022-03-13T09:40:17.447985Z"},"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\": 10,\n          \"model_name\": \"tf_efficientnet_b0_ns\",\n          \"embedding_size\": 768,\n          \"num_classes\": 152,\n          \"train_batch_size\": 64,\n          \"valid_batch_size\": 128,\n          \"learning_rate\": 1e-4,\n          \"scheduler\": 'CosineAnnealingLR',\n          \"min_lr\": 1e-6,\n          \"T_max\": 500,\n          \"weight_decay\": 1e-6,\n          \"n_fold\": 5,\n          \"n_accumulate\": 1,\n          \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n          \"competition\": \"BirdCLEF22\",\n          \"_wandb_kernel\": \"deb\",\n          # Audio Specific\n          \"sample_rate\": 32_000,\n          \"max_time\": 5,\n          \"n_mels\": 224,\n          \"n_fft\": 1024,\n          }","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:40:17.449758Z","iopub.execute_input":"2022-03-13T09:40:17.449961Z","iopub.status.idle":"2022-03-13T09:40:17.499029Z","shell.execute_reply.started":"2022-03-13T09:40:17.449935Z","shell.execute_reply":"2022-03-13T09:40:17.498199Z"},"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-03-13T09:40:17.501903Z","iopub.execute_input":"2022-03-13T09:40:17.502627Z","iopub.status.idle":"2022-03-13T09:40:17.513104Z","shell.execute_reply.started":"2022-03-13T09:40:17.502587Z","shell.execute_reply":"2022-03-13T09:40:17.512331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '../input/birdclef-2022'\nTRAIN_DIR = '../input/birdclef-2022/train_audio'\nTEST_DIR = '../input/birdclef-2022/test_soundscapes'","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:40:17.515724Z","iopub.execute_input":"2022-03-13T09:40:17.516216Z","iopub.status.idle":"2022-03-13T09:40:17.520355Z","shell.execute_reply.started":"2022-03-13T09:40:17.516177Z","shell.execute_reply":"2022-03-13T09:40:17.519518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_file_path(filename):\n    return f\"{TRAIN_DIR}/{filename}\"","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:40:17.521974Z","iopub.execute_input":"2022-03-13T09:40:17.522663Z","iopub.status.idle":"2022-03-13T09:40:17.529548Z","shell.execute_reply.started":"2022-03-13T09:40:17.522617Z","shell.execute_reply":"2022-03-13T09:40:17.528833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <h1 style = \"font-family: garamond; font-size: 40px; font-style: normal; letter-spcaing: 3px; background-color: #f6f5f5; color :#fe346e; border-radius: 100px 100px; text-align:center\">Read the Data 📖</h1>","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f\"{ROOT_DIR}/train_metadata.csv\")\ndf['file_path'] = df['filename'].apply(get_train_file_path)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:40:17.531120Z","iopub.execute_input":"2022-03-13T09:40:17.531764Z","iopub.status.idle":"2022-03-13T09:40:17.661285Z","shell.execute_reply.started":"2022-03-13T09:40:17.531725Z","shell.execute_reply":"2022-03-13T09:40:17.660526Z"},"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\">Visualize Images</h1>","metadata":{}},{"cell_type":"code","source":"run = wandb.init(project=CONFIG['competition'],\n                 job_type='Visualization',\n                 name='Audio Visualization',\n                 anonymous='must')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:40:17.662598Z","iopub.execute_input":"2022-03-13T09:40:17.662864Z","iopub.status.idle":"2022-03-13T09:40:24.284168Z","shell.execute_reply.started":"2022-03-13T09:40:17.662830Z","shell.execute_reply":"2022-03-13T09:40:24.283462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preview_table = wandb.Table(columns=['Audio', 'Label', 'Rating', 'Time'])\n\ntemp_df = df.sample(5000).reset_index(drop=True)\n\nfor i in tqdm(range(len(temp_df))):\n    row = temp_df.loc[i]\n    audio = wandb.Audio(row.file_path, sample_rate=CONFIG['sample_rate'])\n    preview_table.add_data(audio,\n                           row.primary_label,\n                           row.rating,\n                           row.time)\n\nwandb.log({'Visualization': preview_table})\nrun.finish()","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:40:24.285839Z","iopub.execute_input":"2022-03-13T09:40:24.286092Z","iopub.status.idle":"2022-03-13T09:43:46.961622Z","shell.execute_reply.started":"2022-03-13T09:40:24.286056Z","shell.execute_reply":"2022-03-13T09:43:46.960983Z"},"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;\"><a href=\"https://wandb.ai/dchanda/BirdCLEF22/runs/2anb5iwu\">View the Complete Table Here ⮕</a></span>","metadata":{}},{"cell_type":"markdown","source":"![](https://i.imgur.com/riIHcZC.gif)","metadata":{}},{"cell_type":"code","source":"encoder = LabelEncoder()\ndf['primary_label'] = encoder.fit_transform(df['primary_label'])\n\nwith open(\"le.pkl\", \"wb\") as fp:\n    joblib.dump(encoder, fp)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:43:46.963112Z","iopub.execute_input":"2022-03-13T09:43:46.963370Z","iopub.status.idle":"2022-03-13T09:43:46.976773Z","shell.execute_reply.started":"2022-03-13T09:43:46.963335Z","shell.execute_reply":"2022-03-13T09:43:46.975381Z"},"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.primary_label)):\n      df.loc[val_ , \"kfold\"] = fold","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:43:46.978169Z","iopub.execute_input":"2022-03-13T09:43:46.978440Z","iopub.status.idle":"2022-03-13T09:43:46.996581Z","shell.execute_reply.started":"2022-03-13T09:43:46.978404Z","shell.execute_reply":"2022-03-13T09:43:46.995979Z"},"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 BirdCLEFDataset(Dataset):\n    def __init__(self, df, target_sample_rate, max_time, image_transforms=None):\n        self.file_paths = df['file_path'].values\n        self.labels = df['primary_label'].values\n        self.target_sample_rate = target_sample_rate\n        num_samples = target_sample_rate * max_time\n        self.num_samples = num_samples\n        self.image_transforms = image_transforms\n        \n    def __len__(self):\n        return len(self.file_paths)\n    \n    def __getitem__(self, index):\n        filepath = self.file_paths[index]\n        audio, sample_rate = torchaudio.load(filepath)\n        audio = self.to_mono(audio)\n        \n        if sample_rate != self.target_sample_rate:\n            resample = Resample(sample_rate, self.target_sample_rate)\n            audio = resample(audio)\n        \n        if audio.shape[0] > self.num_samples:\n            audio = self.crop_audio(audio)\n            \n        if audio.shape[0] < self.num_samples:\n            audio = self.pad_audio(audio)\n            \n        mel_spectogram = MelSpectrogram(sample_rate=self.target_sample_rate, \n                                        n_mels=CONFIG['n_mels'], \n                                        n_fft=CONFIG['n_fft'])\n        mel = mel_spectogram(audio)\n        label = torch.tensor(self.labels[index])\n        \n        # Convert to Image\n        image = torch.stack([mel, mel, mel])\n        \n        # Normalize Image\n        max_val = torch.abs(image).max()\n        image = image / max_val\n        \n        return {\n            \"image\": image, \n            \"label\": label\n        }\n            \n    def pad_audio(self, audio):\n        pad_length = self.num_samples - audio.shape[0]\n        last_dim_padding = (0, pad_length)\n        audio = F.pad(audio, last_dim_padding)\n        return audio\n        \n    def crop_audio(self, audio):\n        return audio[:self.num_samples]\n        \n    def to_mono(self, audio):\n        return torch.mean(audio, axis=0)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:43:46.997884Z","iopub.execute_input":"2022-03-13T09:43:46.998337Z","iopub.status.idle":"2022-03-13T09:43:47.010753Z","shell.execute_reply.started":"2022-03-13T09:43:46.998301Z","shell.execute_reply":"2022-03-13T09:43:47.009996Z"},"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-03-13T09:43:47.015065Z","iopub.execute_input":"2022-03-13T09:43:47.015815Z","iopub.status.idle":"2022-03-13T09:43:47.024237Z","shell.execute_reply.started":"2022-03-13T09:43:47.015778Z","shell.execute_reply":"2022-03-13T09:43:47.023528Z"},"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 BirdCLEFModel(nn.Module):\n    def __init__(self, model_name, embedding_size, pretrained=True):\n        super(BirdCLEFModel, 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.embedding = nn.Linear(in_features, embedding_size)\n        self.fc = nn.Linear(embedding_size, CONFIG['num_classes'])\n\n    def forward(self, images):\n        features = self.model(images)\n        pooled_features = self.pooling(features).flatten(1)\n        embedding = self.embedding(pooled_features)\n        output = self.fc(embedding)\n        return output\n    \nmodel = BirdCLEFModel(CONFIG['model_name'], CONFIG['embedding_size'])\nmodel.to(CONFIG['device']);","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:43:47.025483Z","iopub.execute_input":"2022-03-13T09:43:47.025753Z","iopub.status.idle":"2022-03-13T09:43:51.357479Z","shell.execute_reply.started":"2022-03-13T09:43:47.025718Z","shell.execute_reply":"2022-03-13T09:43:51.356747Z"},"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-03-13T09:43:51.360728Z","iopub.execute_input":"2022-03-13T09:43:51.360941Z","iopub.status.idle":"2022-03-13T09:43:51.365233Z","shell.execute_reply.started":"2022-03-13T09:43:51.360913Z","shell.execute_reply":"2022-03-13T09:43:51.364589Z"},"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 = model(images)\n        loss = criterion(outputs, labels)\n        loss = loss / CONFIG['n_accumulate']\n            \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-03-13T09:43:51.366714Z","iopub.execute_input":"2022-03-13T09:43:51.367312Z","iopub.status.idle":"2022-03-13T09:43:51.377827Z","shell.execute_reply.started":"2022-03-13T09:43:51.367226Z","shell.execute_reply":"2022-03-13T09:43:51.377080Z"},"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    LABELS = []\n    PREDS = []\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 = model(images)\n        _, preds = torch.max(outputs, 1)\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        PREDS.append(preds.view(-1).cpu().detach().numpy())\n        LABELS.append(labels.view(-1).cpu().detach().numpy())\n        \n        bar.set_postfix(Epoch=epoch, Valid_Loss=epoch_loss,\n                        LR=optimizer.param_groups[0]['lr'])   \n    \n    LABELS = np.concatenate(LABELS)\n    PREDS = np.concatenate(PREDS)\n    val_f1 = f1_score(LABELS, PREDS, average='macro')\n    gc.collect()\n    \n    return epoch_loss, val_f1","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:43:51.379235Z","iopub.execute_input":"2022-03-13T09:43:51.379556Z","iopub.status.idle":"2022-03-13T09:43:51.391071Z","shell.execute_reply.started":"2022-03-13T09:43:51.379519Z","shell.execute_reply":"2022-03-13T09:43:51.390448Z"},"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_f1 = 0\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, val_epoch_f1 = valid_one_epoch(model, valid_loader, \n                                                       device=CONFIG['device'], \n                                                       epoch=epoch)\n    \n        history['Train Loss'].append(train_epoch_loss)\n        history['Valid Loss'].append(val_epoch_loss)\n        history['Valid F1'].append(val_epoch_f1)\n        \n        # Log the metrics\n        wandb.log({\"Train Loss\": train_epoch_loss})\n        wandb.log({\"Valid Loss\": val_epoch_loss})\n        wandb.log({\"Valid F1\": val_epoch_f1})\n        \n        # deep copy the model\n        if val_epoch_f1 >= best_epoch_f1:\n            print(f\"{b_}Validation F1 Improved ({best_epoch_f1} ---> {val_epoch_f1})\")\n            best_epoch_f1 = val_epoch_f1\n            run.summary[\"Best F1 Score\"] = best_epoch_f1\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = \"F1{:.4f}_epoch{:.0f}.bin\".format(best_epoch_f1, 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 F1: {:.4f}\".format(best_epoch_f1))\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-03-13T09:43:51.392423Z","iopub.execute_input":"2022-03-13T09:43:51.392780Z","iopub.status.idle":"2022-03-13T09:43:51.406592Z","shell.execute_reply.started":"2022-03-13T09:43:51.392680Z","shell.execute_reply":"2022-03-13T09:43:51.405552Z"},"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    elif CONFIG['scheduler'] == None:\n        return None\n        \n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:43:51.408002Z","iopub.execute_input":"2022-03-13T09:43:51.408397Z","iopub.status.idle":"2022-03-13T09:43:51.418751Z","shell.execute_reply.started":"2022-03-13T09:43:51.408361Z","shell.execute_reply":"2022-03-13T09:43:51.417983Z"},"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 = BirdCLEFDataset(df_train, target_sample_rate=CONFIG['sample_rate'], max_time=CONFIG['max_time'])\n    valid_dataset = BirdCLEFDataset(df_valid, target_sample_rate=CONFIG['sample_rate'], max_time=CONFIG['max_time'])\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-03-13T09:43:51.420080Z","iopub.execute_input":"2022-03-13T09:43:51.420425Z","iopub.status.idle":"2022-03-13T09:43:51.429950Z","shell.execute_reply.started":"2022-03-13T09:43:51.420390Z","shell.execute_reply":"2022-03-13T09:43:51.429203Z"},"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-03-13T09:43:51.432443Z","iopub.execute_input":"2022-03-13T09:43:51.432907Z","iopub.status.idle":"2022-03-13T09:43:51.449088Z","shell.execute_reply.started":"2022-03-13T09:43:51.432876Z","shell.execute_reply":"2022-03-13T09:43:51.447901Z"},"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 Optimizer and Scheduler</span>","metadata":{}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=CONFIG['learning_rate'], \n                       weight_decay=CONFIG['weight_decay'])\nscheduler = fetch_scheduler(optimizer)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:43:51.451748Z","iopub.execute_input":"2022-03-13T09:43:51.452191Z","iopub.status.idle":"2022-03-13T09:43:51.458869Z","shell.execute_reply.started":"2022-03-13T09:43:51.452149Z","shell.execute_reply":"2022-03-13T09:43:51.458073Z"},"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=CONFIG['competition'], \n                 config=CONFIG,\n                 job_type='Train',\n                 tags=['gem-pooling', CONFIG['model_name']],\n                 anonymous='must')","metadata":{"execution":{"iopub.status.busy":"2022-03-13T09:43:51.460565Z","iopub.execute_input":"2022-03-13T09:43:51.460868Z","iopub.status.idle":"2022-03-13T09:43:57.894965Z","shell.execute_reply.started":"2022-03-13T09:43:51.460829Z","shell.execute_reply":"2022-03-13T09:43:57.894113Z"},"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-03-13T09:43:57.896336Z","iopub.execute_input":"2022-03-13T09:43:57.896594Z","iopub.status.idle":"2022-03-13T12:28:47.996001Z","shell.execute_reply.started":"2022-03-13T09:43:57.896556Z","shell.execute_reply":"2022-03-13T12:28:47.995089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run.finish()","metadata":{"execution":{"iopub.status.busy":"2022-03-13T12:28:47.999987Z","iopub.execute_input":"2022-03-13T12:28:48.000391Z","iopub.status.idle":"2022-03-13T12:28:53.322420Z","shell.execute_reply.started":"2022-03-13T12:28:48.000351Z","shell.execute_reply":"2022-03-13T12:28:53.321763Z"},"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/BirdCLEF22/runs/1tpuvu3f\">View the Complete Dashboard Here ⮕</a></span>","metadata":{}},{"cell_type":"markdown","source":"![img](https://i.imgur.com/IrGUTI7.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":{}}]}