{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":8090934,"sourceType":"datasetVersion","datasetId":4776799},{"sourceId":154204277,"sourceType":"kernelVersion"},{"sourceId":167220511,"sourceType":"kernelVersion"}],"dockerImageVersionId":30684,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BirdCLEF 2024 [Train]\n\nThis is the baseline of EfficientNetB0 with PyTorch. I strive to simplify the process to get started faster. Therefore, I only use data from BirdCLEF 2024 (no extended data from BirdCLEF'23, 22, 21, and other sources) and unlabeled soundscapes are also not used. Besides, PyTorch-Lightning is employed to organize the training.\n\nHope this notebook is useful for you!!\n\n* [Pre-Processing](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-speed-up-audio-to-spec-via-cupy)\n* [The inference Notebook](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-efficientnetb0-pytorch-inference)\n\n## Features\n- Implement with PyTorch and PyTorch-Lightning\n- Speed up audio-to-spec. via CuPy\n- Use EfficientNetB0 from torchvision\n\n## Table of Contents\n\n- [Import Packages](#Import-Packages)\n- [Configuration](#Configuration)\n- [Dataset & Dataloader](#Dataset-&-Dataloader)\n- [Model](#Model)\n- [Functions of Training Loop](#Functions-of-Training-Loop)\n- [Training](#Training)\n\n## Update\n\n- V3: fix bug - After validation, self.validation_step_outputs should be cleared.\n- V4: use XYMasking","metadata":{}},{"cell_type":"markdown","source":"# Import packages\n\nImport all required packages.","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport cv2\nimport math\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\nimport librosa\nfrom scipy import signal as sci_signal\n\nimport torch\nfrom torch import nn\nfrom torchvision.models import efficientnet\n\nimport albumentations as albu\n\nimport pytorch_lightning as pl\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\nfrom pytorch_lightning.callbacks import ModelCheckpoint, TQDMProgressBar\n\n# import score function of BirdCLEF\nsys.path.append('/kaggle/input/birdclef-roc-auc')\nsys.path.append('/kaggle/usr/lib/kaggle_metric_utilities')\nfrom metric import score","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:06.955086Z","iopub.execute_input":"2024-04-14T05:17:06.955487Z","iopub.status.idle":"2024-04-14T05:17:19.597205Z","shell.execute_reply.started":"2024-04-14T05:17:06.955456Z","shell.execute_reply":"2024-04-14T05:17:19.596360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration\n\nHyper-paramters","metadata":{}},{"cell_type":"code","source":"class config:\n    \n    # == global config ==\n    SEED = 2024  # random seed\n    DEVICE = 'cuda'  # device to be used\n    MIXED_PRECISION = False  # whether to use mixed-16 precision\n    OUTPUT_DIR = '/kaggle/working/'  # output folder\n    \n    # == data config ==\n    DATA_ROOT = '/kaggle/input/birdclef-2024'  # root folder\n    PREPROCESSED_DATA_ROOT = '/kaggle/input/birdclef24-spectrograms-via-cupy'\n    LOAD_DATA = True  # whether to load data from pre-processed dataset\n    FS = 32000  # sample rate\n    N_FFT = 1095  # n FFT of Spec.\n    WIN_SIZE = 412  # WIN_SIZE of Spec.\n    WIN_LAP = 100  # overlap of Spec.\n    MIN_FREQ = 40  # min frequency\n    MAX_FREQ = 15000  # max frequency\n    \n    # == model config ==\n    MODEL_TYPE = 'efficientnet_b0'  # model type\n    \n    # == dataset config ==\n    BATCH_SIZE = 32  # batch size of each step\n    N_WORKERS = 4  # number of workers\n    \n    # == AUG ==\n    USE_XYMASKING = True  # whether use XYMasking\n    \n    # == training config ==\n    FOLDS = 10  # n fold\n    EPOCHS = 15  # max epochs\n    LR = 1e-3  # learning rate\n    WEIGHT_DECAY = 1e-5  # weight decay of optimizer\n    \n    # == other config ==\n    VISUALIZE = True  # whether to visualize data and batch\n    \nprint('fix seed')\npl.seed_everything(config.SEED, workers=True)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:19.598896Z","iopub.execute_input":"2024-04-14T05:17:19.599343Z","iopub.status.idle":"2024-04-14T05:17:19.614475Z","shell.execute_reply.started":"2024-04-14T05:17:19.599317Z","shell.execute_reply":"2024-04-14T05:17:19.613549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# labels\nlabel_list = sorted(os.listdir(os.path.join(config.DATA_ROOT, 'train_audio')))\nlabel_id_list = list(range(len(label_list)))\nlabel2id = dict(zip(label_list, label_id_list))\nid2label = dict(zip(label_id_list, label_list))","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:19.615619Z","iopub.execute_input":"2024-04-14T05:17:19.615899Z","iopub.status.idle":"2024-04-14T05:17:19.645200Z","shell.execute_reply.started":"2024-04-14T05:17:19.615876Z","shell.execute_reply":"2024-04-14T05:17:19.644476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1\"\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint('Using', torch.cuda.device_count(), 'GPU(s)')","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:19.647091Z","iopub.execute_input":"2024-04-14T05:17:19.647358Z","iopub.status.idle":"2024-04-14T05:17:19.686967Z","shell.execute_reply.started":"2024-04-14T05:17:19.647336Z","shell.execute_reply":"2024-04-14T05:17:19.686165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset & Dataloader\n\n1. [Load Metadata](#Load-Metadata): Load metadata from dataset\n2. [Pre-Processing](#Pre-Processing): The function to convert audio to spectrograms.\n3. [Dataset](#Dataset): Yield samples\n4. [Augmentation](#Augmentation): Data augmentation\n5. [Verify](#Verify): Verify the dataset and dataloader work well","metadata":{}},{"cell_type":"markdown","source":"## Load Metadata","metadata":{}},{"cell_type":"code","source":"metadata_df = pd.read_csv(f'{config.DATA_ROOT}/train_metadata.csv')\nmetadata_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:19.687838Z","iopub.execute_input":"2024-04-14T05:17:19.688151Z","iopub.status.idle":"2024-04-14T05:17:19.876619Z","shell.execute_reply.started":"2024-04-14T05:17:19.688129Z","shell.execute_reply":"2024-04-14T05:17:19.875653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = metadata_df[['primary_label', 'rating', 'filename']].copy()\n\n# create target\ntrain_df['target'] = train_df.primary_label.map(label2id)\n# create filepath\ntrain_df['filepath'] = config.DATA_ROOT + '/train_audio/' + train_df.filename\n# create new sample name\ntrain_df['samplename'] = train_df.filename.map(lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0])\n\nprint(f'find {len(train_df)} samples')\n\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:19.878065Z","iopub.execute_input":"2024-04-14T05:17:19.878440Z","iopub.status.idle":"2024-04-14T05:17:19.941268Z","shell.execute_reply.started":"2024-04-14T05:17:19.878408Z","shell.execute_reply":"2024-04-14T05:17:19.940384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pre-Processing\n\nTo speed up audio-to-spectrogram, we employ CuPy. `CuPy is a NumPy/SciPy-compatible array library for GPU-accelerated computing with Python,` which can significant improve the efficiency of conversion. For more detailed analysis, you can refer to this [notebook](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-speed-up-audio-to-spec-via-cupy).\n\nPlease note, in this notebook, we only use the **center 5 sec** of each audio. By default (`Load_DATA=True`), pre-processed data will be loaded from the [dataset](https://www.kaggle.com/datasets/zijiangyang1116/birdclef24-spectrograms-via-cupy). If `Load_DATA` is set to `False`, spectrograms will be create with `CuPy` (about 30 minites).","metadata":{}},{"cell_type":"code","source":"def oog2spec_via_scipy(audio_data):\n    # handles NaNs\n    mean_signal = np.nanmean(audio_data)\n    audio_data = np.nan_to_num(audio_data, nan=mean_signal) if np.isnan(audio_data).mean() < 1 else np.zeros_like(audio_data)\n    \n    # to spec.\n    frequencies, times, spec_data = sci_signal.spectrogram(\n        audio_data, \n        fs=config.FS, \n        nfft=config.N_FFT, \n        nperseg=config.WIN_SIZE, \n        noverlap=config.WIN_LAP, \n        window='hann'\n    )\n    \n    # Filter frequency range\n    valid_freq = (frequencies >= config.MIN_FREQ) & (frequencies <= config.MAX_FREQ)\n    spec_data = spec_data[valid_freq, :]\n    \n    # Log\n    spec_data = np.log10(spec_data + 1e-20)\n    \n    # min/max normalize\n    spec_data = spec_data - spec_data.min()\n    spec_data = spec_data / spec_data.max()\n    \n    return spec_data","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:19.942441Z","iopub.execute_input":"2024-04-14T05:17:19.942697Z","iopub.status.idle":"2024-04-14T05:17:19.949542Z","shell.execute_reply.started":"2024-04-14T05:17:19.942676Z","shell.execute_reply":"2024-04-14T05:17:19.948572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def oog2spec_via_cupy(audio_data):\n    \n    import cupy as cp\n    from cupyx.scipy import signal as cupy_signal\n    \n    audio_data = cp.array(audio_data)\n    \n    # handles NaNs\n    mean_signal = cp.nanmean(audio_data)\n    audio_data = cp.nan_to_num(audio_data, nan=mean_signal) if cp.isnan(audio_data).mean() < 1 else cp.zeros_like(audio_data)\n    \n    # to spec.\n    frequencies, times, spec_data = cupy_signal.spectrogram(\n        audio_data, \n        fs=config.FS, \n        nfft=config.N_FFT, \n        nperseg=config.WIN_SIZE, \n        noverlap=config.WIN_LAP, \n        window='hann'\n    )\n    \n    # Filter frequency range\n    valid_freq = (frequencies >= config.MIN_FREQ) & (frequencies <= config.MAX_FREQ)\n    spec_data = spec_data[valid_freq, :]\n    \n    # Log\n    spec_data = cp.log10(spec_data + 1e-20)\n    \n    # min/max normalize\n    spec_data = spec_data - spec_data.min()\n    spec_data = spec_data / spec_data.max()\n    \n    return spec_data.get()","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:19.950567Z","iopub.execute_input":"2024-04-14T05:17:19.950817Z","iopub.status.idle":"2024-04-14T05:17:19.960820Z","shell.execute_reply.started":"2024-04-14T05:17:19.950791Z","shell.execute_reply":"2024-04-14T05:17:19.960103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.LOAD_DATA:\n    print('load from file')\n    all_bird_data = np.load(f'{config.PREPROCESSED_DATA_ROOT}/spec_center_5sec_256_256.npy', allow_pickle=True).item()\nelse:\n    all_bird_data = dict()\n    for i, row_metadata in tqdm(train_df.iterrows()):\n\n        # load ogg\n        audio_data, _ = librosa.load(row_metadata.filepath, sr=config.FS)\n\n        # crop\n        n_copy = math.ceil(5 * config.FS / len(audio_data))\n        if n_copy > 1: audio_data = np.concatenate([audio_data]*n_copy)\n\n        start_idx = int(len(audio_data) / 2 - 2.5 * config.FS)\n        end_idx = int(start_idx + 5.0 * config.FS)\n        input_audio = audio_data[start_idx:end_idx]\n\n        # ogg to spec.\n        input_spec = oog2spec_via_cupy(input_audio)\n        \n        input_spec = cv2.resize(input_spec, (256, 256), interpolation=cv2.INTER_AREA)\n\n        all_bird_data[row_metadata.samplename] = input_spec.astype(np.float32)\n\n    # save to file\n    np.save(os.path.join(config.OUTPUT_DIR, f'spec_center_5sec_256_256.npy'), all_bird_data)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:17:19.961894Z","iopub.execute_input":"2024-04-14T05:17:19.962184Z","iopub.status.idle":"2024-04-14T05:18:22.734979Z","shell.execute_reply.started":"2024-04-14T05:17:19.962162Z","shell.execute_reply":"2024-04-14T05:18:22.733902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset\n\nTo yield samples.","metadata":{}},{"cell_type":"code","source":"class BirdDataset(torch.utils.data.Dataset):\n    \n    def __init__(\n        self,\n        metadata,\n        augmentation=None,\n        mode='train'\n    ):\n        super().__init__()\n        self.metadata = metadata\n        self.augmentation = augmentation\n        self.mode = mode\n    \n    def __len__(self):\n        return len(self.metadata)\n    \n    def __getitem__(self, index):\n        \n        row_metadata = self.metadata.iloc[index]\n        \n        # load spec. data\n        input_spec = all_bird_data[row_metadata.samplename]\n        \n        # aug\n        if self.augmentation is not None:\n            input_spec = self.augmentation(image=input_spec)['image']\n        \n        # target\n        target = row_metadata.target\n        \n        return torch.tensor(input_spec, dtype=torch.float32), torch.tensor(target, dtype=torch.long)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:22.738876Z","iopub.execute_input":"2024-04-14T05:18:22.739212Z","iopub.status.idle":"2024-04-14T05:18:22.747208Z","shell.execute_reply.started":"2024-04-14T05:18:22.739188Z","shell.execute_reply":"2024-04-14T05:18:22.746281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Augmentation","metadata":{}},{"cell_type":"code","source":"def get_transforms(_type):\n    \n    if _type == 'train':\n        return albu.Compose([\n            albu.HorizontalFlip(0.5),\n            albu.XYMasking(\n                p=0.3,\n                num_masks_x=(1, 3),\n                num_masks_y=(1, 3),\n                mask_x_length=(1, 10),\n                mask_y_length=(1, 20),\n            ) if config.USE_XYMASKING else albu.NoOp()\n        ])\n    elif _type == 'valid':\n        return albu.Compose([])","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:22.748392Z","iopub.execute_input":"2024-04-14T05:18:22.748721Z","iopub.status.idle":"2024-04-14T05:18:22.767310Z","shell.execute_reply.started":"2024-04-14T05:18:22.748694Z","shell.execute_reply":"2024-04-14T05:18:22.766573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Verify","metadata":{}},{"cell_type":"code","source":"def show_batch(ds, row=3, col=3):\n    fig = plt.figure(figsize=(10, 10))\n    img_index = np.random.randint(0, len(ds)-1, row*col)\n    \n    for i in range(len(img_index)):\n        img, label = dummy_dataset[img_index[i]]\n        \n        if isinstance(img, torch.Tensor):\n            img = img.detach().numpy()\n        \n        ax = fig.add_subplot(row, col, i + 1, xticks=[], yticks=[])\n        ax.imshow(img, cmap='jet')\n        ax.set_title(f'ID: {img_index[i]}; Target: {label}')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:24:18.200951Z","iopub.execute_input":"2024-04-14T05:24:18.201611Z","iopub.status.idle":"2024-04-14T05:24:18.208680Z","shell.execute_reply.started":"2024-04-14T05:24:18.201578Z","shell.execute_reply":"2024-04-14T05:24:18.207695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dummy_dataset = BirdDataset(train_df, get_transforms('train'))\n\ntest_input, test_target = dummy_dataset[0]\nprint(test_input.detach().numpy().shape)\n\nif config.VISUALIZE:\n    show_batch(dummy_dataset)\n\ndel dummy_dataset\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:24:20.089433Z","iopub.execute_input":"2024-04-14T05:24:20.090277Z","iopub.status.idle":"2024-04-14T05:24:21.589273Z","shell.execute_reply.started":"2024-04-14T05:24:20.090238Z","shell.execute_reply":"2024-04-14T05:24:21.588393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"markdown","source":"## Network","metadata":{}},{"cell_type":"code","source":"class EffNet(nn.Module):\n    \n    def __init__(self, model_type, n_classes, pretrained=True):\n        super().__init__()\n        \n        if model_type == 'efficientnet_b0':\n            if pretrained: weights = efficientnet.EfficientNet_B0_Weights.DEFAULT\n            else: weights = None\n            self.base_model = efficientnet.efficientnet_b0(weights=weights)\n        elif model_type == 'efficientnet_b1':\n            if pretrained: weights = efficientnet.EfficientNet_B1_Weights.DEFAULT\n            else: weights = None\n            self.base_model = efficientnet.efficientnet_b1(weights=weights)\n        elif model_type == 'efficientnet_b2':\n            if pretrained: weights = efficientnet.EfficientNet_B2_Weights.DEFAULT\n            else: weights = None\n            self.base_model = efficientnet.efficientnet_b2(weights=weights)\n        elif model_type == 'efficientnet_b3':\n            if pretrained: weights = efficientnet.EfficientNet_B3_Weights.DEFAULT\n            else: weights = None\n            self.base_model = efficientnet.efficientnet_b3(weights=weights)\n        else:\n            raise ValueError('model type not supported')\n        \n        self.base_model.classifier[1] = nn.Linear(self.base_model.classifier[1].in_features, n_classes, dtype=torch.float32)\n    \n    def forward(self, x):\n        x = x.unsqueeze(-1)\n        x = torch.cat([x, x, x], dim=3).permute(0, 3, 1, 2)\n        return self.base_model(x)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:23.451247Z","iopub.execute_input":"2024-04-14T05:18:23.451620Z","iopub.status.idle":"2024-04-14T05:18:23.463499Z","shell.execute_reply.started":"2024-04-14T05:18:23.451589Z","shell.execute_reply":"2024-04-14T05:18:23.462615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dummy_model = EffNet(config.MODEL_TYPE, n_classes=len(label_list))\n\ndummy_input = torch.randn(2, 256, 256)\nprint(dummy_model(dummy_input).shape)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:23.464555Z","iopub.execute_input":"2024-04-14T05:18:23.464781Z","iopub.status.idle":"2024-04-14T05:18:24.376348Z","shell.execute_reply.started":"2024-04-14T05:18:23.464760Z","shell.execute_reply":"2024-04-14T05:18:24.375375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model by PyTorch-Lightning","metadata":{}},{"cell_type":"code","source":"class BirdModel(pl.LightningModule):\n    \n    def __init__(self):\n        super().__init__()\n        \n        # == backbone ==\n        self.backbone = EffNet(config.MODEL_TYPE, n_classes=len(label_list))\n        \n        # == loss function ==\n        self.loss_fn = nn.CrossEntropyLoss()\n        \n        # == record ==\n        self.validation_step_outputs = []\n        \n    def forward(self, images):\n        return self.backbone(images)\n    \n    def configure_optimizers(self):\n        \n        # == define optimizer ==\n        model_optimizer = torch.optim.Adam(\n            filter(lambda p: p.requires_grad, self.parameters()),\n            lr=config.LR,\n            weight_decay=config.WEIGHT_DECAY\n        )\n        \n        # == define learning rate scheduler ==\n        lr_scheduler = CosineAnnealingWarmRestarts(\n            model_optimizer,\n            T_0=config.EPOCHS,\n            T_mult=1,\n            eta_min=1e-6,\n            last_epoch=-1\n        )\n        \n        return {\n            'optimizer': model_optimizer,\n            'lr_scheduler': {\n                'scheduler': lr_scheduler,\n                'interval': 'epoch',\n                'monitor': 'val_loss',\n                'frequency': 1\n            }\n        }\n    \n    def training_step(self, batch, batch_idx):\n        \n        # == obtain input and target ==\n        image, target = batch\n        image = image.to(self.device)\n        target = target.to(self.device)\n        \n        # == pred ==\n        y_pred = self(image)\n        \n        # == compute loss ==\n        train_loss = self.loss_fn(y_pred, target)\n        \n        # == record ==\n        self.log('train_loss', train_loss, True)\n        \n        return train_loss\n    \n    def validation_step(self, batch, batch_idx):\n        \n        # == obtain input and target ==\n        image, target = batch\n        image = image.to(self.device)\n        target = target.to(self.device)\n        \n        # == pred ==\n        with torch.no_grad():\n            y_pred = self(image)\n            \n        self.validation_step_outputs.append({\"logits\": y_pred, \"targets\": target})\n        \n    def train_dataloader(self):\n        return self._train_dataloader\n\n    def validation_dataloader(self):\n        return self._validation_dataloader\n    \n    def on_validation_epoch_end(self):\n        \n        # = merge batch data =\n        outputs = self.validation_step_outputs\n        \n        output_val = nn.Softmax(dim=1)(torch.cat([x['logits'] for x in outputs], dim=0)).cpu().detach()\n        target_val = torch.cat([x['targets'] for x in outputs], dim=0).cpu().detach()\n        \n        # = compute validation loss =\n        val_loss = self.loss_fn(output_val, target_val)\n        \n        # target to one-hot\n        target_val = torch.nn.functional.one_hot(target_val, len(label_list))\n        \n        # = val with ROC AUC =\n        gt_df = pd.DataFrame(target_val.numpy().astype(np.float32), columns=label_list)\n        pred_df = pd.DataFrame(output_val.numpy().astype(np.float32), columns=label_list)\n        \n        gt_df['id'] = [f'id_{i}' for i in range(len(gt_df))]\n        pred_df['id'] = [f'id_{i}' for i in range(len(pred_df))]\n        \n        val_score = score(gt_df, pred_df, row_id_column_name='id')\n        \n        self.log(\"val_score\", val_score, True)\n        \n        # clear validation outputs\n        self.validation_step_outputs = list()\n        \n        return {'val_loss': val_loss, 'val_score': val_score}","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:24.378420Z","iopub.execute_input":"2024-04-14T05:18:24.378700Z","iopub.status.idle":"2024-04-14T05:18:24.396352Z","shell.execute_reply.started":"2024-04-14T05:18:24.378676Z","shell.execute_reply":"2024-04-14T05:18:24.395439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions of Training Loop","metadata":{}},{"cell_type":"code","source":"def predict(data_loader, model):\n    model.to(config.DEVICE)\n    model.eval()\n    predictions = []\n    gts = []\n    for batch in tqdm(data_loader):\n        with torch.no_grad():\n            x, y = batch\n            x = x.cuda()\n            outputs = model(x)\n            outputs = nn.Softmax(dim=1)(outputs)\n        predictions.append(outputs.detach().cpu())\n        gts.append(y.detach().cpu())\n    \n    predictions = torch.cat(predictions, dim=0).cpu().detach()\n    gts = torch.cat(gts, dim=0).cpu().detach()\n    gts = torch.nn.functional.one_hot(gts, len(label_list))\n    \n    return predictions.numpy().astype(np.float32), gts.numpy().astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:24.397572Z","iopub.execute_input":"2024-04-14T05:18:24.398311Z","iopub.status.idle":"2024-04-14T05:18:24.411364Z","shell.execute_reply.started":"2024-04-14T05:18:24.398280Z","shell.execute_reply":"2024-04-14T05:18:24.410512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(fold_id, total_df):\n    print('================================================================')\n    print(f\"==== Running training for fold {fold_id} ====\")\n    \n    # == create dataset and dataloader ==\n    train_df = total_df[total_df['fold'] != fold_id].copy()\n    valid_df = total_df[total_df['fold'] == fold_id].copy()\n    \n    print(f'Train Samples: {len(train_df)}')\n    print(f'Valid Samples: {len(valid_df)}')\n    \n    train_ds = BirdDataset(train_df, get_transforms('train'), 'train')\n    val_ds = BirdDataset(valid_df, get_transforms('valid'), 'valid')\n    \n    train_dl = torch.utils.data.DataLoader(\n        train_ds,\n        batch_size=config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=config.N_WORKERS,\n        pin_memory=True,\n        persistent_workers=True\n    )\n    \n    val_dl = torch.utils.data.DataLoader(\n        val_ds,\n        batch_size=config.BATCH_SIZE * 2,\n        shuffle=False,\n        num_workers=config.N_WORKERS,\n        pin_memory=True,\n        persistent_workers=True\n    )\n    \n    # == init model ==\n    bird_model = BirdModel()\n    \n    # == init callback ==\n    checkpoint_callback = ModelCheckpoint(monitor='val_score',\n                                          dirpath=config.OUTPUT_DIR,\n                                          save_top_k=1,\n                                          save_last=False,\n                                          save_weights_only=True,\n                                          filename=f\"fold_{fold_id}\",\n                                          mode='max')\n    callbacks_to_use = [checkpoint_callback, TQDMProgressBar(refresh_rate=1)]\n    \n    # == init trainer ==\n    trainer = pl.Trainer(\n        max_epochs=config.EPOCHS,\n        val_check_interval=0.5,\n        callbacks=callbacks_to_use,\n        enable_model_summary=False,\n        accelerator=\"gpu\",\n        deterministic=True,\n        precision='16-mixed' if config.MIXED_PRECISION else 32,\n    )\n    \n    # == Training ==\n    trainer.fit(bird_model, train_dataloaders=train_dl, val_dataloaders=val_dl)\n    \n    # == Prediction ==\n    best_model_path = checkpoint_callback.best_model_path\n    weights = torch.load(best_model_path)['state_dict']\n    bird_model.load_state_dict(weights)\n    \n    preds, gts = predict(val_dl, bird_model)\n    \n    # = create dataframe =\n    pred_df = pd.DataFrame(preds, columns=label_list)\n    pred_df['id'] = np.arange(len(pred_df))\n    gt_df = pd.DataFrame(gts, columns=label_list)\n    gt_df['id'] = np.arange(len(gt_df))\n    \n    # = compute score =\n    val_score = score(gt_df, pred_df, row_id_column_name='id')\n    \n    # == save to file ==\n    pred_cols = [f'pred_{t}' for t in label_list]\n    valid_df = pd.concat([valid_df.reset_index(), pd.DataFrame(np.zeros((len(valid_df), len(label_list)*2)).astype(np.float32), columns=label_list+pred_cols)], axis=1)\n    valid_df[label_list] = gts\n    valid_df[pred_cols] = preds\n    valid_df.to_csv(f\"{config.OUTPUT_DIR}/pred_df_f{fold_id}.csv\", index=False)\n    \n    return preds, gts, val_score","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:24.412699Z","iopub.execute_input":"2024-04-14T05:18:24.412983Z","iopub.status.idle":"2024-04-14T05:18:24.428194Z","shell.execute_reply.started":"2024-04-14T05:18:24.412961Z","shell.execute_reply":"2024-04-14T05:18:24.427415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"markdown","source":"## KFold","metadata":{}},{"cell_type":"code","source":"kf = KFold(n_splits=config.FOLDS, shuffle=True, random_state=config.SEED)\ntrain_df['fold'] = 0\nfor fold, (train_idx, val_idx) in enumerate(kf.split(train_df)):\n    train_df.loc[val_idx, 'fold'] = fold","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:24.429191Z","iopub.execute_input":"2024-04-14T05:18:24.429456Z","iopub.status.idle":"2024-04-14T05:18:24.457363Z","shell.execute_reply.started":"2024-04-14T05:18:24.429434Z","shell.execute_reply":"2024-04-14T05:18:24.456439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training Loop","metadata":{}},{"cell_type":"code","source":"# training\ntorch.set_float32_matmul_precision('high')\n\n# record\nfold_val_score_list = list()\noof_df = train_df.copy()\npred_cols = [f'pred_{t}' for t in label_list]\noof_df = pd.concat([oof_df, pd.DataFrame(np.zeros((len(oof_df), len(pred_cols)*2)).astype(np.float32), columns=label_list+pred_cols)], axis=1)\n\nfor f in range(config.FOLDS):\n    \n    # get validation index\n    val_idx = list(train_df[train_df['fold'] == f].index)\n    \n    # main loop of f-fold\n    val_preds, val_gts, val_score = run_training(f, train_df)\n    \n    # record\n    oof_df.loc[val_idx, label_list] = val_gts\n    oof_df.loc[val_idx, pred_cols] = val_preds\n    fold_val_score_list.append(val_score)\n    \n    # only training one fold\n    break\n\n\nfor idx, val_score in enumerate(fold_val_score_list):\n    print(f'Fold {idx} Val Score: {val_score:.5f}')\n\n# oof_gt_df = oof_df[['samplename'] + label_list].copy()\n# oof_pred_df = oof_df[['samplename'] + pred_cols].copy()\n# oof_pred_df.columns = ['samplename'] + label_list\n# oof_score = score(oof_gt_df, oof_pred_df, 'samplename')\n# print(f'OOF Score: {oof_score:.5f}')\n\noof_df.to_csv(f\"{config.OUTPUT_DIR}/oof_pred.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:18:24.458534Z","iopub.execute_input":"2024-04-14T05:18:24.458785Z","iopub.status.idle":"2024-04-14T05:23:17.532140Z","shell.execute_reply.started":"2024-04-14T05:18:24.458764Z","shell.execute_reply":"2024-04-14T05:23:17.530819Z"},"trusted":true},"execution_count":null,"outputs":[]}]}