{"metadata":{"kaggle":{"accelerator":"none","dataSources":[{"sourceId":44224,"databundleVersionId":5188730,"sourceType":"competition"},{"sourceId":5148212,"sourceType":"datasetVersion","datasetId":2991134},{"sourceId":121796742,"sourceType":"kernelVersion"},{"sourceId":160017136,"sourceType":"kernelVersion"}],"dockerImageVersionId":30635,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"papermill":{"default_parameters":{},"duration":58.402956,"end_time":"2024-01-22T18:46:19.721856","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-01-22T18:45:21.318900","version":"2.4.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Welcome to my Kaggle inference notebook! In this notebook, I'll guide you through the process of making predictions using a model trained for the BIRDClef 2023 competition. The model has been trained using supervised contrastive learning and implemented with PyTorch Lightning.\n\n## Notebook Overview\n\n- **Inference Approach:** Leveraging the trained model to make predictions on new data.\n- **Framework:** PyTorch Lightning is utilized to streamline the inference process.\n\n## How to Use This Notebook\n\nFeel free to explore the inference code, try out predictions on different data, and adapt the model for your specific needs. If you find this notebook helpful, consider giving it an upvote to show your support.\n\nFor more detailed information on the training process and model architecture, check out my [BirdCLEF23 Supervised Contrastive Loss Training](https://www.kaggle.com/code/vijayravichander/supervised-contrastive-learning-pytorchlightning).\n\n\n**Notebook Credits:**\n- This notebook builds upon the work of [Nischay Dhankhar](https://www.kaggle.com/nischaydnk/). I have adapted and extended their code for the inference phase of the BIRDClef 2023 competition.\n\n## Upvote if Useful\n\nIf you find this inference notebook valuable, please consider giving it an upvote. Your feedback and support are highly appreciated!\n\nHappy predicting!\n","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn.functional as F\nfrom torchvision import datasets, transforms\nimport torchvision.transforms as transforms\nfrom torch import nn, optim\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nfrom torch.utils.data import random_split\nimport pytorch_lightning as pl\nimport torchmetrics\nfrom torchmetrics import Metric\nimport timm\nimport torch.nn as nn\nimport pandas as pd\nimport numpy as np\nfrom pathlib import Path","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2024-01-22T18:45:24.461212Z","iopub.status.busy":"2024-01-22T18:45:24.460762Z","iopub.status.idle":"2024-01-22T18:45:34.380681Z","shell.execute_reply":"2024-01-22T18:45:34.379635Z"},"papermill":{"duration":9.930957,"end_time":"2024-01-22T18:45:34.383415","exception":false,"start_time":"2024-01-22T18:45:24.452458","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    num_classes = 264\n    batch_size = 12\n    PRECISION = 16    \n    seed = 2023\n    model = \"resnet50\"\n    pretrained = False\n    use_mixup = False\n    mixup_alpha = 0.2   \n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')    \n\n    data_root = \"/kaggle/input/birdclef-2023/\"\n    train_images = \"/kaggle/input/split-creating-melspecs-stage-1/specs/train/\"\n    valid_images = \"/kaggle/input/split-creating-melspecs-stage-1/specs/valid/\"\n    train_path = \"/kaggle/input/bc2023-train-val-df/train.csv\"\n    valid_path = \"/kaggle/input/bc2023-train-val-df/valid.csv\"\n    \n    test_path = '/kaggle/input/birdclef-2023/test_soundscapes/'\n    SR = 32000\n    DURATION = 5\n    LR = 5e-4\n    \n    model_ckpt = '/kaggle/input/birdclef23-supervised-contrastive-loss-training/birdclef_supconmodel.ckpt'","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.397653Z","iopub.status.busy":"2024-01-22T18:45:34.397081Z","iopub.status.idle":"2024-01-22T18:45:34.404426Z","shell.execute_reply":"2024-01-22T18:45:34.403378Z"},"papermill":{"duration":0.016909,"end_time":"2024-01-22T18:45:34.406624","exception":false,"start_time":"2024-01-22T18:45:34.389715","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pl.seed_everything(Config.seed, workers=True)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.420258Z","iopub.status.busy":"2024-01-22T18:45:34.419828Z","iopub.status.idle":"2024-01-22T18:45:34.435136Z","shell.execute_reply":"2024-01-22T18:45:34.434029Z"},"papermill":{"duration":0.025021,"end_time":"2024-01-22T18:45:34.437685","exception":false,"start_time":"2024-01-22T18:45:34.412664","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def config_to_dict(cfg):\n    return dict((name, getattr(cfg, name)) for name in dir(cfg) if not name.startswith('__'))","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.452448Z","iopub.status.busy":"2024-01-22T18:45:34.451992Z","iopub.status.idle":"2024-01-22T18:45:34.457395Z","shell.execute_reply":"2024-01-22T18:45:34.456398Z"},"papermill":{"duration":0.015777,"end_time":"2024-01-22T18:45:34.459960","exception":false,"start_time":"2024-01-22T18:45:34.444183","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_melspec(y, sr, n_mels, fmin, fmax):\n    \"\"\"\n    Computes a mel-spectrogram and puts it at decibel scale\n    Arguments:\n        y {np array} -- signal\n        params {AudioParams} -- Parameters to use for the spectrogram. Expected to have the attributes sr, n_mels, f_min, f_max\n    Returns:\n        np array -- Mel-spectrogram\n    \"\"\"\n    melspec = lb.feature.melspectrogram(\n        y=y, sr=sr, n_mels=n_mels, fmin=fmin, fmax=fmax,\n    )\n\n    melspec = lb.power_to_db(melspec).astype(np.float32)\n    return melspec\n\ndef mono_to_color(X, eps=1e-6, mean=None, std=None):\n    mean = mean or X.mean()\n    std = std or X.std()\n    X = (X - mean) / (std + eps)\n    \n    _min, _max = X.min(), X.max()\n\n    if (_max - _min) > eps:\n        V = np.clip(X, _min, _max)\n        V = 255 * (V - _min) / (_max - _min)\n        V = V.astype(np.uint8)\n    else:\n        V = np.zeros_like(X, dtype=np.uint8)\n\n    return V\n\ndef crop_or_pad(y, length, is_train=True, start=None):\n    if len(y) < length:\n        y = np.concatenate([y, np.zeros(length - len(y))])\n        \n        n_repeats = length // len(y)\n        epsilon = length % len(y)\n        \n        y = np.concatenate([y]*n_repeats + [y[:epsilon]])\n        \n    elif len(y) > length:\n        if not is_train:\n            start = start or 0\n        else:\n            start = start or np.random.randint(len(y) - length)\n\n        y = y[start:start + length]\n\n    return y","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.474311Z","iopub.status.busy":"2024-01-22T18:45:34.473592Z","iopub.status.idle":"2024-01-22T18:45:34.483831Z","shell.execute_reply":"2024-01-22T18:45:34.482999Z"},"papermill":{"duration":0.019537,"end_time":"2024-01-22T18:45:34.485751","exception":false,"start_time":"2024-01-22T18:45:34.466214","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(Config.train_path)\nConfig.num_classes = len(df_train.primary_label.unique())","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.500881Z","iopub.status.busy":"2024-01-22T18:45:34.499793Z","iopub.status.idle":"2024-01-22T18:45:34.683871Z","shell.execute_reply":"2024-01-22T18:45:34.682663Z"},"papermill":{"duration":0.194555,"end_time":"2024-01-22T18:45:34.686460","exception":false,"start_time":"2024-01-22T18:45:34.491905","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for path in Path(Config.test_path).glob(\"*.ogg\"):\n    print(path)\n    print(path.stem)\n    print(path.stem.split(\"_\"))","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.702237Z","iopub.status.busy":"2024-01-22T18:45:34.701063Z","iopub.status.idle":"2024-01-22T18:45:34.717570Z","shell.execute_reply":"2024-01-22T18:45:34.716478Z"},"papermill":{"duration":0.027033,"end_time":"2024-01-22T18:45:34.719920","exception":false,"start_time":"2024-01-22T18:45:34.692887","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test = pd.DataFrame(\n     [(path.stem, *path.stem.split(\"_\"), path) for path in Path(Config.test_path).glob(\"*.ogg\")],\n    columns = [\"filename\", \"name\" ,\"id\", \"path\"]\n)\nprint(df_test.shape)\ndf_test.head()","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.735126Z","iopub.status.busy":"2024-01-22T18:45:34.734437Z","iopub.status.idle":"2024-01-22T18:45:34.752999Z","shell.execute_reply":"2024-01-22T18:45:34.752201Z"},"papermill":{"duration":0.028301,"end_time":"2024-01-22T18:45:34.754789","exception":false,"start_time":"2024-01-22T18:45:34.726488","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import librosa as lb\nimport librosa.display as lbd\nimport soundfile as sf\nfrom  soundfile import SoundFile \nfrom torch.utils.data import Dataset, DataLoader\n\nclass BirdDataset(Dataset):\n    def __init__(self, data, sr=Config.SR, n_mels=128, fmin=0, fmax=None, duration=Config.DURATION, step=None, res_type=\"kaiser_fast\", resample=True):\n        \n        self.data = data\n        \n        self.sr = sr\n        self.n_mels = n_mels\n        self.fmin = fmin\n        self.fmax = fmax or self.sr//2\n\n        self.duration = duration\n        self.audio_length = self.duration*self.sr\n        self.step = step or self.audio_length\n        \n        self.res_type = res_type\n        self.resample = resample\n\n    def __len__(self):\n        return len(self.data)\n    \n    @staticmethod\n    def normalize(image):\n        image = image.astype(\"float32\", copy=False) / 255.0\n        image = np.stack([image, image, image])\n        return image\n    \n    \n    def audio_to_image(self, audio):\n        melspec = compute_melspec(audio, self.sr, self.n_mels, self.fmin, self.fmax) \n        image = mono_to_color(melspec)\n        image = self.normalize(image)\n        return image\n\n    def read_file(self, filepath):\n        audio, orig_sr = sf.read(filepath, dtype=\"float32\")\n\n        if self.resample and orig_sr != self.sr:\n            audio = lb.resample(audio, orig_sr, self.sr, res_type=self.res_type)\n          \n        audios = []\n        for i in range(self.audio_length, len(audio) + self.step, self.step):\n            start = max(0, i - self.audio_length)\n            end = start + self.audio_length\n            audios.append(audio[start:end])\n            \n        if len(audios[-1]) < self.audio_length:\n            audios = audios[:-1]\n            \n        images = [self.audio_to_image(audio) for audio in audios]\n        images = np.stack(images)\n        \n        return images\n    \n        \n    def __getitem__(self, idx):\n        return self.read_file(self.data.loc[idx, \"path\"])","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.770130Z","iopub.status.busy":"2024-01-22T18:45:34.769251Z","iopub.status.idle":"2024-01-22T18:45:34.824708Z","shell.execute_reply":"2024-01-22T18:45:34.823543Z"},"papermill":{"duration":0.066558,"end_time":"2024-01-22T18:45:34.828032","exception":false,"start_time":"2024-01-22T18:45:34.761474","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_test = BirdDataset(\n    df_test, \n    sr = Config.SR,\n    duration = Config.DURATION,\n)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.843348Z","iopub.status.busy":"2024-01-22T18:45:34.842932Z","iopub.status.idle":"2024-01-22T18:45:34.848412Z","shell.execute_reply":"2024-01-22T18:45:34.846759Z"},"papermill":{"duration":0.015908,"end_time":"2024-01-22T18:45:34.850810","exception":false,"start_time":"2024-01-22T18:45:34.834902","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_test[0].shape","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:34.865806Z","iopub.status.busy":"2024-01-22T18:45:34.865436Z","iopub.status.idle":"2024-01-22T18:45:52.984805Z","shell.execute_reply":"2024-01-22T18:45:52.983744Z"},"papermill":{"duration":18.130836,"end_time":"2024-01-22T18:45:52.988201","exception":false,"start_time":"2024-01-22T18:45:34.857365","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sklearn.metrics\n\ndef padded_cmap(solution, submission, padding_factor=5):\n    solution = solution#.drop(['row_id'], axis=1, errors='ignore')\n    submission = submission#.drop(['row_id'], axis=1, errors='ignore')\n    new_rows = []\n    for i in range(padding_factor):\n        new_rows.append([1 for i in range(len(solution.columns))])\n    new_rows = pd.DataFrame(new_rows)\n    new_rows.columns = solution.columns\n    padded_solution = pd.concat([solution, new_rows]).reset_index(drop=True).copy()\n    padded_submission = pd.concat([submission, new_rows]).reset_index(drop=True).copy()\n    score = sklearn.metrics.average_precision_score(\n        padded_solution.values,\n        padded_submission.values,\n        average='macro',\n    )\n    return score","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:53.012285Z","iopub.status.busy":"2024-01-22T18:45:53.011474Z","iopub.status.idle":"2024-01-22T18:45:53.198629Z","shell.execute_reply":"2024-01-22T18:45:53.197365Z"},"papermill":{"duration":0.202009,"end_time":"2024-01-22T18:45:53.201256","exception":false,"start_time":"2024-01-22T18:45:52.999247","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# SupConLoss: https://github.com/HobbitLong/SupContrast/blob/master/losses.py\nclass SupConLoss(nn.Module):\n    \"\"\"Supervised Contrastive Learning: https://arxiv.org/pdf/2004.11362.pdf.\n    It also supports the unsupervised contrastive loss in SimCLR\"\"\"\n    def __init__(self, temperature=0.07, contrast_mode='all',\n                 base_temperature=0.07):\n        super(SupConLoss, self).__init__()\n        self.temperature = temperature\n        self.contrast_mode = contrast_mode\n        self.base_temperature = base_temperature\n\n    def forward(self, features, labels=None, mask=None):\n        \"\"\"Compute loss for model. If both `labels` and `mask` are None,\n        it degenerates to SimCLR unsupervised loss:\n        https://arxiv.org/pdf/2002.05709.pdf\n\n        Args:\n            features: hidden vector of shape [bsz, n_views, ...].\n            labels: ground truth of shape [bsz].\n            mask: contrastive mask of shape [bsz, bsz], mask_{i,j}=1 if sample j\n                has the same class as sample i. Can be asymmetric.\n        Returns:\n            A loss scalar.\n        \"\"\"\n        device = (torch.device('cuda')\n                  if features.is_cuda\n                  else torch.device('cpu'))\n\n        if len(features.shape) < 3:\n            raise ValueError('`features` needs to be [bsz, n_views, ...],'\n                             'at least 3 dimensions are required')\n        if len(features.shape) > 3:\n            features = features.view(features.shape[0], features.shape[1], -1)\n\n        batch_size = features.shape[0]\n        if labels is not None and mask is not None:\n            raise ValueError('Cannot define both `labels` and `mask`')\n        elif labels is None and mask is None:\n            mask = torch.eye(batch_size, dtype=torch.float32).to(device)\n        elif labels is not None:\n            labels = labels.contiguous().view(-1, 1)\n            if labels.shape[0] != batch_size:\n                raise ValueError('Num of labels does not match num of features')\n            mask = torch.eq(labels, labels.T).float().to(device)\n        else:\n            mask = mask.float().to(device)\n\n        contrast_count = features.shape[1]\n        contrast_feature = torch.cat(torch.unbind(features, dim=1), dim=0)\n        if self.contrast_mode == 'one':\n            anchor_feature = features[:, 0]\n            anchor_count = 1\n        elif self.contrast_mode == 'all':\n            anchor_feature = contrast_feature\n            anchor_count = contrast_count\n        else:\n            raise ValueError('Unknown mode: {}'.format(self.contrast_mode))\n\n        # compute logits\n        anchor_dot_contrast = torch.div(\n            torch.matmul(anchor_feature, contrast_feature.T),\n            self.temperature)\n        # for numerical stability\n        logits_max, _ = torch.max(anchor_dot_contrast, dim=1, keepdim=True)\n        logits = anchor_dot_contrast - logits_max.detach()\n\n        # tile mask\n        mask = mask.repeat(anchor_count, contrast_count)\n        # mask-out self-contrast cases\n        logits_mask = torch.scatter(\n            torch.ones_like(mask),\n            1,\n            torch.arange(batch_size * anchor_count).view(-1, 1).to(device),\n            0\n        )\n        mask = mask * logits_mask\n\n        # compute log_prob\n        exp_logits = torch.exp(logits) * logits_mask\n        log_prob = logits - torch.log(exp_logits.sum(1, keepdim=True))\n\n        # compute mean of log-likelihood over positive\n        # modified to handle edge cases when there is no positive pair\n        # for an anchor point.\n        # Edge case e.g.:-\n        # features of shape: [4,1,...]\n        # labels:            [0,1,1,2]\n        # loss before mean:  [nan, ..., ..., nan]\n        mask_pos_pairs = mask.sum(1)\n        mask_pos_pairs = torch.where(mask_pos_pairs < 1e-6, 1, mask_pos_pairs)\n        mean_log_prob_pos = (mask * log_prob).sum(1) / mask_pos_pairs\n\n        # loss\n        loss = - (self.temperature / self.base_temperature) * mean_log_prob_pos\n        loss = loss.view(anchor_count, batch_size).mean()\n\n        return loss","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:53.217563Z","iopub.status.busy":"2024-01-22T18:45:53.216832Z","iopub.status.idle":"2024-01-22T18:45:53.234959Z","shell.execute_reply":"2024-01-22T18:45:53.234025Z"},"papermill":{"duration":0.028948,"end_time":"2024-01-22T18:45:53.237359","exception":false,"start_time":"2024-01-22T18:45:53.208411","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Encoder(pl.LightningModule):\n    def __init__(self, model_name, emb_dim):\n        super().__init__()\n        self.backbone = timm.create_model(\"resnet50\", pretrained = False)\n        self.in_features = self.backbone.fc.in_features\n        self.backbone.fc = nn.Linear(self.in_features, emb_dim)\n        self.loss_fn = SupConLoss(0.07, 'one', 0.07)\n\n\n    def forward(self, x):\n        emb = self.backbone(x)\n        return emb\n\n    # Difference between Normal and Lightning: The train, valid and test steps is written here inside the class\n    def training_step(self, batch, batch_idx):\n        images, labels = batch\n\n        bsz = len(labels)\n\n        images = torch.cat([images[0], images[1]], dim=0)\n\n        #print(images.shape)\n\n        features = self.forward(images)\n\n        # Manipulating the features for SupConLoss\n        f1, f2 = torch.split(features, [bsz, bsz], dim=0)\n        features = torch.cat([f1.unsqueeze(1), f2.unsqueeze(1)], dim=1)\n\n        # Calculating SupConLoss\n        loss = self.loss_fn(features,labels)\n\n        return loss\n\n    # We have training_epoch_end function\n    # def on_train_epoch_end(self):\n    #     #print(\"Epoch Done\")\n\n    def validation_step(self, batch , batch_idx):\n\n        images, labels = batch\n\n        bsz = len(labels)\n\n        images = torch.cat([images[0], images[1]], dim=0)\n        \n        #print(images.shape)\n\n        features = self.forward(images)\n\n        # Manipulating the arrangment of features for SupConLoss\n        f1, f2 = torch.split(features, [bsz, bsz], dim=0)\n\n        features = torch.cat([f1.unsqueeze(1), f2.unsqueeze(1)], dim=1)\n\n        # Calculating SupConLoss\n        loss = self.loss_fn(features,labels)\n\n        return loss\n\n    def test_step(self, batch, batch_idx):\n        images, labels = batch\n        \n        bsz = len(labels)\n\n        images = torch.cat([images[0], images[1]], dim=0)\n        \n\n        features = self.forward(images)\n\n        # Manipulating the arrangment of features for SupConLoss\n        f1, f2 = torch.split(features, [bsz, bsz], dim=0)\n\n        features = torch.cat([f1.unsqueeze(1), f2.unsqueeze(1)], dim=1)\n\n        # Calculating SupConLoss\n        loss = self.loss_fn(features,labels)\n        return loss\n\n    # We can add schedulers to this method\n    def configure_optimizers(self):\n        return optim.Adam(self.parameters(), lr = 0.001)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:53.254355Z","iopub.status.busy":"2024-01-22T18:45:53.253235Z","iopub.status.idle":"2024-01-22T18:45:53.267338Z","shell.execute_reply":"2024-01-22T18:45:53.266207Z"},"papermill":{"duration":0.025684,"end_time":"2024-01-22T18:45:53.270017","exception":false,"start_time":"2024-01-22T18:45:53.244333","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport pickle\nimport pytorch_lightning as pl\nfrom torch.optim import Adam\n\n\n\nclass SupConCE(pl.LightningModule):\n    def __init__(self,):\n        super().__init__()\n        self.loss_fn = nn.CrossEntropyLoss()\n        self.accuracy = torchmetrics.Accuracy(task = 'multiclass', num_classes = 264)\n        self.f1_score = torchmetrics.F1Score(task = 'multiclass', num_classes = 264)\n        backbone = 'resnet50'\n        model_path = '/kaggle/input/birdclef23-supervised-contrastive-loss-training/birdclef_supconencoder.ckpt'\n        pretrained_model = Encoder.load_from_checkpoint(model_path, model_name = backbone, emb_dim = 128)\n\n\n        #Freezing all the encoder layers\n        for param in pretrained_model.parameters():\n            param.requires_grad = False\n\n\n        #Trainging only the last layer\n        pretrained_model.backbone.fc = nn.Linear(in_features=pretrained_model.backbone.fc.in_features, out_features=264)\n\n        pretrained_model.backbone.fc.requires_grad = True\n\n        self.model = pretrained_model\n\n\n    def forward(self, x):\n        logits = self.model(x)\n        return logits\n\n    def training_step(self, batch, batch_idx):\n        images, labels = batch\n        y_pred = self.forward(images)\n        loss = self.loss_fn(y_pred,labels)\n        accuracy = self.accuracy(y_pred,labels)\n        f1_score = self.f1_score(y_pred,labels)\n        self.log_dict({'train_loss': loss, 'train_accuracy': accuracy, 'train_f1_score': f1_score},\n                      on_step = False, on_epoch = True, prog_bar = True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        images, labels = batch\n        y_pred = self.forward(images)\n        loss = self.loss_fn(y_pred,labels)\n        accuracy = self.accuracy(y_pred,labels)\n        f1_score = self.f1_score(y_pred,labels)\n        \n        one_hot_target = F.one_hot(labels, num_classes=264)\n        \n        y_pred = pd.DataFrame(y_pred.cpu().detach().numpy())\n        y_true = pd.DataFrame(one_hot_target.cpu().detach().numpy())\n        \n        cmap_score = padded_cmap(y_true, y_pred)\n        \n        self.log_dict({'valid_loss': loss, 'valid_accuracy': accuracy, 'valid_f1_score': f1_score, 'cmap_score': cmap_score},\n                      on_step = False, on_epoch = True, prog_bar = True)\n        return loss\n\n    def configure_optimizers(self):\n        return optim.Adam(self.parameters(), lr = 0.001)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:53.285449Z","iopub.status.busy":"2024-01-22T18:45:53.285077Z","iopub.status.idle":"2024-01-22T18:45:53.297895Z","shell.execute_reply":"2024-01-22T18:45:53.296600Z"},"papermill":{"duration":0.023216,"end_time":"2024-01-22T18:45:53.300133","exception":false,"start_time":"2024-01-22T18:45:53.276917","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(data_loader, model):\n        \n    model.to('cpu')\n    model.eval()    \n    predictions = []\n    for en in range(len(ds_test)):\n        print(en)\n        images = torch.from_numpy(ds_test[en])\n        print(images.shape)\n        with torch.no_grad():\n            outputs = model(images).sigmoid().detach().cpu().numpy()\n            print(outputs.shape)\n#             pred_batch.extend(outputs.detach().cpu().numpy())\n#         pred_batch = np.vstack(pred_batch)\n        predictions.append(outputs)\n            \n    \n    return predictions","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:53.315472Z","iopub.status.busy":"2024-01-22T18:45:53.315078Z","iopub.status.idle":"2024-01-22T18:45:53.322248Z","shell.execute_reply":"2024-01-22T18:45:53.320670Z"},"papermill":{"duration":0.017912,"end_time":"2024-01-22T18:45:53.324823","exception":false,"start_time":"2024-01-22T18:45:53.306911","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\nprint(f\"Create Dataloader...\")\n\nds_test = BirdDataset(\n    df_test, \n    sr = Config.SR,\n    duration = Config.DURATION,\n)\n\n\naudio_model = SupConCE()\n\nprint(\"Model Creation\")\n\nmodel = SupConCE.load_from_checkpoint(Config.model_ckpt, train_dataloader=None,validation_dataloader=None) \nprint(\"Running Inference..\")\n\npreds = predict(ds_test, model)   \n\n#gc.collect()\n#torch.cuda.empty_cache()","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:45:53.340942Z","iopub.status.busy":"2024-01-22T18:45:53.340530Z","iopub.status.idle":"2024-01-22T18:46:16.393544Z","shell.execute_reply":"2024-01-22T18:46:16.392719Z"},"papermill":{"duration":23.063749,"end_time":"2024-01-22T18:46:16.395899","exception":false,"start_time":"2024-01-22T18:45:53.332150","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filenames = df_test.filename.values.tolist()\n\nbird_cols = list(pd.get_dummies(df_train['primary_label']).columns)\nsub_df = pd.DataFrame(columns=['row_id']+bird_cols)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:46:16.412803Z","iopub.status.busy":"2024-01-22T18:46:16.411593Z","iopub.status.idle":"2024-01-22T18:46:16.460355Z","shell.execute_reply":"2024-01-22T18:46:16.459053Z"},"papermill":{"duration":0.060172,"end_time":"2024-01-22T18:46:16.463190","exception":false,"start_time":"2024-01-22T18:46:16.403018","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:46:16.480018Z","iopub.status.busy":"2024-01-22T18:46:16.479607Z","iopub.status.idle":"2024-01-22T18:46:16.493985Z","shell.execute_reply":"2024-01-22T18:46:16.492945Z"},"papermill":{"duration":0.025774,"end_time":"2024-01-22T18:46:16.496810","exception":false,"start_time":"2024-01-22T18:46:16.471036","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, file in enumerate(filenames):\n    pred = preds[i]\n    num_rows = len(pred)\n    row_ids = [f'{file}_{(i+1)*5}' for i in range(num_rows)]\n    df = pd.DataFrame(columns=['row_id']+bird_cols)\n    \n    df['row_id'] = row_ids\n    df[bird_cols] = pred\n    \n    sub_df = pd.concat([sub_df,df]).reset_index(drop=True)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:46:16.513476Z","iopub.status.busy":"2024-01-22T18:46:16.512727Z","iopub.status.idle":"2024-01-22T18:46:16.587876Z","shell.execute_reply":"2024-01-22T18:46:16.587073Z"},"papermill":{"duration":0.086198,"end_time":"2024-01-22T18:46:16.590288","exception":false,"start_time":"2024-01-22T18:46:16.504090","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:46:16.606446Z","iopub.status.busy":"2024-01-22T18:46:16.605804Z","iopub.status.idle":"2024-01-22T18:46:16.633996Z","shell.execute_reply":"2024-01-22T18:46:16.632968Z"},"papermill":{"duration":0.039213,"end_time":"2024-01-22T18:46:16.636683","exception":false,"start_time":"2024-01-22T18:46:16.597470","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv('submission.csv',index=False)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T18:46:16.654051Z","iopub.status.busy":"2024-01-22T18:46:16.653416Z","iopub.status.idle":"2024-01-22T18:46:16.695191Z","shell.execute_reply":"2024-01-22T18:46:16.694208Z"},"papermill":{"duration":0.053371,"end_time":"2024-01-22T18:46:16.697760","exception":false,"start_time":"2024-01-22T18:46:16.644389","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.007223,"end_time":"2024-01-22T18:46:16.712630","exception":false,"start_time":"2024-01-22T18:46:16.705407","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}