{"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":8108072,"sourceType":"datasetVersion","datasetId":4789213},{"sourceId":8120971,"sourceType":"datasetVersion","datasetId":4779991},{"sourceId":8136153,"sourceType":"datasetVersion","datasetId":4809690},{"sourceId":8141043,"sourceType":"datasetVersion","datasetId":4813248},{"sourceId":8143129,"sourceType":"datasetVersion","datasetId":4814832},{"sourceId":8146572,"sourceType":"datasetVersion","datasetId":4817505},{"sourceId":8150382,"sourceType":"datasetVersion","datasetId":4820295},{"sourceId":172220540,"sourceType":"kernelVersion"}],"dockerImageVersionId":30684,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BirdCLEF 2024 [Inference with ONNX]\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\n     * Train: [BirdCLEF'24 | ConvNeXt PyTorch [Train]](https://www.kaggle.com/code/kmatsu01/birdclef-24-convnext-pytorch-train/)\n     * Infer: This notebook\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 the ConvNeXt model</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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"# Install ONNX","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/onnxruntime/humanfriendly-10.0-py2.py3-none-any.whl --no-index --find-links /kaggle/input/onnxruntime\n!pip install /kaggle/input/onnxruntime/coloredlogs-15.0.1-py2.py3-none-any.whl --no-index --find-links /kaggle/input/onnxruntime\n!pip install /kaggle/input/onnxruntime/onnxruntime-1.17.3-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl --no-index --find-links /kaggle/input/onnxruntime","metadata":{"execution":{"iopub.status.busy":"2024-04-17T23:43:08.410534Z","iopub.execute_input":"2024-04-17T23:43:08.411217Z","iopub.status.idle":"2024-04-17T23:43:54.572180Z","shell.execute_reply.started":"2024-04-17T23:43:08.411182Z","shell.execute_reply":"2024-04-17T23:43:54.571005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\nimport scipy\nfrom scipy import signal as sci_signal\n\nimport torch\nfrom torch import nn\nfrom torchvision.models import efficientnet\nfrom torchvision.models import convnext\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-17T23:43:54.574225Z","iopub.execute_input":"2024-04-17T23:43:54.574554Z","iopub.status.idle":"2024-04-17T23:44:07.868872Z","shell.execute_reply.started":"2024-04-17T23:43:54.574524Z","shell.execute_reply":"2024-04-17T23:44:07.867745Z"},"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_b1'  # model type\n    MODEL_TYPE = 'convnext_tiny'  #\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/convnext-tiny-24017/20240417203444_0003_convnext_tiny'\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-17T23:44:07.870798Z","iopub.execute_input":"2024-04-17T23:44:07.871523Z","iopub.status.idle":"2024-04-17T23:44:07.888500Z","shell.execute_reply.started":"2024-04-17T23:44:07.871482Z","shell.execute_reply":"2024-04-17T23:44:07.887321Z"},"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-17T23:44:07.891278Z","iopub.execute_input":"2024-04-17T23:44:07.892036Z","iopub.status.idle":"2024-04-17T23:44:07.910942Z","shell.execute_reply.started":"2024-04-17T23:44:07.892004Z","shell.execute_reply":"2024-04-17T23:44:07.909805Z"},"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-17T23:44:07.912980Z","iopub.execute_input":"2024-04-17T23:44:07.914096Z","iopub.status.idle":"2024-04-17T23:44:07.923504Z","shell.execute_reply.started":"2024-04-17T23:44:07.914037Z","shell.execute_reply":"2024-04-17T23:44:07.922139Z"},"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-17T23:44:07.925211Z","iopub.execute_input":"2024-04-17T23:44:07.925698Z","iopub.status.idle":"2024-04-17T23:44:32.573832Z","shell.execute_reply.started":"2024-04-17T23:44:07.925667Z","shell.execute_reply":"2024-04-17T23:44:32.572505Z"},"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":{"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":{"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":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"markdown","source":"## Network","metadata":{}},{"cell_type":"code","source":"class ConvNeXtModel(nn.Module):\n    def __init__(self, model_type, n_classes, pretrained=False):\n        super().__init__()\n        if model_type == 'convnext_tiny':\n            if pretrained: weights = convnext.ConvNeXt_Tiny_Weights.DEFAULT\n            else: weights = None\n            self.base_model = convnext.convnext_tiny(weights=weights)\n        elif model_type == 'convnext_small':\n            if pretrained: weights = convnext.ConvNeXt_Small_Weights.DEFAULT\n            else: weights = None\n            self.base_model = convnext.convnext_small(weights=weights)\n        elif model_type == 'convnext_base':\n            if pretrained: weights = convnext.ConvNeXt_Base_Weights.DEFAULT\n            else: weights = None\n            self.base_model = convnext.convnext_base(weights=weights)\n        else:\n            raise ValueError('model type not supported')\n\n        num_features = self.base_model.classifier[2].in_features\n        self.base_model.classifier[2] = nn.Linear(num_features, n_classes, bias=True, dtype=torch.float32)\n    \n    def forward(self, x):\n        x = x.unsqueeze(1)\n        x = x.expand(-1, 3, -1, -1)\n        return self.base_model(x)","metadata":{"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 = ConvNeXtModel(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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ONNX","metadata":{}},{"cell_type":"markdown","source":"## ckpts","metadata":{}},{"cell_type":"code","source":"ckpt_list = [f'/kaggle/input/convnext-tiny-24017/20240417203444_0003_convnext_tiny/fold_0.ckpt']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config of ONNX","metadata":{}},{"cell_type":"code","source":"# == setting of onnx ==\n\ninput_tensor = torch.randn(config.BATCH_SIZE, 256, 256)  # input shape\ninput_names = ['x']\noutput_names = ['output']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Convert","metadata":{}},{"cell_type":"code","source":"onnx_ckpt_list = list()\nfor ckpt_path in ckpt_list:\n    ckpt_name = os.path.basename(ckpt_path).split('.')[0]\n    # == init model ==\n    bird_model = BirdModel()\n    \n    # == load ckpt ==\n    weights = torch.load(ckpt_path, map_location=torch.device('cpu'))['state_dict']\n    bird_model.load_state_dict(weights)\n    bird_model.eval()\n    \n    # == convert to onnx ==\n    torch.onnx.export(bird_model.backbone, input_tensor, f\"{ckpt_name}.onnx\", verbose=False, input_names=input_names, output_names=output_names)\n    onnx_ckpt_list.append(f\"{ckpt_name}.onnx\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions of Inference Loop","metadata":{}},{"cell_type":"code","source":"def predict(data_loader, onnx_model):\n    pred = []\n    for batch in tqdm(data_loader):\n        with torch.no_grad():\n            x = batch\n            n_pad = 0\n            \n            # == make sure the batch_size equal to setting\n            if x.shape[0] < config.BATCH_SIZE:\n                n_pad = config.BATCH_SIZE - x.shape[0]\n                zero_tensor = torch.zeros((n_pad, 256, 256))\n                x = torch.cat([x, zero_tensor], dim=0)\n            \n            outputs = onnx_model.run(output_names, {input_names[0]: x.numpy()})[0]\n            outputs = scipy.special.softmax(outputs[:config.BATCH_SIZE-n_pad, ...], axis=1)\n        pred.append(outputs)\n    \n    return np.concatenate(pred, axis=0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference & Submision","metadata":{}},{"cell_type":"markdown","source":"## Main Loop","metadata":{}},{"cell_type":"code","source":"import onnx\nimport onnxruntime as ort","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\n\nfor ckpt in onnx_ckpt_list:\n    \n    # == init ONNX model ==\n    onnx_model = onnx.load(ckpt)\n    onnx_model_graph = onnx_model.graph\n    onnx_session = ort.InferenceSession(onnx_model.SerializeToString())\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, onnx_session))\n    gc.collect()\n\npredictions = np.mean(predictions, axis=0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"about x2~3 faster than PyTorch","metadata":{}},{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}