{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"none","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\n# Set the path to the Kaggle input directory\nkaggle_input_dir = \"/kaggle/input\"\n\n# Function to list all files in the directory\ndef list_files(directory):\n    for root, dirs, files in os.walk(directory):\n        for file in files:\n            if not file.endswith(\"ogg\"):\n                print(os.path.join(root, file))\n\n# List all files in the Kaggle input directory\nlist_files(kaggle_input_dir)","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:07:49.767230Z","iopub.execute_input":"2024-06-13T13:07:49.767892Z","iopub.status.idle":"2024-06-13T13:07:51.813853Z","shell.execute_reply.started":"2024-06-13T13:07:49.767855Z","shell.execute_reply":"2024-06-13T13:07:51.812688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"[Mostly this notebook based on Efficentnet BirdCLEF starter](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-efficientnetb0-pytorch-train)\n\nThank to the authors","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\n\nimport math\nimport random\n\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import KFold\n\nimport librosa\nimport torch\nfrom torchvision.models import efficientnet\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint, TQDMProgressBar\n\nimport cv2\nimport matplotlib.pyplot as plt\nfrom scipy import signal as sci_signal\n\nimport albumentations as albu\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:07:51.815658Z","iopub.execute_input":"2024-06-13T13:07:51.816014Z","iopub.status.idle":"2024-06-13T13:07:51.823238Z","shell.execute_reply.started":"2024-06-13T13:07:51.815986Z","shell.execute_reply":"2024-06-13T13:07:51.821659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class cfg:\n       \n       KAGGLE = True\n       LOAD_DATA = False\n       USE_CUPY = True\n\n       SEED = 42\n       FOLDS = 5\n\n       VISUALIZE = True\n\n       # mel specs constants\n       MIN_FREQ = 20\n       MAX_FREQ = 20000\n\n       FS = 32000\n       N_FFT = 412\n       WIN_SIZE = 300\n       WIN_LAP = 0\n\n       MODEL_TYPE = \"efficientnet_b0\" \n       LR = 0.001\n       WEIGHT_DECAY = 0\n       BATCH_SIZE = 128\n       EPOCHS = 5\n       N_WORKERS = 4\n       MIXED_PRECISION = False\n\n       USE_XYMASKING = True\n\n       if KAGGLE:\n              DATA_DIR = \"/kaggle/input/birdclef-2024\"\n              PREPROCESSED_DATA_ROOT = \"/kaggle/input/birdclef24-spectrograms-via-cupy\"\n              OUTPUT_DIR = \"/kaggle/working/\"\n       else:\n              DATA_DIR = \"datka/main\"\n              PREPROCESSED_DATA_ROOT = \"datka/spectrograms_via_cupy\"\n              OUTPUT_DIR = \"datka/spectrograms_via_cupy\"\n\nprint('fix seed')\npl.seed_everything(cfg.SEED, workers=True)   \n\nmeta_df = pd.read_csv(os.path.join(cfg.DATA_DIR, \"train_metadata.csv\"))\n\nmeta_df[\"rating\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:07:51.824639Z","iopub.execute_input":"2024-06-13T13:07:51.825030Z","iopub.status.idle":"2024-06-13T13:07:52.057199Z","shell.execute_reply.started":"2024-06-13T13:07:51.825002Z","shell.execute_reply":"2024-06-13T13:07:52.055781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The part of 5 rating data of all training data","metadata":{}},{"cell_type":"code","source":"meta_df[meta_df.rating >= 5].__len__() / meta_df.__len__()","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:07:52.059593Z","iopub.execute_input":"2024-06-13T13:07:52.059975Z","iopub.status.idle":"2024-06-13T13:07:52.076058Z","shell.execute_reply.started":"2024-06-13T13:07:52.059945Z","shell.execute_reply":"2024-06-13T13:07:52.074630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Distribution of classes in our data","metadata":{}},{"cell_type":"code","source":"all_data_counts = meta_df.primary_label.value_counts()\n\nhigh_quality_counts = meta_df[meta_df.rating == 5].primary_label.value_counts()\n\npd.DataFrame([all_data_counts, high_quality_counts],\n             index=[\"all_data\", \"high_quality\"]).T","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:07:52.077648Z","iopub.execute_input":"2024-06-13T13:07:52.078691Z","iopub.status.idle":"2024-06-13T13:07:52.111055Z","shell.execute_reply.started":"2024-06-13T13:07:52.078649Z","shell.execute_reply":"2024-06-13T13:07:52.109380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Ok. We don't have much labels of rare classes in high quality data, but let's try our magis","metadata":{}},{"cell_type":"code","source":"labels_list = sorted(os.listdir(os.path.join(cfg.DATA_DIR, \"train_audio\")))\n\nlabels_id_list = list(range(len(labels_list)))\nlabel2id = dict(zip(labels_list, labels_id_list))\nid2label = dict(zip(labels_id_list, labels_list))","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:07:52.113360Z","iopub.execute_input":"2024-06-13T13:07:52.113848Z","iopub.status.idle":"2024-06-13T13:07:52.122818Z","shell.execute_reply.started":"2024-06-13T13:07:52.113808Z","shell.execute_reply":"2024-06-13T13:07:52.120855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = meta_df[[\"primary_label\", \"rating\", \"filename\"]].copy()\n\ntrain_df[\"target\"] = train_df[\"primary_label\"].map(label2id)\ntrain_df[\"filepath\"] = cfg.DATA_DIR + \"/train_audio/\" + train_df[\"filename\"]\ntrain_df[\"samplename\"] = train_df.filename.map(lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0])\n\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:07:52.125312Z","iopub.execute_input":"2024-06-13T13:07:52.125827Z","iopub.status.idle":"2024-06-13T13:07:52.197619Z","shell.execute_reply.started":"2024-06-13T13:07:52.125788Z","shell.execute_reply":"2024-06-13T13:07:52.196077Z"},"trusted":true},"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    # to spec.\n    frequencies, times, spec_data = sci_signal.spectrogram(\n        audio_data, \n        fs=cfg.FS, \n        nfft=cfg.N_FFT, \n        nperseg=cfg.WIN_SIZE, \n        noverlap=cfg.WIN_LAP, \n        window='hann'\n    )\n    \n    # Filter frequency range\n    valid_freq = (frequencies >= cfg.MIN_FREQ) & (frequencies <= cfg.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\n\ndef 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=cfg.FS, \n        nfft=cfg.N_FFT, \n        nperseg=cfg.WIN_SIZE, \n        noverlap=cfg.WIN_LAP, \n        window='hann'\n    )\n    \n    # Filter frequency range\n    valid_freq = (frequencies >= cfg.MIN_FREQ) & (frequencies <= cfg.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-06-13T13:07:52.199351Z","iopub.execute_input":"2024-06-13T13:07:52.200513Z","iopub.status.idle":"2024-06-13T13:07:52.215626Z","shell.execute_reply.started":"2024-06-13T13:07:52.200464Z","shell.execute_reply":"2024-06-13T13:07:52.213977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if cfg.LOAD_DATA:\n    print('load from file')\n    all_bird_data = np.load(f'{cfg.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=cfg.FS)\n\n        # crop\n        n_copy = math.ceil(5 * cfg.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 * cfg.FS)\n        end_idx = int(start_idx + 5.0 * cfg.FS)\n        input_audio = audio_data[start_idx:end_idx]\n\n        # ogg to spec.\n        if cfg.USE_CUPY:\n            input_spec = oog2spec_via_cupy(input_audio)\n        else:\n            input_spec = oog2spec_via_scipy(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\nif not os.path.exists(cfg.OUTPUT_DIR):\n    os.makedirs(cfg.OUTPUT_DIR)\nnp.save(os.path.join(cfg.OUTPUT_DIR, f'spec_center_5sec_256_256.npy'), all_bird_data)","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:07:52.217635Z","iopub.execute_input":"2024-06-13T13:07:52.218161Z","iopub.status.idle":"2024-06-13T13:08:10.003840Z","shell.execute_reply.started":"2024-06-13T13:07:52.218082Z","shell.execute_reply":"2024-06-13T13:08:09.999779Z"},"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-06-13T13:08:10.005556Z","iopub.status.idle":"2024-06-13T13:08:10.006537Z","shell.execute_reply.started":"2024-06-13T13:08:10.006256Z","shell.execute_reply":"2024-06-13T13:08:10.006281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Dataset","metadata":{}},{"cell_type":"code","source":"def get_transforms(_type):\n    \n    if _type == 'train':\n        return albu.Compose([\n            albu.HorizontalFlip(p=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 cfg.USE_XYMASKING else albu.NoOp()\n        ])\n    elif _type == 'valid':\n        return albu.Compose([])\n    \ndef 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 = ds[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()\n\nclass 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)\n\ndummy_dataset = BirdDataset(train_df, get_transforms('train'))\n\ntest_input, test_target = dummy_dataset[0]\nprint(test_input.detach().numpy().shape)\n\nif cfg.VISUALIZE:\n    show_batch(dummy_dataset)\n\ndel dummy_dataset\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:08:10.008248Z","iopub.status.idle":"2024-06-13T13:08:10.009302Z","shell.execute_reply.started":"2024-06-13T13:08:10.009008Z","shell.execute_reply":"2024-06-13T13:08:10.009032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Score","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\nfrom typing import Union\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\nclass HostVisibleError(Exception):\n    pass\n\n\ndef treat_as_participant_error(error_message: str, solution: Union[pd.DataFrame, np.ndarray]) -> bool:\n    ''' Many metrics can raise more errors than can be handled manually. This function attempts\n    to identify errors that can be treated as ParticipantVisibleError without leaking any competition data.\n\n    If the solution is purely numeric, and there are no numbers in the error message,\n    then the error message is sufficiently unlikely to leak usable data and can be shown to participants.\n\n    We expect this filter to reject many safe messages. It's intended only to reduce the number of errors we need to manage manually.\n    '''\n    # This check treats bools as numeric\n    if isinstance(solution, pd.DataFrame):\n        solution_is_all_numeric = all([pd.api.types.is_numeric_dtype(x) for x in solution.dtypes.values])\n        solution_has_bools = any([pd.api.types.is_bool_dtype(x) for x in solution.dtypes.values])\n    elif isinstance(solution, np.ndarray):\n        solution_is_all_numeric = pd.api.types.is_numeric_dtype(solution)\n        solution_has_bools = pd.api.types.is_bool_dtype(solution)\n\n    if not solution_is_all_numeric:\n        return False\n\n    for char in error_message:\n        if char.isnumeric():\n            return False\n    if solution_has_bools:\n        if 'true' in error_message.lower() or 'false' in error_message.lower():\n            return False\n    return True\n\n\ndef safe_call_score(metric_function, solution, submission, **metric_func_kwargs):\n    '''\n    Call score. If that raises an error and that already been specifically handled, just raise it.\n    Otherwise make a conservative attempt to identify potential participant visible errors.\n    '''\n    try:\n        score_result = metric_function(solution, submission, **metric_func_kwargs)\n    except Exception as err:\n        error_message = str(err)\n        if err.__class__.__name__ == 'ParticipantVisibleError':\n            raise ParticipantVisibleError(error_message)\n        elif err.__class__.__name__ == 'HostVisibleError':\n            raise HostVisibleError(error_message)\n        else:\n            if treat_as_participant_error(error_message, solution):\n                raise ParticipantVisibleError(error_message)\n            else:\n                raise err\n    return score_result\n\n\ndef verify_valid_probabilities(df: pd.DataFrame, df_name: str):\n    \"\"\" Verify that the dataframe contains valid probabilities.\n\n    The dataframe must be limited to the target columns; do not pass in any ID columns.\n    \"\"\"\n    if not pd.api.types.is_numeric_dtype(df.values):\n        raise ParticipantVisibleError(f'All target values in {df_name} must be numeric')\n\n    if df.min().min() < 0:\n        raise ParticipantVisibleError(f'All target values in {df_name} must be at least zero')\n\n    if df.max().max() > 1:\n        raise ParticipantVisibleError(f'All target values in {df_name} must be no greater than one')\n\n    if not np.allclose(df.sum(axis=1), 1):\n        raise ParticipantVisibleError(f'Target values in {df_name} do not add to one within all rows')\n\n\ndef score(solution: pd.DataFrame, submission: pd.DataFrame): #, row_id_column_name: str) -> float:\n    '''\n    Version of macro-averaged ROC-AUC score that ignores all classes that have no true positive labels.\n    '''\n    #del solution[row_id_column_name]\n    #del submission[row_id_column_name]\n\n    if not pd.api.types.is_numeric_dtype(submission.values):\n        bad_dtypes = {x: submission[x].dtype  for x in submission.columns if not pd.api.types.is_numeric_dtype(submission[x])}\n        raise ParticipantVisibleError(f'Invalid submission data types found: {bad_dtypes}')\n\n    solution_sums = solution.sum(axis=0)\n    scored_columns = list(solution_sums[solution_sums > 0].index.values)\n    assert len(scored_columns) > 0\n\n    return safe_call_score(roc_auc_score, solution[scored_columns].values, submission[scored_columns].values, average='macro')","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:08:10.010766Z","iopub.status.idle":"2024-06-13T13:08:10.011303Z","shell.execute_reply.started":"2024-06-13T13:08:10.011033Z","shell.execute_reply":"2024-06-13T13:08:10.011054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Model","metadata":{}},{"cell_type":"code","source":"class EffNet(torch.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] = torch.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)\n    \ndummy_model = EffNet(cfg.MODEL_TYPE, n_classes=len(labels_list))\n\ndummy_input = torch.randn(2, 256, 256)\nprint(dummy_model(dummy_input).shape)","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:08:10.012696Z","iopub.status.idle":"2024-06-13T13:08:10.013262Z","shell.execute_reply.started":"2024-06-13T13:08:10.012980Z","shell.execute_reply":"2024-06-13T13:08:10.013003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BirdModel(pl.LightningModule):\n    \n    def __init__(self):\n        super().__init__()\n        \n        # == backbone ==\n        self.backbone = EffNet(cfg.MODEL_TYPE, n_classes=len(labels_list))\n        \n        # == loss function ==\n        self.loss_fn = torch.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=cfg.LR,\n            weight_decay=cfg.WEIGHT_DECAY\n        )\n        \n        # == define learning rate scheduler ==\n        lr_scheduler = CosineAnnealingWarmRestarts(\n            model_optimizer,\n            T_0=cfg.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 = torch.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(labels_list))\n        \n        # = val with ROC AUC =\n        gt_df = pd.DataFrame(target_val.numpy().astype(np.float32), columns=labels_list)\n        pred_df = pd.DataFrame(output_val.numpy().astype(np.float32), columns=labels_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-06-13T13:08:10.015162Z","iopub.status.idle":"2024-06-13T13:08:10.015750Z","shell.execute_reply.started":"2024-06-13T13:08:10.015444Z","shell.execute_reply":"2024-06-13T13:08:10.015466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(data_loader, model):\n    model.to(cfg.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 = torch.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(labels_list))\n    \n    return predictions.numpy().astype(np.float32), gts.numpy().astype(np.float32)\n\ndef 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=cfg.BATCH_SIZE,\n        shuffle=True,\n        num_workers=cfg.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=cfg.BATCH_SIZE * 2,\n        shuffle=False,\n        num_workers=cfg.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=cfg.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=cfg.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 cfg.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=labels_list)\n    pred_df['id'] = np.arange(len(pred_df))\n    gt_df = pd.DataFrame(gts, columns=labels_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 labels_list]\n    valid_df = pd.concat([valid_df.reset_index(), pd.DataFrame(np.zeros((len(valid_df), len(labels_list)*2)).astype(np.float32), columns=label_list+pred_cols)], axis=1)\n    valid_df[labels_list] = gts\n    valid_df[pred_cols] = preds\n    valid_df.to_csv(f\"{cfg.OUTPUT_DIR}/pred_df_f{fold_id}.csv\", index=False)\n    \n    return preds, gts, val_score","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:08:10.017273Z","iopub.status.idle":"2024-06-13T13:08:10.017817Z","shell.execute_reply.started":"2024-06-13T13:08:10.017536Z","shell.execute_reply":"2024-06-13T13:08:10.017557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Training","metadata":{}},{"cell_type":"code","source":"kf = KFold(n_splits=cfg.FOLDS, shuffle=True, random_state=cfg.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-06-13T13:08:10.019105Z","iopub.status.idle":"2024-06-13T13:08:10.019625Z","shell.execute_reply.started":"2024-06-13T13:08:10.019369Z","shell.execute_reply":"2024-06-13T13:08:10.019391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 labels_list]\noof_df = pd.concat([oof_df, pd.DataFrame(np.zeros((len(oof_df), len(pred_cols)*2)).astype(np.float32), columns=labels_list+pred_cols)], axis=1)\n\nfor f in range(cfg.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, labels_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\"{cfg.OUTPUT_DIR}/oof_pred.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-06-13T13:08:10.021758Z","iopub.status.idle":"2024-06-13T13:08:10.022321Z","shell.execute_reply.started":"2024-06-13T13:08:10.022032Z","shell.execute_reply":"2024-06-13T13:08:10.022053Z"},"trusted":true},"execution_count":null,"outputs":[]}]}