{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":70203,"databundleVersionId":8068726,"isSourceIdPinned":false},{"sourceType":"competition","sourceId":129329,"databundleVersionId":15996945,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":8090934,"datasetId":4776799,"databundleVersionId":8207598},{"sourceType":"datasetVersion","sourceId":16014749,"datasetId":10271576,"databundleVersionId":16979086,"isSourceIdPinned":true},{"sourceType":"kernelVersion","sourceId":167220511,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":189366517,"isSourceIdPinned":false}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\nimport librosa\nfrom scipy import signal as sci_signal\nfrom pathlib import Path\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\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:49:52.236772Z","iopub.execute_input":"2026-05-02T15:49:52.237081Z","iopub.status.idle":"2026-05-02T15:50:00.524481Z","shell.execute_reply.started":"2026-05-02T15:49:52.237050Z","shell.execute_reply":"2026-05-02T15:50:00.523741Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class config:\n    \n    # == global config ==\n    SEED = 2026  # 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/competitions/birdclef-2026' # root folder\n    PREPROCESSED_DATA_ROOT ='/kaggle/input/birdclef26-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 = 1  # 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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:00.525601Z","iopub.execute_input":"2026-05-02T15:50:00.526387Z","iopub.status.idle":"2026-05-02T15:50:00.539442Z","shell.execute_reply.started":"2026-05-02T15:50:00.526354Z","shell.execute_reply":"2026-05-02T15:50:00.538767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"test\")\n# 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))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:00.541986Z","iopub.execute_input":"2026-05-02T15:50:00.542529Z","iopub.status.idle":"2026-05-02T15:50:00.555708Z","shell.execute_reply.started":"2026-05-02T15:50:00.542493Z","shell.execute_reply":"2026-05-02T15:50:00.554786Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:00.556830Z","iopub.execute_input":"2026-05-02T15:50:00.557154Z","iopub.status.idle":"2026-05-02T15:50:00.609606Z","shell.execute_reply.started":"2026-05-02T15:50:00.557125Z","shell.execute_reply":"2026-05-02T15:50:00.608794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"metadata_df = pd.read_csv(f'{config.DATA_ROOT}/train.csv')\nmetadata_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:00.610617Z","iopub.execute_input":"2026-05-02T15:50:00.610955Z","iopub.status.idle":"2026-05-02T15:50:00.756105Z","shell.execute_reply.started":"2026-05-02T15:50:00.610915Z","shell.execute_reply":"2026-05-02T15:50:00.755419Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:00.757000Z","iopub.execute_input":"2026-05-02T15:50:00.757351Z","iopub.status.idle":"2026-05-02T15:50:00.805196Z","shell.execute_reply.started":"2026-05-02T15:50:00.757283Z","shell.execute_reply":"2026-05-02T15:50:00.804345Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:00.806262Z","iopub.execute_input":"2026-05-02T15:50:00.806612Z","iopub.status.idle":"2026-05-02T15:50:00.813520Z","shell.execute_reply.started":"2026-05-02T15:50:00.806572Z","shell.execute_reply":"2026-05-02T15:50:00.812576Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:00.814550Z","iopub.execute_input":"2026-05-02T15:50:00.814819Z","iopub.status.idle":"2026-05-02T15:50:00.824892Z","shell.execute_reply.started":"2026-05-02T15:50:00.814795Z","shell.execute_reply":"2026-05-02T15:50:00.824085Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:00.825793Z","iopub.execute_input":"2026-05-02T15:50:00.826132Z","iopub.status.idle":"2026-05-02T15:50:54.296854Z","shell.execute_reply.started":"2026-05-02T15:50:00.826095Z","shell.execute_reply":"2026-05-02T15:50:54.296108Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:54.297898Z","iopub.execute_input":"2026-05-02T15:50:54.298608Z","iopub.status.idle":"2026-05-02T15:50:54.304728Z","shell.execute_reply.started":"2026-05-02T15:50:54.298563Z","shell.execute_reply":"2026-05-02T15:50:54.303926Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:54.305927Z","iopub.execute_input":"2026-05-02T15:50:54.306382Z","iopub.status.idle":"2026-05-02T15:50:54.320865Z","shell.execute_reply.started":"2026-05-02T15:50:54.306326Z","shell.execute_reply":"2026-05-02T15:50:54.320086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_batch(dummy_dataset, row=3, col=3):\n    fig = plt.figure(figsize=(10, 10))\n    img_index = np.random.randint(0, len(dummy_dataset)-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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:54.321827Z","iopub.execute_input":"2026-05-02T15:50:54.322120Z","iopub.status.idle":"2026-05-02T15:50:54.331791Z","shell.execute_reply.started":"2026-05-02T15:50:54.322094Z","shell.execute_reply":"2026-05-02T15:50:54.330940Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:54.334741Z","iopub.execute_input":"2026-05-02T15:50:54.335097Z","iopub.status.idle":"2026-05-02T15:50:55.498185Z","shell.execute_reply.started":"2026-05-02T15:50:54.335051Z","shell.execute_reply":"2026-05-02T15:50:55.497219Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:55.499308Z","iopub.execute_input":"2026-05-02T15:50:55.499608Z","iopub.status.idle":"2026-05-02T15:50:55.507340Z","shell.execute_reply.started":"2026-05-02T15:50:55.499581Z","shell.execute_reply":"2026-05-02T15:50:55.506398Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:55.508445Z","iopub.execute_input":"2026-05-02T15:50:55.508855Z","iopub.status.idle":"2026-05-02T15:50:56.165959Z","shell.execute_reply.started":"2026-05-02T15:50:55.508825Z","shell.execute_reply":"2026-05-02T15:50:56.165199Z"}},"outputs":[],"execution_count":null},{"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        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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:56.166997Z","iopub.execute_input":"2026-05-02T15:50:56.167486Z","iopub.status.idle":"2026-05-02T15:50:56.179603Z","shell.execute_reply.started":"2026-05-02T15:50:56.167454Z","shell.execute_reply":"2026-05-02T15:50:56.178766Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:56.180634Z","iopub.execute_input":"2026-05-02T15:50:56.180964Z","iopub.status.idle":"2026-05-02T15:50:56.197485Z","shell.execute_reply.started":"2026-05-02T15:50:56.180924Z","shell.execute_reply":"2026-05-02T15:50:56.196708Z"}},"outputs":[],"execution_count":null},{"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=(config.N_WORKERS > 0)\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=(config.N_WORKERS > 0)\n    )\n    print(\"finish loading data\")\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]\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        enable_progress_bar=True,          \n       \n    )\n    # == Training ==\n    trainer.fit(bird_model, train_dataloaders=train_dl, val_dataloaders=val_dl)\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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:56.198308Z","iopub.execute_input":"2026-05-02T15:50:56.198633Z","iopub.status.idle":"2026-05-02T15:50:56.213878Z","shell.execute_reply.started":"2026-05-02T15:50:56.198590Z","shell.execute_reply":"2026-05-02T15:50:56.213225Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:56.214693Z","iopub.execute_input":"2026-05-02T15:50:56.214987Z","iopub.status.idle":"2026-05-02T15:50:56.244138Z","shell.execute_reply.started":"2026-05-02T15:50:56.214951Z","shell.execute_reply":"2026-05-02T15:50:56.243573Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T15:50:56.244825Z","iopub.execute_input":"2026-05-02T15:50:56.245012Z","iopub.status.idle":"2026-05-02T15:52:04.741144Z","shell.execute_reply.started":"2026-05-02T15:50:56.244991Z","shell.execute_reply":"2026-05-02T15:52:04.739899Z"}},"outputs":[],"execution_count":null}]}