{"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":8231800,"sourceType":"datasetVersion","datasetId":4881848},{"sourceId":154204277,"sourceType":"kernelVersion"},{"sourceId":167220511,"sourceType":"kernelVersion"}],"dockerImageVersionId":30683,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BirdCLEF'24 | inference with Noise Reduction\n\n* This notebook is a fork of the excellent notebook by Koolo\n     * Train: [BirdCLEF'24 | EfficientNetB0 PyTorch [Train]](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-efficientnetb0-pytorch-train)\n     * Infer: [BirdCLEF'24 | Inference with ONNX](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-inference-with-onnx)\n\n* My notebook (ConvNeXt)\n     * Train: [BirdCLEF'24 | ConvNeXt PyTorch [Train]](https://www.kaggle.com/code/kmatsu01/birdclef-24-convnext-pytorch-train/)\n     * Infer: [BirdCLEF'24 | ConvNeXt [inference with ONNX]](https://www.kaggle.com/code/kmatsu01/birdclef-24-convnext-inference-with-onnx/)\n\n* My notebook (Noise Reduction)\n     * Train: This notebook\n     * Infer: [BirdCLEF'24 | infererence with Noise Reduction](https://www.kaggle.com/code/kmatsu01/birdclef-24-infererence-with-noise-reduction/)\n\n\nAs only CPU notebooks are allowed and there are only two hours for inference, speeding up inference (by quantization, pruning, and distillation) is important. This notebook implements the inference and submission with ONNX. You can find the common inference in this [notebook](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-efficientnetb0-pytorch-inference).\n\n## Features\n- PyTorch's Dataset & Dataloader\n- Use PyTorch-Lightning for building model\n- Data slice is based on @MARK WIJKHUIZEN's [notebook](https://www.kaggle.com/code/markwijkhuizen/birdclef-2024-efficientvit-inference).\n- Inference with ONNX\n- <span style=\"color:red;\">Use of Noise Reduction</span>\n\n## Table of Contents\n\n- [Install ONNX](#Install-ONNX)\n- [Import Packages](#Import-Packages)\n- [Configuration](#Configuration)\n- [Dataset & Dataloader](#Dataset-&-Dataloader)\n- [Model](#Model)\n- [ONNX](#ONNX)\n- [Functions of Inference Loop](#Functions-of-Inference-Loop)\n- [Inference & Submision](#Inference-&-Submision)","metadata":{}},{"cell_type":"markdown","source":"## Install NoiseReduction","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/noisereduce3/noisereduce-3.0.2-py3-none-any.whl","metadata":{},"execution_count":null,"outputs":[]},{"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.59636Z"},"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.6452Z","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.87844Z","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":"import noisereduce as nr","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def SpectralNoiseReduction(audio_data, sr, min_length_sec=5):\n    # Skip processing if the length of the audio data is less than the specified minimum length in seconds\n    if len(audio_data) < sr * min_length_sec:\n        return audio_data\n\n    # Calculate the transition of noise levels across the audio data using a window size of 3 seconds and an overlap of 1.5 seconds\n    hop_length = int(sr * 1.5)  # 1.5 seconds overlap\n    win_length = int(sr * 3)    # 3 seconds window size\n    rms = librosa.feature.rms(y=audio_data, frame_length=win_length, hop_length=hop_length)\n\n    # Identify the time with the smallest noise level\n    noise_sec = 1  # noise reference length in seconds\n    min_rms_idx = np.argmin(rms)  # index of the minimum RMS value\n    start_idx = min_rms_idx * hop_length\n    end_idx = start_idx + sr * noise_sec  # Extract 1 second of data around the time of minimum noise\n\n    # Adjust the indices to make sure they are within the bounds of the audio data\n    start_idx = max(0, start_idx)  # Ensure start index is not negative\n    end_idx = min(len(audio_data), end_idx)  # Ensure end index does not exceed the length of the audio data\n\n    # Use the extracted data as the reference noise data\n    noise_data = audio_data[start_idx:end_idx]\n\n    # Perform noise reduction\n    return nr.reduce_noise(y=audio_data, sr=sr, y_noise=noise_data)","metadata":{},"execution_count":null,"outputs":[]},{"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    # Noise Reduction\n    audio_data = SpectralNoiseReduction(audio_data, config.FS)    \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.96082Z","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.76731Z","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.20868Z","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.45162Z","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.46476Z","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.37842Z","iopub.execute_input":"2024-04-14T05:18:24.3787Z","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.39828Z","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.53214Z","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":[]}]}