{"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":"none","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":8120971,"sourceType":"datasetVersion","datasetId":4779991}],"dockerImageVersionId":30684,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BirdCLEF 2024 [Inference]\n\nThis notebook implements the inference and submission. You can find the pre-processing and training in the following notebooks:\n\n* [Pre-Processing](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-speed-up-audio-to-spec-via-cupy)\n* [The Training Notebook](https://www.kaggle.com/code/zijiangyang1116/birdclef-24-efficientnetb0-pytorch-train)\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\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 Inference Loop](#Functions-of-Inference-Loop)\n- [Inference & Submision](#Inference-&-Submision)\n\n## Update\n\n- V4: Inference of New Version (Train V3)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"# Import Packages","metadata":{}},{"cell_type":"code","source":"import re\nimport 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\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","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:02.599657Z","iopub.execute_input":"2024-04-15T02:20:02.599992Z","iopub.status.idle":"2024-04-15T02:20:08.827486Z","shell.execute_reply.started":"2024-04-15T02:20:02.599966Z","shell.execute_reply":"2024-04-15T02:20:08.826397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"class config:\n    \n    # == global config ==\n    SEED = 2024  # random seed\n    DEVICE = 'cpu'  # 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    Load_DATA = False  # 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 = 64  # batch size of each step\n    N_WORKERS = 4  # number of workers\n    \n    # == inference config ==\n    CKPT_ROOT = '/kaggle/input/birdclef24-baseline-checkpoints'\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-15T02:20:08.831653Z","iopub.execute_input":"2024-04-15T02:20:08.832226Z","iopub.status.idle":"2024-04-15T02:20:08.846137Z","shell.execute_reply.started":"2024-04-15T02:20:08.832195Z","shell.execute_reply":"2024-04-15T02:20:08.845446Z"},"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-15T02:20:08.847112Z","iopub.execute_input":"2024-04-15T02:20:08.847741Z","iopub.status.idle":"2024-04-15T02:20:08.854650Z","shell.execute_reply.started":"2024-04-15T02:20:08.847714Z","shell.execute_reply":"2024-04-15T02:20:08.853825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset & Dataloader","metadata":{}},{"cell_type":"markdown","source":"## Pre-Processing","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-15T02:20:08.855718Z","iopub.execute_input":"2024-04-15T02:20:08.856419Z","iopub.status.idle":"2024-04-15T02:20:08.863794Z","shell.execute_reply.started":"2024-04-15T02:20:08.856392Z","shell.execute_reply":"2024-04-15T02:20:08.862755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_bird_data = dict()\n\n# https://www.kaggle.com/code/markwijkhuizen/birdclef-2024-efficientvit-inference\nif len(glob(f'{config.DATA_ROOT}/test_soundscapes/*.ogg')) > 0:\n    ogg_file_paths = glob(f'{config.DATA_ROOT}/test_soundscapes/*.ogg')\nelse:\n    ogg_file_paths = sorted(glob(f'{config.DATA_ROOT}/unlabeled_soundscapes/*.ogg'))[:10]\n\nfor i, file_path in tqdm(enumerate(ogg_file_paths)):\n    row_id = re.search(r'/([^/]+)\\.ogg$', file_path).group(1)  # filename\n    audio_data, _ = librosa.load(file_path, sr=config.FS)\n    \n    # to spec.\n    spec = oog2spec_via_scipy(audio_data)\n    \n    # pad\n    pad = 512 - (spec.shape[1] % 512)\n    if pad > 0:\n        spec = np.pad(spec, ((0,0), (0,pad)))\n    \n    # reshape\n    spec = spec.reshape(512,-1,512).transpose([0, 2, 1])\n    spec = cv2.resize(spec, (256, 256), interpolation=cv2.INTER_AREA)\n    \n    for j in range(48):\n        all_bird_data[f'{row_id}_{(j+1)*5}'] = spec[:, :, j]","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:08.866934Z","iopub.execute_input":"2024-04-15T02:20:08.867294Z","iopub.status.idle":"2024-04-15T02:20:21.532063Z","shell.execute_reply.started":"2024-04-15T02:20:08.867265Z","shell.execute_reply":"2024-04-15T02:20:21.531044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class BirdDataset(torch.utils.data.Dataset):\n    \n    def __init__(\n        self,\n        bird_data,\n        augmentation=None,\n    ):\n        super().__init__()\n        self.bird_data = bird_data\n        self.keys_list = list(bird_data.keys())\n        self.augmentation = augmentation\n    \n    def __len__(self):\n        return len(self.bird_data)\n    \n    def __getitem__(self, index):\n        \n        _spec = self.bird_data[self.keys_list[index]]\n        \n        if self.augmentation is not None:\n            _spec = self.augmentation(image=_spec)['image'] \n        \n        return torch.tensor(_spec, dtype=torch.float32)","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:21.533123Z","iopub.execute_input":"2024-04-15T02:20:21.534294Z","iopub.status.idle":"2024-04-15T02:20:21.540193Z","shell.execute_reply.started":"2024-04-15T02:20:21.534264Z","shell.execute_reply":"2024-04-15T02:20:21.539514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Augmentation","metadata":{}},{"cell_type":"code","source":"def get_transforms(_type):\n    \n    if _type == 'test':\n        return albu.Compose([])","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:21.541182Z","iopub.execute_input":"2024-04-15T02:20:21.541638Z","iopub.status.idle":"2024-04-15T02:20:21.552403Z","shell.execute_reply.started":"2024-04-15T02:20:21.541613Z","shell.execute_reply":"2024-04-15T02:20:21.551722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Verify","metadata":{}},{"cell_type":"code","source":"def show_batch(ds, row=2, col=2):\n    fig = plt.figure(figsize=(6, 6))\n    img_index = np.random.randint(0, len(ds)-1, row*col)\n    \n    for i in range(len(img_index)):\n        img = dummy_dataset[img_index[i]]\n        \n        if isinstance(img, torch.Tensor):\n            img = img.detach().numpy()\n        \n        ax = fig.add_subplot(2, 2, i + 1, xticks=[], yticks=[])\n        ax.imshow(img, cmap='jet')\n        ax.set_title(f'ID: {img_index[i]}')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:21.553885Z","iopub.execute_input":"2024-04-15T02:20:21.554153Z","iopub.status.idle":"2024-04-15T02:20:21.562778Z","shell.execute_reply.started":"2024-04-15T02:20:21.554129Z","shell.execute_reply":"2024-04-15T02:20:21.561868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dummy_dataset = BirdDataset(all_bird_data, get_transforms('test'))\n\ntest_input = 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-15T02:20:21.564069Z","iopub.execute_input":"2024-04-15T02:20:21.564350Z","iopub.status.idle":"2024-04-15T02:20:22.360860Z","shell.execute_reply.started":"2024-04-15T02:20:21.564328Z","shell.execute_reply":"2024-04-15T02:20:22.360084Z"},"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=False):\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-15T02:20:22.362046Z","iopub.execute_input":"2024-04-15T02:20:22.362817Z","iopub.status.idle":"2024-04-15T02:20:22.371940Z","shell.execute_reply.started":"2024-04-15T02:20:22.362787Z","shell.execute_reply":"2024-04-15T02:20:22.371007Z"},"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        return {'val_loss': val_loss, 'val_score': val_score}","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:22.373433Z","iopub.execute_input":"2024-04-15T02:20:22.374059Z","iopub.status.idle":"2024-04-15T02:20:22.390051Z","shell.execute_reply.started":"2024-04-15T02:20:22.374031Z","shell.execute_reply":"2024-04-15T02:20:22.388985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions of Inference Loop","metadata":{}},{"cell_type":"code","source":"def predict(data_loader, model):\n    model.to(config.DEVICE)\n    model.eval()\n    pred = []\n    for batch in tqdm(data_loader):\n        with torch.no_grad():\n            x = batch\n            outputs = model(x)\n            outputs = nn.Softmax(dim=1)(outputs)\n        pred.append(outputs.detach().cpu())\n    \n    pred = torch.cat(pred, dim=0).cpu().detach()\n    \n    return pred.numpy().astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:22.391440Z","iopub.execute_input":"2024-04-15T02:20:22.392033Z","iopub.status.idle":"2024-04-15T02:20:22.402593Z","shell.execute_reply.started":"2024-04-15T02:20:22.391999Z","shell.execute_reply":"2024-04-15T02:20:22.401429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference & Submision","metadata":{}},{"cell_type":"markdown","source":"## ckpts","metadata":{}},{"cell_type":"code","source":"# ckpt_list = glob(f'{config.CKPT_ROOT}/*.ckpt')\n# print(f'find {len(ckpt_list)} ckpts in {config.CKPT_ROOT}.')\n\nckpt_list = [f'/kaggle/input/birdclef24-baseline-checkpoints/fold_0.ckpt']","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:22.404267Z","iopub.execute_input":"2024-04-15T02:20:22.404732Z","iopub.status.idle":"2024-04-15T02:20:22.411839Z","shell.execute_reply.started":"2024-04-15T02:20:22.404687Z","shell.execute_reply":"2024-04-15T02:20:22.410807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Main Loop","metadata":{}},{"cell_type":"code","source":"predictions = []\n\nfor ckpt in ckpt_list:\n    \n    # == init model ==\n    bird_model = BirdModel()\n    \n    # == load ckpt ==\n    weights = torch.load(ckpt, map_location=torch.device('cpu'))['state_dict']\n    bird_model.load_state_dict(weights)\n    \n    # == create dataset & dataloader ==\n    test_dataset = BirdDataset(all_bird_data, get_transforms('test'))\n    test_loader = torch.utils.data.DataLoader(\n        test_dataset,\n        batch_size=config.BATCH_SIZE,\n        num_workers=config.N_WORKERS,\n        shuffle=False,\n        drop_last=False\n    )\n    \n    predictions.append(predict(test_loader, bird_model))\n    gc.collect()\n\npredictions = np.mean(predictions, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:22.415968Z","iopub.execute_input":"2024-04-15T02:20:22.416271Z","iopub.status.idle":"2024-04-15T02:20:45.394894Z","shell.execute_reply.started":"2024-04-15T02:20:22.416246Z","shell.execute_reply":"2024-04-15T02:20:45.393826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_pred = pd.DataFrame(predictions, columns=label_list)\nsub_id = pd.DataFrame({'row_id': list(all_bird_data.keys())})\n\nsub = pd.concat([sub_id, sub_pred], axis=1)\n\nsub.to_csv('submission.csv',index=False)\nprint(f'Submissionn shape: {sub.shape}')\nsub.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-04-15T02:20:45.396148Z","iopub.execute_input":"2024-04-15T02:20:45.396441Z","iopub.status.idle":"2024-04-15T02:20:45.555213Z","shell.execute_reply.started":"2024-04-15T02:20:45.396415Z","shell.execute_reply":"2024-04-15T02:20:45.553215Z"},"trusted":true},"execution_count":null,"outputs":[]}]}