{"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":"# 1) Getting Setup!","metadata":{"id":"-a3epi32BS7E"}},{"cell_type":"code","source":"!pip install /kaggle/input/nnaudio031/nnAudio-0.3.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:03.906082Z","iopub.execute_input":"2023-05-11T15:22:03.906442Z","iopub.status.idle":"2023-05-11T15:22:16.663956Z","shell.execute_reply.started":"2023-05-11T15:22:03.906413Z","shell.execute_reply":"2023-05-11T15:22:16.662758Z"},"id":"ymj4uDciBS7I","outputId":"f79af700-c566-4a1d-b175-02cf5ceef636","collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport os","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:16.667775Z","iopub.execute_input":"2023-05-11T15:22:16.668839Z","iopub.status.idle":"2023-05-11T15:22:16.674091Z","shell.execute_reply.started":"2023-05-11T15:22:16.66878Z","shell.execute_reply":"2023-05-11T15:22:16.67309Z"},"id":"jzUU_zY3BS7L","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset, random_split\nimport torch.nn.functional as F\nimport torchaudio.transforms as T\nimport torchaudio\nfrom torchaudio import transforms\nfrom IPython.display import Audio\nimport torchvision\nfrom sklearn.preprocessing import OneHotEncoder, LabelEncoder\nfrom sklearn.model_selection import train_test_split\nimport sklearn\nimport numpy as np\nimport pandas as pd\nimport math, random\nfrom matplotlib import pyplot as plt\nimport librosa\nimport timm\nfrom tqdm.notebook import tqdm\nfrom nnAudio import features\nimport wandb\n%matplotlib inline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-11T15:22:16.675909Z","iopub.execute_input":"2023-05-11T15:22:16.676427Z","iopub.status.idle":"2023-05-11T15:22:22.790087Z","shell.execute_reply.started":"2023-05-11T15:22:16.67639Z","shell.execute_reply":"2023-05-11T15:22:22.789132Z"},"id":"_6X093D4BS7M","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed: int):\n    random.seed(seed)\n    os.environ['PYTHONHASHseed'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:22.792697Z","iopub.execute_input":"2023-05-11T15:22:22.793105Z","iopub.status.idle":"2023-05-11T15:22:22.799271Z","shell.execute_reply.started":"2023-05-11T15:22:22.79307Z","shell.execute_reply":"2023-05-11T15:22:22.798355Z"},"id":"cqLdMP4EBS7N","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dict_from_class(cls):\n    return dict((key, value) for (key, value) in cls.__dict__.items() if not \"__\" in key )","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:22.80055Z","iopub.execute_input":"2023-05-11T15:22:22.801436Z","iopub.status.idle":"2023-05-11T15:22:22.813025Z","shell.execute_reply.started":"2023-05-11T15:22:22.801402Z","shell.execute_reply":"2023-05-11T15:22:22.811945Z"},"id":"LH-fzSKrBS7N","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    # General configurations\n    seed = 14\n    device = \"cuda\"\n    batch_size = 32\n    base_path = \"/kaggle/input/birdclef-2023\"\n    \n    #WandB configurations\n    name = \"Overfitting\"\n    final_metric_name = \"Padded cMAP\"\n    running_metric_name = \"Micro AP\"\n    model_name = \"DenseNet121 PANN CQT Long + Data Augmentation\"\n    mode = \"maximize\"\n    \n    #Training Hyperparameters\n    epochs = 20\n    patience = 6\n    grad_accum = 1\n    n_folds = 5\n    lr = 5e-4\n    max_lr = 1e-2\n    optimizer = \"Adam\"\n    scheduler = \"OneCycleLR\"\n    weight_decay = 0.0\n    \n    \n    #Data Hyperparameters\n    sample_rate = 32000\n    duration = 15000\n    channels = 3\n    test_size = 0.2\n    num_classes = 264\n    n_mels = 128\n    n_bins = 84\n    \n    #Augmentation Hyperparameters\n    freq_mixup = 5\n    p_freq_block = 0.5\n    p_time_block = 0.5\n    p_audio_aug = 0.5","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:22.814269Z","iopub.execute_input":"2023-05-11T15:22:22.814677Z","iopub.status.idle":"2023-05-11T15:22:22.825888Z","shell.execute_reply.started":"2023-05-11T15:22:22.814646Z","shell.execute_reply":"2023-05-11T15:22:22.824509Z"},"id":"tguBW7YQBS7O","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(config.seed)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:22.829612Z","iopub.execute_input":"2023-05-11T15:22:22.831031Z","iopub.status.idle":"2023-05-11T15:22:22.841042Z","shell.execute_reply.started":"2023-05-11T15:22:22.831Z","shell.execute_reply":"2023-05-11T15:22:22.840078Z"},"id":"_KtN-h0BBS7P","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"wandb\")\n\n!wandb login $secret_value_0","metadata":{"id":"X4jVewp76BiV","outputId":"9c08332e-3d53-4b19-bce5-1e4740a2b031","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1.1) Trackers","metadata":{"id":"d_9rC77IBS7Q"}},{"cell_type":"code","source":"def padded_cmap(solution, submission, padding_factor=5):\n    new_rows = []\n    for i in range(padding_factor):\n        new_rows.append([1 for i in range(len(solution[0]))])\n    \n    padded_solution = np.concatenate((solution, new_rows), axis = 0)\n    padded_submission = np.concatenate((submission, new_rows), axis = 0)\n    \n    score = sklearn.metrics.average_precision_score(\n        padded_solution,\n        padded_submission,\n        average = \"macro\",\n    )\n    return score\n\ndef micro_ap_score(solution, submission):\n    solution = solution\n    submission = submission\n    score = sklearn.metrics.average_precision_score(\n        solution,\n        submission,\n        average='micro',\n    )\n    return score","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.519593Z","iopub.execute_input":"2023-05-11T15:22:27.520028Z","iopub.status.idle":"2023-05-11T15:22:27.532414Z","shell.execute_reply.started":"2023-05-11T15:22:27.519987Z","shell.execute_reply":"2023-05-11T15:22:27.531245Z"},"id":"_Fo0Unw3BS7Q","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AverageMeter:\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2023-05-11T15:22:27.538891Z","iopub.execute_input":"2023-05-11T15:22:27.539255Z","iopub.status.idle":"2023-05-11T15:22:27.547019Z","shell.execute_reply.started":"2023-05-11T15:22:27.53923Z","shell.execute_reply":"2023-05-11T15:22:27.545881Z"},"id":"xVOpvOIdBS7R","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MetricTracker():\n    def __init__(self, metric):\n        self.y_hat = []\n        self.y = []\n        self.metric = metric\n\n    def update(self, y_hat, y):\n        self.y_hat.extend(y_hat.detach().cpu().numpy())\n        self.y.extend(y.detach().cpu().numpy())\n    \n    def score(self):\n        return self.metric(self.y_hat, self.y)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.548211Z","iopub.execute_input":"2023-05-11T15:22:27.549543Z","iopub.status.idle":"2023-05-11T15:22:27.557657Z","shell.execute_reply.started":"2023-05-11T15:22:27.549503Z","shell.execute_reply":"2023-05-11T15:22:27.556645Z"},"id":"h_MHs0dVBS7R","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ModelTracker():\n    def __init__(self, model, path, optimizer, scheduler):\n        self.patience = config.patience\n        self.base_path = \"./\"\n        self.mode = config.mode\n        self.missed = 0\n        self.path = path\n        self.model = model\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n        self.metric = float(\"-inf\") if self.mode == \"maximize\" else float(\"inf\")\n        self.metric_name = config.final_metric_name\n        \n    def save_helper(self, epoch):\n        torch.save({\n                    \"epoch\": epoch, \n                    \"model_state_dict\": self.model.state_dict(), \n                    \"optimizer_state_dict\": self.optimizer.state_dict(),\n                    \"scheduler\": self.scheduler.state_dict()\n                }, f\"{self.base_path}/{self.path}\")\n\n        print(f\"Saved to model to {self.base_path}{self.path}!\")\n        \n    def save_model(self, epoch):\n        self.save_helper(epoch)\n        \n\n    def update(self, value, epoch):\n        if self.mode == \"maximize\":\n            if value >= self.metric:\n                print(f\"Validation {self.metric_name} rose from {self.metric:.4f} to {value:.4f} on epoch {epoch}\")\n                self.metric = value\n                self.save_model(epoch)    \n                self.missed = 0\n\n            else:\n                print(f\"Validation {self.metric_name} fell from {self.metric:.4f} to {value:.4f} on epoch {epoch}\")\n                print(f\"Model did not improve on epoch {epoch}\")\n                self.missed += 1\n        else:\n            if value <= self.metric:\n                print(f\"Validation {self.metric_name} fell from {self.metric:.4f} to {value:.4f} on epoch {epoch}\")\n                self.metric = value\n                self.save_model(epoch) \n                self.missed = 0\n\n            else:\n                print(f\"Validation {self.metric_name} rose from {self.metric:.4f} to {value:.4f} on epoch {epoch}\")\n                print(f\"Model did not improve on epoch {epoch}\")\n                self.missed += 1\n\n    def get_full_path(self):\n        return f\"{self.base_path}{self.path}\"\n        \n    def check_improvement(self):\n        return self.missed < self.patience","metadata":{"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-05-11T15:22:27.559272Z","iopub.execute_input":"2023-05-11T15:22:27.559662Z","iopub.status.idle":"2023-05-11T15:22:27.577037Z","shell.execute_reply.started":"2023-05-11T15:22:27.559628Z","shell.execute_reply":"2023-05-11T15:22:27.576018Z"},"id":"3RYFaYs9BS7S","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2) Data!","metadata":{"id":"FdX7HtAaBS7S"}},{"cell_type":"markdown","source":"### Since we need fast inference, simple stratified split is likely best.","metadata":{"id":"1Du1ztpnBS7T"}},{"cell_type":"code","source":"def birds_stratified_split(df):\n    class_counts = df[\"labels\"].value_counts()\n    low_count_classes = class_counts[class_counts < 2].index.tolist() ### Birds with single counts\n\n    df['train'] = df[\"labels\"].isin(low_count_classes)\n\n    train_df, val_df = train_test_split(df[~df['train']], test_size=config.test_size, stratify=df[~df['train']][\"labels\"], random_state=config.seed)\n\n    train_df = pd.concat([train_df, df[df['train']]], axis=0).reset_index(drop=True)\n\n    # Remove the 'valid' column\n    train_df.drop('train', axis=1, inplace=True)\n    val_df.drop('train', axis=1, inplace=True)\n\n    return train_df, val_df","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.580531Z","iopub.execute_input":"2023-05-11T15:22:27.580913Z","iopub.status.idle":"2023-05-11T15:22:27.59331Z","shell.execute_reply.started":"2023-05-11T15:22:27.580885Z","shell.execute_reply":"2023-05-11T15:22:27.592202Z"},"id":"P7QRYD63BS7T","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def noise_psd(N, psd = lambda f: 1):\n        X_white = np.fft.rfft(np.random.randn(N));\n        S = psd(np.fft.rfftfreq(N))\n        # Normalize S\n        S = S / np.sqrt(np.mean(S**2))\n        X_shaped = X_white * S;\n        return np.fft.irfft(X_shaped);\n\ndef PSDGenerator(f):\n    return lambda N: noise_psd(N, f)\n\ndef check_prob(p):\n    return np.random.uniform() <= p","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.594539Z","iopub.execute_input":"2023-05-11T15:22:27.595088Z","iopub.status.idle":"2023-05-11T15:22:27.609354Z","shell.execute_reply.started":"2023-05-11T15:22:27.595035Z","shell.execute_reply":"2023-05-11T15:22:27.608273Z"},"id":"Gp62vwLiBS7T","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AudioUtil():\n    @staticmethod\n    def open(file):\n        sig, sr = torchaudio.load(file)\n        return sig, sr\n    \n    @staticmethod\n    def rechannel(aud, new_channel):\n        sig, sr = aud\n        if (sig.shape[0] == new_channel):\n            return aud\n        \n        if (new_channel == 1):\n            return sig[0], sr\n        \n        else:\n            return torch.cat([sig, sig, sig]), sr\n        \n    @staticmethod\n    def resample(aud, newsr):\n        sig, sr = aud\n        if (sr == newsr):\n            return aud\n        \n        else:\n            return transforms.Resample(sample_rate, newsr)(sig), sr\n        \n    @staticmethod\n    def pad_trunc(aud, max_ms):\n        sig, sr = aud\n        num_rows, sig_len = sig.shape\n        max_len = max_ms // 1000 * sr\n        \n        if (sig_len > max_len):\n            sig = sig[:, :max_len]\n        \n        elif (sig_len < max_len):\n            right_pad_len = max_len - sig_len\n            right_padding = torch.zeros((num_rows, right_pad_len))\n            sig = torch.cat((sig, right_padding), 1)\n            \n        return sig, sr\n    \n    @staticmethod\n    def get_mel_spec(aud, n_mels=128, n_fft=1024, hop_len = None):\n        sig, sr = aud\n        top_db = 80\n        \n        spec = transforms.MelSpectrogram(sr, n_fft=n_fft, hop_length=hop_len, n_mels = n_mels)(sig)\n        spec = transforms.AmplitudeToDB(top_db=top_db)(spec)\n        \n        return spec\n            \n    @staticmethod\n    def show_spec(spec):\n        if spec.shape[0] > 1:\n            plt.imshow(spec[0], aspect = \"auto\")\n        else:\n            plt.imshow(spec, aspect = \"auto\")\n            \n    @staticmethod\n    @PSDGenerator\n    def white_noise(f):\n        return 1;\n    \n    @staticmethod\n    @PSDGenerator\n    def blue_noise(f):\n        return np.sqrt(f);\n    \n    @staticmethod\n    @PSDGenerator\n    def violet_noise(f):\n        return f;\n\n    @staticmethod\n    @PSDGenerator\n    def brownian_noise(f):\n        return 1/np.where(f == 0, float('inf'), f)\n\n    @staticmethod\n    @PSDGenerator\n    def pink_noise(f):\n        return 1/np.where(f == 0, float('inf'), np.sqrt(f))\n\n    \n    @staticmethod\n    def audio_aug(waveform, color, p, min_snr = 1.0, max_snr = 20.0):\n        if check_prob(p):\n            length = waveform.shape[-1]\n            noise = color(length)\n\n            snr = np.random.uniform(min_snr, max_snr)\n            a_signal = np.sqrt(waveform ** 2).mean()\n            a_noise = a_signal / (10 ** (snr / 20))\n\n            a_color = np.sqrt(noise ** 2).mean()\n\n    #         augmented = (waveform + torch.from_numpy(noise) * 1 / a_color * a_noise).to(waveform.dtype)\n            augmented = (waveform + torch.from_numpy(noise)).to(waveform.dtype)\n\n            return augmented\n        return waveform","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.611086Z","iopub.execute_input":"2023-05-11T15:22:27.611757Z","iopub.status.idle":"2023-05-11T15:22:27.634458Z","shell.execute_reply.started":"2023-05-11T15:22:27.611721Z","shell.execute_reply":"2023-05-11T15:22:27.633544Z"},"id":"xCjYwl8JBS7U","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BirdDS(Dataset):\n#     def __init__(self, df, mixup = False, audio_aug = False, spec_aug = False):\n    def __init__(self, df, mixup = False, audio_aug = False):\n        self.df = df\n#         self.inp = features.mel.MelSpectrogram(n_fft=2048, n_mels = 128, sr=config.sample_rate) # Initializing the model\n#         self.inp = features.cqt.CQT(sr=config.sample_rate, n_bins = 84)\n#         self.db_transform = transforms.AmplitudeToDB(top_db=80)\n    \n        self.mixup = mixup\n        self.audio_aug = audio_aug\n#         self.spec_aug = spec_aug\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def preprocess(self, file):\n        x = AudioUtil.open(file)\n        x = AudioUtil.resample(x, config.sample_rate)\n        x = AudioUtil.rechannel(x, 1)\n        x, sr = AudioUtil.pad_trunc(x, config.duration)\n        \n        return x, sr\n    \n    def __getitem__(self, idx, mixup = False, audio_aug = False):\n        file = f\"{config.base_path}/train_audio/{self.df.filename.iloc[idx]}\"\n        x, _ = self.preprocess(file)\n        y = torch.tensor(self.df.labels.iloc[idx])\n        \n        if self.mixup:\n            if idx % config.freq_mixup == 0:\n                lam = np.random.beta(0.2, 0.2)\n                mixup_idx = np.random.randint(low = 0, high = len(self.df))\n                mixup_file = f\"{config.base_path}/train_audio/{train_df.filename.iloc[mixup_idx]}\"\n                mixup_label = torch.tensor(train_df.labels.iloc[mixup_idx])\n\n                mixup_waveform, _ = self.preprocess(mixup_file)\n\n                x = lam * x + (1 - lam) * mixup_waveform\n                y = lam * y + (1 - lam) * mixup_label\n        \n        if self.audio_aug:\n            x = AudioUtil.audio_aug(x, AudioUtil.brownian_noise, config.p_audio_aug)\n#         x = self.inp(x)\n#         x = self.db_transform(x)\n        \n#         if self.spec_aug:\n#             x = AudioUtil.spec_aug(x, config.p_freq_block, config.p_time_block)\n        \n        return x, y","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.637123Z","iopub.execute_input":"2023-05-11T15:22:27.637476Z","iopub.status.idle":"2023-05-11T15:22:27.652307Z","shell.execute_reply.started":"2023-05-11T15:22:27.63745Z","shell.execute_reply":"2023-05-11T15:22:27.651016Z"},"id":"JH1oZuwmBS7U","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataModule():\n\n    def __init__(self, train, val = None, weights = None):\n        self.train, self.val = train, val\n        \n    def train_dataloader(self):\n        train_loader = DataLoader(self.train, batch_size = config.batch_size, shuffle = True, pin_memory=True, num_workers = 2)\n        return train_loader\n\n    def val_dataloader(self):\n        val_loader = DataLoader(self.val, batch_size = config.batch_size, pin_memory=True, num_workers = 2)\n        return val_loader","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.653949Z","iopub.execute_input":"2023-05-11T15:22:27.654337Z","iopub.status.idle":"2023-05-11T15:22:27.667735Z","shell.execute_reply.started":"2023-05-11T15:22:27.654305Z","shell.execute_reply":"2023-05-11T15:22:27.666661Z"},"id":"EI9oGBZSBS7V","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3) Model Creation!","metadata":{"id":"JbdYALK0BS7V"}},{"cell_type":"code","source":"def init_layer(layer):\n    nn.init.xavier_uniform_(layer.weight)\n\n    if hasattr(layer, \"bias\"):\n        if layer.bias is not None:\n            layer.bias.data.fill_(0.)\n\n\ndef init_bn(bn):\n    bn.bias.data.fill_(0.)\n    bn.weight.data.fill_(1.0)\n\n\ndef interpolate(x: torch.Tensor, ratio: int):\n    \"\"\"Interpolate data in time domain. This is used to compensate the\n    resolution reduction in downsampling of a CNN.\n\n    Args:\n      x: (batch_size, time_steps, classes_num)\n      ratio: int, ratio to interpolate\n    Returns:\n      upsampled: (batch_size, time_steps * ratio, classes_num)\n    \"\"\"\n    (batch_size, time_steps, classes_num) = x.shape\n    upsampled = x[:, :, None, :].repeat(1, 1, ratio, 1)\n    upsampled = upsampled.reshape(batch_size, time_steps * ratio, classes_num)\n    return upsampled\n\n\ndef pad_framewise_output(framewise_output: torch.Tensor, frames_num: int):\n    \"\"\"Pad framewise_output to the same length as input frames. The pad value\n    is the same as the value of the last frame.\n    Args:\n      framewise_output: (batch_size, frames_num, classes_num)\n      frames_num: int, number of frames to pad\n    Outputs:\n      output: (batch_size, frames_num, classes_num)\n    \"\"\"\n    pad = framewise_output[:, -1:, :].repeat(\n        1, frames_num - framewise_output.shape[1], 1)\n    \"\"\"tensor for padding\"\"\"\n\n    output = torch.cat((framewise_output, pad), dim=1)\n    \"\"\"(batch_size, frames_num, classes_num)\"\"\"\n\n    return output\n\nclass AttBlock(nn.Module):\n    def __init__(self,\n                 in_features: int,\n                 out_features: int,\n                 activation=\"linear\",\n                 temperature=1.0):\n        super().__init__()\n\n        self.activation = activation\n        self.temperature = temperature\n        self.att = nn.Conv1d(\n            in_channels=in_features,\n            out_channels=out_features,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True)\n        self.cla = nn.Conv1d(\n            in_channels=in_features,\n            out_channels=out_features,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True)\n\n        self.bn_att = nn.BatchNorm1d(out_features)\n        self.init_weights()\n\n    def init_weights(self):\n        init_layer(self.att)\n        init_layer(self.cla)\n        init_bn(self.bn_att)\n\n    def forward(self, x):\n        # x: (n_samples, n_in, n_time)\n        norm_att = torch.softmax(torch.tanh(self.att(x)), dim=-1)\n        cla = self.nonlinear_transform(self.cla(x))\n        x = torch.sum(norm_att * cla, dim=2)\n        return x, norm_att, cla\n\n    def nonlinear_transform(self, x):\n        if self.activation == 'linear':\n            return x\n        elif self.activation == 'sigmoid':\n            return torch.sigmoid(x)\n        \nclass Model(nn.Module):\n    def __init__(self, spec_aug = True):\n        super().__init__()\n        self.spec = spec_aug\n        \n        self.interpolate_ratio = 32  # Downsampled ratio\n        self.inp = features.cqt.CQT(sr=config.sample_rate, n_bins = 84, verbose = False, trainable=False)\n#         self.inp = features.mel.MelSpectrogram(n_fft=2048, n_mels = 128, sr=config.sample_rate) # Initializing the model\n        self.db = transforms.AmplitudeToDB(top_db=80)\n\n        self.bn0 = nn.BatchNorm2d(config.n_bins)\n\n        self.fc1 = nn.Linear(1024, 1024, bias=True)\n        self.att_block = AttBlock(1024, config.num_classes, activation='sigmoid')\n\n        self.model = timm.create_model('densenet121', pretrained=False).features\n\n        self.init_weight()\n\n    def init_weight(self):\n        init_bn(self.bn0)\n        init_layer(self.fc1)\n        \n    def cnn_feature_extractor(self, x):\n        x = self.model(x)\n        return x\n    \n    def spec_aug(self, spec, p_freq_block, p_time_block):\n        X = spec\n\n        masking = T.TimeMasking(time_mask_param = 80)\n        freq_masking = T.FrequencyMasking(freq_mask_param = 20)\n\n        if check_prob(p_freq_block):\n            X = freq_masking(X)\n\n        if check_prob(p_time_block):\n            X = masking(X)\n\n        return X\n        \n    def preprocess(self, input_x):\n      x = self.inp(input_x)\n      x = self.db(x)\n        \n      if self.spec:\n          x = self.spec_aug(x, config.p_freq_block, config.p_time_block)\n          \n      x = torch.stack([x, x, x]).transpose(0, 1)\n      x = x.transpose(2, 3)  # (batch_size, 1, time_steps, freq_bins)\n\n      frames_num = x.shape[2]\n\n      x = x.transpose(1, 3)\n      x = self.bn0(x)\n      x = x.transpose(1, 3)\n\n      return x, frames_num\n      \n\n    def forward(self, input_data):\n        input_x = input_data\n        \"\"\"\n        Input: (batch_size, data_length)\"\"\"\n        x, frames_num = self.preprocess(input_x)\n        b = x.shape[0]\n        c = 1\n        \n        # Output shape (batch size, channels, time, frequency)\n        x = x.expand(x.shape[0], 3, x.shape[2], x.shape[3])\n        x = self.cnn_feature_extractor(x)\n        \n        # Aggregate in frequency axis\n        x = torch.mean(x, dim=3)\n\n        x1 = F.max_pool1d(x, kernel_size=3, stride=1, padding=1)\n        x2 = F.avg_pool1d(x, kernel_size=3, stride=1, padding=1)\n        x = x1 + x2\n\n        x = F.dropout(x, p=0.5, training=self.training)\n        x = x.transpose(1, 2)\n        x = F.relu_(self.fc1(x))\n        x = x.transpose(1, 2)\n        x = F.dropout(x, p=0.5, training=self.training)\n\n        (clipwise_output, norm_att, segmentwise_output) = self.att_block(x)\n        segmentwise_output = segmentwise_output.transpose(1, 2)\n\n        # Get framewise output\n        framewise_output = interpolate(segmentwise_output, self.interpolate_ratio)\n        framewise_output = pad_framewise_output(framewise_output, frames_num)\n        frame_shape = framewise_output.shape\n        clip_shape = clipwise_output.shape\n        output_dict = {\n            'framewise_output': framewise_output.reshape(b, c, frame_shape[1], frame_shape[2]),\n            'clipwise_output': clipwise_output.reshape(b, c, clip_shape[1]),\n        }\n\n        return output_dict","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.669506Z","iopub.execute_input":"2023-05-11T15:22:27.670305Z","iopub.status.idle":"2023-05-11T15:22:27.70101Z","shell.execute_reply.started":"2023-05-11T15:22:27.67027Z","shell.execute_reply":"2023-05-11T15:22:27.699836Z"},"id":"jncCUGc5BS7V","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PANNsLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.loss = torch.nn.BCELoss()\n\n    def forward(self, inp, target):\n        input_ = torch.where(torch.isnan(inp),\n                             torch.zeros_like(inp),\n                             inp)\n        input_ = torch.where(torch.isinf(input_),\n                             torch.zeros_like(input_),\n                             input_)\n        \n        return self.loss(input_, target)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.702432Z","iopub.execute_input":"2023-05-11T15:22:27.70297Z","iopub.status.idle":"2023-05-11T15:22:27.717042Z","shell.execute_reply.started":"2023-05-11T15:22:27.70293Z","shell.execute_reply":"2023-05-11T15:22:27.715732Z"},"id":"tD6L9bsIBS7W","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4) Training Loop 💀","metadata":{"id":"naKDc-ywBS7W"}},{"cell_type":"markdown","source":"## 4.1) Prepping Data","metadata":{"id":"mP7Tsi4mBS7X"}},{"cell_type":"code","source":"encoder = OneHotEncoder(sparse=False)\n    \ndf = pd.read_csv(f\"{config.base_path}/train_metadata.csv\")\nlabels = encoder.fit_transform(df['primary_label'].to_numpy().reshape(-1,1))\ndf['labels'] = pd.DataFrame(labels).apply(lambda x: list(x), axis = 1)\ntrain_df, val_df = birds_stratified_split(df)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:27.718943Z","iopub.execute_input":"2023-05-11T15:22:27.719326Z","iopub.status.idle":"2023-05-11T15:22:32.37574Z","shell.execute_reply.started":"2023-05-11T15:22:27.719282Z","shell.execute_reply":"2023-05-11T15:22:32.37457Z"},"id":"nST2J8uOBS7b","outputId":"eb704b0f-b141-4af0-cdcf-b611c9519764","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds, val_ds = BirdDS(train_df, mixup = True, audio_aug = True), BirdDS(val_df)\ndm = DataModule(train_ds, val_ds)\ntrain_loader, val_loader = dm.train_dataloader(), dm.val_dataloader()","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:32.377202Z","iopub.execute_input":"2023-05-11T15:22:32.377577Z","iopub.status.idle":"2023-05-11T15:22:32.634846Z","shell.execute_reply.started":"2023-05-11T15:22:32.377542Z","shell.execute_reply":"2023-05-11T15:22:32.633559Z"},"id":"F9jnRrEqBS7c","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.2) Setting Up Criterions + Optimizers + Schedulers","metadata":{"id":"aRy8zbUbBS7c"}},{"cell_type":"code","source":"model = Model(spec_aug = True)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:32.636415Z","iopub.execute_input":"2023-05-11T15:22:32.637146Z","iopub.status.idle":"2023-05-11T15:22:32.867683Z","shell.execute_reply.started":"2023-05-11T15:22:32.637109Z","shell.execute_reply":"2023-05-11T15:22:32.866584Z"},"id":"U4jwuo6xBS7c","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = PANNsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr = config.lr)\nscheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, config.max_lr, epochs = config.epochs, steps_per_epoch = len(train_loader))","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:32.869493Z","iopub.execute_input":"2023-05-11T15:22:32.870006Z","iopub.status.idle":"2023-05-11T15:22:32.880906Z","shell.execute_reply.started":"2023-05-11T15:22:32.869968Z","shell.execute_reply":"2023-05-11T15:22:32.879272Z"},"id":"bBqO5thUBS7c","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.3) Training Utilities","metadata":{"id":"5xse4gDUBS7d"}},{"cell_type":"code","source":"def train_fn(model, epoch, train_loader, optimizer, scheduler, criterion):\n    model.train()\n    \n    losses = AverageMeter()\n#     running_train_metric = MetricTracker(micro_ap_score)\n#     final_train_metric = MetricTracker(padded_cmap)\n    latest_loss = None\n    latest_metric = 0.0\n    \n    pbar = tqdm(train_loader, desc = f\"Training Loop Epoch: {epoch}\", mininterval=0, position=0, leave = True)\n        \n        \n    for batch_idx, (X, Y) in enumerate(pbar):\n        X = X.to(config.device)\n        Y = Y.to(config.device)\n        model = model.to(config.device)\n\n        batch_size = Y.size(0)\n\n        y_hat = model(X)\n        y_hat = torch.squeeze(y_hat[\"clipwise_output\"])\n        train_loss = criterion(y_hat, Y)\n\n        scaled_loss = train_loss / config.grad_accum\n        losses.update(scaled_loss.item(), batch_size)\n\n        y_hat_probs = torch.nn.functional.softmax(y_hat, dim = 1)\n#         running_train_metric.update(Y, y_hat_probs)\n#         final_train_metric.update(Y, y_hat_probs)\n\n        scaled_loss.backward()\n            \n        if (batch_idx + 1) % config.grad_accum == 0:\n            # Training\n            optimizer.step()\n            optimizer.zero_grad()\n            \n            scheduler.step()\n\n            # Logging\n            latest_avg = f\"{losses.avg:.4f}\"\n#             latest_metric = f\"{running_train_metric.score():.4f}\"\n            for i, lr in enumerate(scheduler.get_last_lr()):\n                wandb.log({f\"Layer {i} Learning Rate\": lr})\n\n            wandb.log({f\"Training Loss Step\": losses.val})\n#             wandb.log({f\"Training {config.running_metric_name} Step\": running_train_metric.score()})\n\n#             text = f\"Epoch: {epoch} | Training {config.running_metric_name}: {latest_metric} | Training Loss Averaged: {latest_avg} | Training Loss Step: {losses.val:.4f} | Learning Rate: {scheduler.get_last_lr()[0]:.4f}\"\n            text = f\"Epoch: {epoch} | Training Loss Averaged: {latest_avg} | Training Loss Step: {losses.val:.4f} | Learning Rate: {scheduler.get_last_lr()[0]:.4f}\"\n            pbar.set_postfix_str(text)\n            pbar.refresh()\n\n    average_loss = losses.avg\n#     average_running_metric = running_train_metric.score()\n#     average_final_metric = final_train_metric.score()\n\n    wandb.log({f\"Training Loss Epoch\": average_loss})\n#     wandb.log({f\"Training {config.running_metric_name} Epoch\": average_running_metric})\n#     wandb.log({f\"Training {config.final_metric_name} Epoch\": average_final_metric})\n\n#     return average_loss, average_running_metric, average_final_metric\n#     return average_loss, average_final_metric\n    return average_loss\n    ","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:32.882749Z","iopub.execute_input":"2023-05-11T15:22:32.883239Z","iopub.status.idle":"2023-05-11T15:22:32.897859Z","shell.execute_reply.started":"2023-05-11T15:22:32.8832Z","shell.execute_reply":"2023-05-11T15:22:32.896629Z"},"id":"ECJJg8KSBS7d","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid_fn(model, epoch, val_loader, criterion):\n    model.eval()\n\n    losses = AverageMeter()\n#     running_val_metric = MetricTracker(micro_ap_score)\n    final_val_metric = MetricTracker(padded_cmap)\n    pbar = tqdm(val_loader, desc = f\"Validation Loop Epoch: {epoch}\", mininterval=0, position=0, leave=True)\n    \n    for batch_idx, (X, Y) in enumerate(pbar):\n        X = X.to(config.device)\n        Y = Y.to(config.device)\n        model = model.to(config.device)\n\n        batch_size = Y.size(0)\n\n        with torch.no_grad():\n            y_hat = model(X)\n            y_hat = torch.squeeze(y_hat[\"clipwise_output\"])\n\n        val_loss = criterion(y_hat, Y).item()\n        losses.update(val_loss, batch_size)\n        \n        y_hat_probs = torch.nn.functional.softmax(y_hat, dim = 1)\n#         running_val_metric.update(Y, y_hat_probs)\n        final_val_metric.update(Y, y_hat_probs)\n\n#         pbar.set_postfix_str(f\"Epoch: {epoch} | Validation Loss Average: {losses.avg:.4f} | Validation {config.running_metric_name}: {running_val_metric.score()}\")\n        pbar.set_postfix_str(f\"Epoch: {epoch} | Validation Loss Average: {losses.avg:.4f}\")\n        pbar.refresh()\n\n    average_loss = losses.avg\n#     average_running_metric = running_val_metric.score()\n    average_final_metric = final_val_metric.score()\n\n    wandb.log({f\"Validation Loss Epoch\": average_loss})\n#     wandb.log({f\"Validation {config.running_metric_name} Epoch\": average_running_metric})\n    wandb.log({f\"Validation {config.final_metric_name} Epoch\": average_final_metric})\n\n#     return average_loss, average_running_metric, average_final_metric\n    return average_loss, average_final_metric","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:32.900099Z","iopub.execute_input":"2023-05-11T15:22:32.900663Z","iopub.status.idle":"2023-05-11T15:22:32.91581Z","shell.execute_reply.started":"2023-05-11T15:22:32.900628Z","shell.execute_reply":"2023-05-11T15:22:32.914763Z"},"id":"OTGJXAaDBS7e","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5) Train!","metadata":{"id":"7nl8q9_nBS7e"}},{"cell_type":"code","source":"def run(model, train_loader, val_loader, criterion, optimizer, scheduler):\n    wandb.init(project=\"BirdCLEF 2023\", entity = \"kagglers\", group = \"Leo's Models\", config = dict_from_class(config), job_type = f\"{config.name}\", reinit = True, name = f\"{config.model_name}\")\n    tracker = ModelTracker(model, f\"{config.model_name}.pt\", optimizer, scheduler)\n    \n    wandb.watch(model)\n    \n    for epoch in range(config.epochs):\n#         train_loss, running_train_metric, final_train_metric = train_fn(model, epoch, train_loader, optimizer, scheduler, criterion)\n        train_loss = train_fn(model, epoch, train_loader, optimizer, scheduler, criterion)\n#         val_loss, running_val_metric, final_val_metric = valid_fn(model, epoch, val_loader, criterion)\n        val_loss, final_val_metric = valid_fn(model, epoch, val_loader, criterion)\n        tracker.update(final_val_metric, epoch)\n        \n        if not tracker.check_improvement():\n            print(\"Model not improving. Early stopping now.\")\n            break\n            \n    checkpoint = tracker.get_full_path()\n    saved = torch.load(checkpoint)\n    \n    wandb.save(checkpoint)\n\n    torch.cuda.empty_cache()\n    \n    wandb.finish()\n    \n    return saved","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:32.919443Z","iopub.execute_input":"2023-05-11T15:22:32.919727Z","iopub.status.idle":"2023-05-11T15:22:32.932013Z","shell.execute_reply.started":"2023-05-11T15:22:32.919703Z","shell.execute_reply":"2023-05-11T15:22:32.930984Z"},"id":"r6KHBzOEBS7e","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%pdb","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:32.933863Z","iopub.execute_input":"2023-05-11T15:22:32.935199Z","iopub.status.idle":"2023-05-11T15:22:32.948407Z","shell.execute_reply.started":"2023-05-11T15:22:32.935164Z","shell.execute_reply":"2023-05-11T15:22:32.947436Z"},"id":"ymbEXE0RBS7f","outputId":"ecdf87a3-d1d6-4d15-da1e-f2dd183c576e","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(config.seed)\nsaved_model = run(model, train_loader, val_loader, criterion, optimizer, scheduler)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:22:32.954081Z","iopub.execute_input":"2023-05-11T15:22:32.954909Z","iopub.status.idle":"2023-05-11T15:25:18.77514Z","shell.execute_reply.started":"2023-05-11T15:22:32.954877Z","shell.execute_reply":"2023-05-11T15:25:18.772051Z"},"id":"CTzDdnzFBS7f","outputId":"97f2cf77-740a-4f63-f4dc-fc2b602e0f48","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(saved_model[\"model_state_dict\"])","metadata":{"execution":{"iopub.status.busy":"2023-05-11T15:25:18.776672Z","iopub.status.idle":"2023-05-11T15:25:18.777419Z","shell.execute_reply.started":"2023-05-11T15:25:18.777164Z","shell.execute_reply":"2023-05-11T15:25:18.777187Z"},"id":"04-xfrOMBS7f","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"zkOMUQSuBS7r"},"execution_count":null,"outputs":[]}]}