{"metadata":{"kernelspec":{"display_name":"Python 3 (ipykernel)","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.16"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BirdSet Fine-Tuning Tutorial [Copied from BirdSet Repository](https://github.com/DBD-research-group/BirdSet)\n\nThis notebook details how to fine-tune models on BirdSet data.","metadata":{}},{"cell_type":"markdown","source":"## Helper Functions\n\nThese functions are just helper functions for the code below. These can be ignored for now and looked up when they are used in the notebook.  \nNot all BirdSet code used, is turned into helper functions for this notebook as there are too many to do that, however, the most abstracting parts are.","metadata":{}},{"cell_type":"markdown","source":"### Event Limiting\n\nThis function limits event extraction per specification. If needed it can limit the extracted events per class, per extracted file, or of both.","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport random\nfrom tqdm import tqdm\nfrom collections import Counter\n\n\ndef smart_sampling(dataset, label_name, class_limit, event_limit):\n    def _unique_identifier(x, labelname):\n        file = x[\"filepath\"]\n        label = x[labelname]\n        return {\"id\": f\"{file}-{label}\"}\n\n    class_limit = class_limit if class_limit else -float(\"inf\")\n    dataset = dataset.map(\n        lambda x: _unique_identifier(x, label_name), desc=\"sampling: unique-identifier\"\n    )\n    df = pd.DataFrame(dataset)\n    path_label_count = df.groupby([\"id\", label_name], as_index=False).size()\n    path_label_count = path_label_count.set_index(\"id\")\n    class_sizes = df.groupby(label_name).size()\n\n    for label in tqdm(class_sizes.index, desc=\"sampling\"):\n        current = path_label_count[path_label_count[label_name] == label]\n        total = current[\"size\"].sum()\n        most = current[\"size\"].max()\n\n        while total > class_limit or most != event_limit:\n            largest_count = current[\"size\"].value_counts()[current[\"size\"].max()]\n            n_largest = current.nlargest(largest_count + 1, \"size\")\n            to_del = n_largest[\"size\"].max() - n_largest[\"size\"].min()\n\n            idxs = n_largest[n_largest[\"size\"] == n_largest[\"size\"].max()].index\n            if (\n                total - (to_del * largest_count) < class_limit\n                or most == event_limit\n                or most == 1\n            ):\n                break\n            for idx in idxs:\n                current.at[idx, \"size\"] = current.at[idx, \"size\"] - to_del\n                path_label_count.at[idx, \"size\"] = (\n                    path_label_count.at[idx, \"size\"] - to_del\n                )\n\n            total = current[\"size\"].sum()\n            most = current[\"size\"].max()\n\n    event_counts = Counter(dataset[\"id\"])\n\n    all_file_indices = {label: [] for label in event_counts.keys()}\n    for idx, label in enumerate(dataset[\"id\"]):\n        all_file_indices[label].append(idx)\n\n    limited_indices = []\n    for file, indices in all_file_indices.items():\n        limit = path_label_count.loc[file][\"size\"]\n        limited_indices.extend(random.sample(indices, limit))\n\n    dataset = dataset.remove_columns(\"id\")\n    return dataset.select(limited_indices)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### One-Hot Encoding\n\nThis function one hot encodes the given labels. This is needed for BirdSet's multilabel training as a `MultilabelMix` transform gets applied that mixes the labels and input values of some inputs into one to simulate multilabel data. Without that multilabel data is only present in the `test_5s` splits of BirdSet's datasets.","metadata":{}},{"cell_type":"code","source":"import torch\n\n\ndef classes_one_hot(batch, num_classes):\n    \"\"\"\n    Converts class labels to one-hot encoding.\n\n    This method takes a batch of data and converts the class labels to one-hot encoding.\n    The one-hot encoding is a binary matrix representation of the class labels.\n\n    Args:\n        batch (dict): A batch of data. The batch should be a dictionary where the keys are the field names and the values are the field data.\n\n    Returns:\n        dict: The batch with the \"labels\" field converted to one-hot encoding. The keys are the field names and the values are the field data.\n    \"\"\"\n    label_list = [y for y in batch[\"labels\"]]\n    class_one_hot_matrix = torch.zeros(\n        (len(label_list), num_classes), dtype=torch.float\n    )\n\n    for class_idx, idx in enumerate(label_list):\n        class_one_hot_matrix[class_idx, idx] = 1\n\n    class_one_hot_matrix = torch.tensor(class_one_hot_matrix, dtype=torch.float32)\n    return {\"labels\": class_one_hot_matrix}","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Basic Transforms Wrapper\n\nThis class will help in handling the application of transforms on the data as well as converting the waveforms to spectrogramms. Transforms are stored in a list and then applied in the list order when a call to this class is made.","metadata":{}},{"cell_type":"code","source":"from dataclasses import dataclass\nfrom torchaudio.transforms import Spectrogram, MelScale\nfrom birdset.datamodule.components.resize import Resizer\nfrom birdset.datamodule.components.augmentations import PowerToDB\nimport numpy as np\nimport torch_audiomentations\nimport torchvision.transforms\n\n\n@dataclass\nclass CustomProcessingConfig:\n    spectrogram_conversion: Spectrogram | None = Spectrogram(\n        n_fft=1024,\n        hop_length=320,\n        power=2.0,\n    )\n    resizer: Resizer | None = (\n        Resizer(\n            db_scale=True,\n        ),\n    )\n    melscale_conversion: MelScale | None = (\n        MelScale(\n            n_mels=128,\n            sample_rate=32000,\n            n_stft=513,  # n_fft//2+1\n        ),\n    )\n    dbscale_conversion: PowerToDB | None = (PowerToDB(),)\n    normalize_spectrogram: bool = (True,)\n    mean: float = (-4.268,)\n    std: float = (-4.569,)\n\n\nclass BasicTransformsWrapper:\n    def __init__(\n        self,\n        wav_transforms,\n        spec_transforms,\n        decoding,\n        feature_extractor,\n        nocall_sampler,\n        processing_config=CustomProcessingConfig,\n    ):\n        self.wav_transforms = torch_audiomentations.Compose(\n            transforms=wav_transforms, output_type=\"object_dict\"\n        )\n        self.spec_transforms = torchvision.transforms.Compose(\n            transforms=spec_transforms\n        )\n        self.processing = processing_config\n        self.nocall_sampler = nocall_sampler\n        self.decoding = decoding\n        self.feature_extractor = feature_extractor\n\n    def __call__(self, batch, **kwds):\n\n        batch = self.decoding(batch)\n\n        waveform_batch = self._get_waveform_batch(batch)\n\n        input_values = waveform_batch[\"input_values\"]\n        input_values = input_values.unsqueeze(1)\n        labels = torch.tensor(batch[\"labels\"])\n\n        input_values, labels = self._waveform_augmentation(input_values, labels)\n\n        if self.nocall_sampler:\n            input_values, labels = self.nocall_sampler(input_values, labels)\n\n        if self.processing.spectrogram_conversion is not None:\n            spectrograms = self.processing.spectrogram_conversion(input_values)\n\n            if self.spec_transforms:\n                spectrograms = self.spec_transforms(spectrograms)\n\n            if self.processing.melscale_conversion:\n                spectrograms = self.processing.melscale_conversion(spectrograms)\n\n            if self.processing.dbscale_conversion:\n                spectrograms = self.processing.dbscale_conversion(spectrograms)\n\n            if self.processing.resizer:\n                spectrograms = self.processing.resizer.resize_spectrogram_batch(\n                    spectrograms\n                )\n\n            if self.processing.normalize_spectrogram:\n                spectrograms = (\n                    spectrograms - self.processing.mean\n                ) / self.processing.std\n\n            input_values = spectrograms\n\n        # values in labels need to be of type float for further use\n        labels = labels.to(torch.float16)\n\n        return {\"input_values\": input_values, \"labels\": labels}\n\n    def _get_waveform_batch(self, batch):\n        waveform_batch = [audio[\"array\"] for audio in batch[\"audio\"]]\n\n        # extract/pad/truncate\n        max_length = int(\n            int(self.feature_extractor.sampling_rate) * int(self.decoding.max_len)\n        )\n        waveform_batch = self.feature_extractor(\n            waveform_batch,\n            padding=\"max_length\",\n            max_length=max_length,\n            truncation=True,\n            return_attention_mask=True,\n        )\n\n        return waveform_batch\n\n    def _waveform_augmentation(self, input_values, labels):\n        labels = labels.unsqueeze(1).unsqueeze(1)\n        output_dict = self.wav_transforms(\n            samples=input_values,\n            sample_rate=self.feature_extractor.sampling_rate,\n            targets=labels,\n        )\n        labels = output_dict.targets.squeeze(1).squeeze(1)\n\n        return output_dict.samples, labels","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Custom Datamodule\n\nWhen using a lightning trainer you can either pass a Datamodule or the Dataloaders needed for the planed task. Typically a Datamodule would load the data, preprocess it and apply transforms. In this case the datamodule will only hold the data and provide Dataloaders. The other steps are done separately so as to not abstract each step too much.  \nFor a more typical implementation of Datamodules you should look into the `BirdSetDataModule` in `birdset/datamodules/birdset_datamodule.py`.","metadata":{}},{"cell_type":"code","source":"import lightning as L\nfrom torch.utils.data import DataLoader\n\n\nclass CustomDatamodule(L.LightningDataModule):\n    def __init__(self, dataset, batch_size, num_workers, num_classes, task):\n        super().__init__()\n        self.dataset = dataset\n        self.batch_size = batch_size\n        self.train_batch_size = batch_size\n        self.num_workers = num_workers\n        self.num_classes = num_classes\n        self.task = task\n        self.len_trainset = len(dataset[\"train\"])\n\n    def setup(self, stage):\n        pass\n\n    def train_dataloader(self):\n        return DataLoader(\n            dataset=self.dataset[\"train\"],\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n        )\n\n    def val_dataloader(self):\n        return DataLoader(\n            dataset=self.dataset[\"valid\"],\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n        )\n\n    def test_dataloader(self):\n        return DataLoader(\n            dataset=self.dataset[\"test\"],\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n        )","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Custom Lightning Module\n\nThis snippet defines a custom LightningModule. It is needed as a wrapper for models loaded from Hugging Face if you want to train using the Lightning trainer, as it expects an instance of `Lightning.LightningModule` and models loaded from Hugging Face are not.","metadata":{}},{"cell_type":"code","source":"from transformers import ASTForAudioClassification, ConvNextForImageClassification\nfrom torchmetrics import AUROC, MetricCollection\nfrom birdset.modules.metrics.multilabel import TopKAccuracy, cmAP\nimport lightning as l\nimport torch.nn as nn\nimport torch\nfrom transformers import AdamW\n\n\nclass ConvNextClassifierLightningModule(l.LightningModule):\n    def __init__(\n        self,\n        num_classes,\n        num_epochs,\n    ):\n        super(ConvNextClassifierLightningModule, self).__init__()\n        self.model = ConvNextForImageClassification.from_pretrained(\n            \"DBD-research-group/ConvNeXT-Base-BirdSet-XCL\",\n            num_labels=num_classes,\n            ignore_mismatched_sizes=True,\n        )\n\n        self.num_classes = num_classes\n        self.num_epochs = num_epochs\n        self.loss = nn.BCEWithLogitsLoss()\n        self.main_metric = cmAP(num_labels=num_classes, thresholds=None)\n        self.other_metrics = MetricCollection(\n            {\n                \"MultilabelAUROC\": AUROC(\n                    task=\"multilabel\",\n                    num_labels=num_classes,\n                    average=\"macro\",\n                    thresholds=None,\n                ),\n                \"T1Accuracy\": TopKAccuracy(topk=1),\n            }\n        )\n\n    def forward(self, pixel_values):\n        outputs = self.model(pixel_values=pixel_values)\n        return outputs.logits\n\n    def common_step(self, batch, batch_idx):\n        values = batch[\"input_values\"]\n        labels = batch[\"labels\"]\n        logits = self(values)\n\n        loss = self.loss(logits, labels)\n        predictions = torch.sigmoid(logits)\n\n        return loss, predictions\n\n    def training_step(self, batch, batch_idx):\n        loss, preds = self.common_step(batch, batch_idx)\n        self.log(\n            f\"train/{self.loss.__class__.__name__}\",\n            loss,\n            on_step=True,\n            on_epoch=True,\n            prog_bar=True,\n        )\n\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        loss, preds = self.common_step(batch, batch_idx)\n        self.log(\n            f\"val/{self.loss.__class__.__name__}\",\n            loss,\n            on_step=True,\n            on_epoch=True,\n            prog_bar=True,\n        )\n\n        return loss\n\n    def test_step(self, batch, batch_idx):\n        loss, preds = self.common_step(batch, batch_idx)\n        self.log(\n            f\"test/{self.loss.__class__.__name__}\",\n            loss,\n            on_step=False,\n            on_epoch=True,\n            prog_bar=True,\n        )\n        self.main_metric(preds, batch[\"labels\"].int())\n        self.log(\n            f\"test/{self.main_metric.__class__.__name__}\",\n            self.main_metric,\n            on_step=False,\n            on_epoch=True,\n            prog_bar=False,\n        )\n        self.other_metrics(preds, batch[\"labels\"].int())\n        self.log_dict(self.other_metrics, on_step=False, on_epoch=True, prog_bar=False)\n\n        return loss\n\n    def configure_optimizers(self):\n        return AdamW(self.parameters(), lr=5e-5)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loading Data\n\nThere are two ways to use BirdSet datasets:\n\n1. using BirdSet's inbuild datamodules which manage the loading, preparing and transformation of data\n2. loading, preparing and transforming the data manually\n\nBoth of these options are shown in this section. The configurations shown here are taken from the configurations that are used in BirdSet's training proccesses, although some things (e.g. batch size) may differ.","metadata":{}},{"cell_type":"markdown","source":"### Using BirdSet Datamodules\n\nThis section shows how to load a dataset with BirdSets datamodules. These are the modules that handle the loading, preparing and transforming of data. They have a great amount of customizability. For more information about what the parameters of the various configs do, a look into the `birdset-pipeline_tutorial` notebook is advised. The configuration of datamodules is explained there.\n\nPlease note that `PretrainDataModule` should be used for the `XCL` and `XCM` datasets as they do not have `test` splits.","metadata":{}},{"cell_type":"markdown","source":"#### Configuring Transforms\n\nHere all the relevant transforms as well as event decoding and extracting is confifigured. For a more granular look into which configurations apply to what please refer to the `Loading Data Manually` section or the `birdset-pipeline-tutorial`.\n\nPlease note that the usage of `NoCallMixer` and `AddBackgroundNoise` requires extra soundfiles that can be used as such. You can utilize a script provided by BirdSet under `resources/utils/download_background_noise.py` to download files that can be used. Alternatively you can set `nocall` to `None` and comment out `AddBackgroundNoise` to not use them for now.","metadata":{}},{"cell_type":"code","source":"from birdset.datamodule.components.transforms import (\n    BirdSetTransformsWrapper,\n    PreprocessingConfig,\n)\nfrom birdset.datamodule.components.event_decoding import EventDecoding\nfrom birdset.datamodule.components.feature_extraction import DefaultFeatureExtractor\nfrom birdset.datamodule.components.augmentations import (\n    NoCallMixer,\n    MultilabelMix,\n    AddBackgroundNoise,\n    PowerToDB,\n)\nfrom birdset.datamodule.components.resize import Resizer\nfrom torch_audiomentations import AddColoredNoise, Gain\nfrom torchaudio.transforms import Spectrogram, MelScale, FrequencyMasking, TimeMasking\nfrom torchvision.transforms import RandomApply\n\ndecoder = EventDecoding(\n    min_len=1, max_len=5, sampling_rate=32000, extension_time=8, extracted_interval=5\n)\n\nfeature_extractor = DefaultFeatureExtractor(\n    feature_size=1, sampling_rate=32000, padding_value=0.0, return_attention_mask=False\n)\n\nnocall = NoCallMixer(\n    directory=\"/mnt/stud/work/rantjuschin/datasets/background_noise/\",\n    p=0.075,\n    sampling_rate=32000,\n    length=5,\n)\n\nwav_transforms = {\n    \"multilabel_mix\": MultilabelMix(\n        p=0.7, min_snr_in_db=3.0, max_snr_in_db=30.0, mix_target=\"union\"\n    ),\n    \"add_background_noise\": AddBackgroundNoise(\n        p=0.5,\n        min_snr_in_db=3,\n        max_snr_in_db=30,\n        sample_rate=32000,\n        target_rate=32000,\n        background_paths=\"/mnt/stud/work/rantjuschin/datasets/background_noise/\",\n    ),\n    \"add_colored_noise\": AddColoredNoise(\n        p=0.2, max_f_decay=2, min_f_decay=-2, max_snr_in_db=30, min_snr_in_db=3\n    ),\n    \"gain\": Gain(p=0.2, min_gain_in_db=-18, max_gain_in_db=6),\n}\n\npreprocessing = PreprocessingConfig(\n    spectrogram_conversion=Spectrogram(n_fft=1024, hop_length=320, power=2.0),\n    resizer=Resizer(db_scale=True, target_height=None, target_width=None),\n    melscale_conversion=MelScale(n_mels=128, sample_rate=32000, n_stft=513),\n    dbscale_conversion=PowerToDB(),\n    normalize_spectrogram=True,\n    mean=-4.268,\n    std=4.569,\n)\n\nspec_transforms = {\n    \"frequency_masking\": RandomApply(\n        p=0.5, transforms=[FrequencyMasking(freq_mask_param=100, iid_masks=True)]\n    ),\n    \"time_masking\": RandomApply(\n        p=0.5, transforms=[TimeMasking(time_mask_param=100, iid_masks=True)]\n    ),\n}\n\n\nbirdset_transforms = BirdSetTransformsWrapper(\n    task=\"multilabel\",\n    sampling_rate=32000,\n    model_type=\"vision\",\n    max_length=5,\n    decoding=decoder,\n    feature_extractor=feature_extractor,\n    nocall_sampler=nocall,\n    waveform_augmentations=wav_transforms,\n    preprocessing=preprocessing,\n    spectrogram_augmentations=spec_transforms,\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Instantiating the Datamodule","metadata":{}},{"cell_type":"code","source":"from birdset.datamodule.birdset_datamodule import BirdSetDataModule\nfrom birdset.configs.datamodule_configs import DatasetConfig, LoadersConfig\nfrom birdset.datamodule.components import XCEventMapping\n\ndataset_config = DatasetConfig(\n    data_dir=\"/mnt/stud/work/rantjuschin/datasets/HSN\",\n    hf_path=\"DBD-research-group/BirdSet\",\n    hf_name=\"HSN\",\n    n_workers=3,\n    val_split=0.2,\n    task=\"multilabel\",\n    classlimit=500,\n    eventlimit=5,\n    sampling_rate=32000,\n    seed=2,\n)\nloaders_config = LoadersConfig()\nmapper_config = XCEventMapping()\n\ndatamodule = BirdSetDataModule(\n    dataset=dataset_config,\n    loaders=loaders_config,\n    transforms=birdset_transforms,\n    mapper=mapper_config,\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The code above only configures the datamodule. To actully use it we need to prepare the data (which also downloads it) and setup the dataloaders.","metadata":{}},{"cell_type":"code","source":"datamodule.prepare_data()\ndatamodule.setup(stage=\"fit\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now the datamodule can be used.","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndl = datamodule.train_dataloader()\nsample = next(iter(dl))[\"input_values\"][0]\n\nplt.imshow(sample.squeeze().numpy())","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loading Data Manually\n\nIf you want to implement your own data pipeline you can also load the datasets through Hugging Face datasets. This section will detail that aproach.  \nSome amount of BirdSet Code will still be used here, so as to not fill up this notebook with helper functions. If you want to, you can look up the used classes and methods under their respective import paths.","metadata":{}},{"cell_type":"code","source":"from datasets import load_dataset\n\ndataset = load_dataset(\n    path=\"DBD-research-group/BirdSet\",\n    name=\"HSN\",\n    cache_dir=\"/mnt/stud/work/rantjuschin/datasets/HSN\",\n    num_proc=4,\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(dataset)\nprint(dataset[\"train\"][0][\"audio\"])\nprint(dataset[\"train\"][0][\"detected_events\"])","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Preprocessing the Data\n\nYou may heave seen that the `audio` column doesn't actually contain any audio data. The loaded dataset as it is, only contains the filepaths of the respective audiosamples. As shown above those files may also contain multiple birdcall events per file which is why it is not encouraged to directly load the whole audio into the sample (it may also just be too big). For BirdSet we typically specify some kind of eventlimit per file and per class and extract and map those events to singular samples in the dataset.\n\n**Please Note:** Only the `train` split needs to be processed this way. Both the `test_5s` and `test` splits do not need to be processed like that.","metadata":{}},{"cell_type":"code","source":"from datasets import Audio\nfrom birdset.datamodule.components.event_mapping import XCEventMapping\n\ndataset[\"train\"] = dataset[\"train\"].cast_column(\n    column=\"audio\",\n    feature=Audio(\n        sampling_rate=32_000,\n        mono=True,\n        decode=False,\n    ),\n)\n\nmapper = XCEventMapping()\ndataset[\"train\"] = dataset[\"train\"].map(\n    mapper,\n    remove_columns=[\"audio\"],\n    batched=True,\n    batch_size=300,\n    num_proc=3,\n    desc=\"Train event mapping\",\n)\n\ndataset[\"train\"] = dataset[\"train\"].remove_columns(\"audio\")\n\nprint(dataset)\nprint(dataset[\"train\"][0][\"filepath\"])\nprint(dataset[\"train\"][0][\"detected_events\"])","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now every event in the train split has been extracted into it's own sample. If you want to, you can limit these to some amount of events per class or per file.","metadata":{}},{"cell_type":"code","source":"# smart sampling is defined in the \"Helper Functions\" section\ndataset[\"train\"] = smart_sampling(\n    dataset=dataset[\"train\"], label_name=\"ebird_code\", class_limit=500, event_limit=5\n)\n\nprint(dataset)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The last step of this part is to remove unnecessary columns or splits. For BirdSet we use the `test_5s` split for multilabel testing. It contains 5 second long soundscape samples that may contain multiple birds per sample.  \nTo use it you need to remove the `test` split and rename the `test_5s` split to `test`.\n\nAdditionaly the dataset contains multiple columns that are not used and as such can be removed to make it more lightweight.","metadata":{}},{"cell_type":"code","source":"from datasets import DatasetDict\n\ndataset = DatasetDict({\"train\": dataset[\"train\"], \"test\": dataset[\"test_5s\"]})\ndataset","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = dataset.rename_column(\"ebird_code_multilabel\", \"labels\")\ndataset","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"columns_to_keep = {\"filepath\", \"labels\", \"detected_events\", \"start_time\", \"end_time\"}\n\nremovable_train_columns = [\n    column for column in dataset[\"train\"].column_names if column not in columns_to_keep\n]\nremovable_test_columns = [\n    column for column in dataset[\"test\"].column_names if column not in columns_to_keep\n]\n\nprint(removable_test_columns, \"\\n\", removable_train_columns)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset[\"train\"] = dataset[\"train\"].remove_columns(removable_train_columns)\ndataset[\"test\"] = dataset[\"test\"].remove_columns(removable_test_columns)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(dataset)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now the labels need to be one hot encoded for multilabel training.","metadata":{}},{"cell_type":"code","source":"# classes_one_hot it defined in the \"Helper Functions\" section\ndataset = dataset.map(\n    lambda batch: classes_one_hot(batch, num_classes=21),\n    batched=True,\n    batch_size=300,\n    load_from_cache_file=True,\n    num_proc=4,\n    desc=f\"One-hot-encoding labels.\",\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Splitting the Data\n\nfor better results the `train` split should be split into splits for training and for validation.","metadata":{}},{"cell_type":"code","source":"from datasets import DatasetDict\n\nsplits = dataset[\"train\"].train_test_split(test_size=0.2)\n\ndataset = DatasetDict(\n    {\"train\": splits[\"train\"], \"valid\": splits[\"test\"], \"test\": dataset[\"test\"]}\n)\n\nprint(dataset)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Applying Transforms\n\nAs the datasets can get very big very quickly we suggest applying transforms on-the-fly. That means that transforms are only applied to a batch when it is requested (e.g. by a trainer). With huggingface datasets that is possible through the `set_transform` method. For this you will need a wrapper class that will be able to handle multiple transforms on the data. Here, only a very basic one is implemented but like before you can look up a more sophisticated one in the BirdSet code (`BirdSetTransformsWrapper` in `birdset/datamodule/components/transforms.py`).","metadata":{}},{"cell_type":"markdown","source":"##### Decoding & Extracting\n\nAs previously mentioned the loaded dataset currently only contains filepaths but to apply transforms or use the data otherwise actual audio data is needed. As the name implies BirdSets `EventDecoding` class takes an event and pulls the associated audio data out of the respective file.","metadata":{}},{"cell_type":"code","source":"from birdset.datamodule.components.event_decoding import EventDecoding\nfrom birdset.datamodule.components.feature_extraction import DefaultFeatureExtractor\n\ndecoder = EventDecoding(\n    min_len=1, max_len=5, sampling_rate=32000, extension_time=8, extracted_interval=5\n)\n\nfeature_extractor = DefaultFeatureExtractor(\n    feature_size=1, sampling_rate=32000, padding_value=0.0, return_attention_mask=False\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### Nocall Sampling\n\nNocall Sampling uses other soundfiles to introduce samples which don't have any birdcalls in them as such samples are not included by default. To use Nocall Sampling you need to download background noise. BirdSet provides such functionality under `resources/utils/download_background_noise.py`. You will also need this background noise to use `AddBackgroundNoise` in the Waveform Transforms section.","metadata":{}},{"cell_type":"code","source":"from birdset.datamodule.components.augmentations import NoCallMixer\n\nnocall = NoCallMixer(\n    directory=\"/mnt/stud/work/rantjuschin/datasets/background_noise/\",\n    p=0.075,\n    sampling_rate=32000,\n    length=5,\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### Waveform Transforms\n\nThe transform in this section change the input data to ensure that the models learned are more robust.","metadata":{}},{"cell_type":"code","source":"wav_transforms = []","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from birdset.datamodule.components.augmentations import MultilabelMix\n\nmultilabel_mix = MultilabelMix(\n    p=0.7, min_snr_in_db=3.0, max_snr_in_db=30.0, mix_target=\"union\"\n)\n\nwav_transforms.append(multilabel_mix)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from birdset.datamodule.components.augmentations import AddBackgroundNoise\n\nbackground_noise = AddBackgroundNoise(\n    p=0.5,\n    min_snr_in_db=3,\n    max_snr_in_db=30,\n    sample_rate=32000,\n    target_rate=32000,\n    background_paths=\"/mnt/stud/work/rantjuschin/datasets/background_noise/\",\n)\n\nwav_transforms.append(background_noise)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch_audiomentations import AddColoredNoise\n\ncolored_noise = AddColoredNoise(\n    p=0.2, max_f_decay=2, min_f_decay=-2, max_snr_in_db=30, min_snr_in_db=3\n)\n\nwav_transforms.append(colored_noise)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch_audiomentations import Gain\n\ngain = Gain(p=0.2, min_gain_in_db=-18, max_gain_in_db=6)\n\nwav_transforms.append(gain)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### Spectrogramm Processing\n\nAs ConvNeXT needs spectrogramm data as input, the waveform need to turned into spectrogramms. These steps are configured here.","metadata":{}},{"cell_type":"code","source":"from torchaudio.transforms import Spectrogram\n\nspectrogramm_conversion = Spectrogram(n_fft=1024, hop_length=320, power=2.0)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchaudio.transforms import MelScale\n\nmelscale_conversion = MelScale(n_mels=128, sample_rate=32000, n_stft=513)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from birdset.datamodule.components.augmentations import PowerToDB\n\ndbscale_conversion = PowerToDB()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from birdset.datamodule.components.resize import Resizer\n\nresizer = Resizer(db_scale=True, target_height=None, target_width=None)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"processing_config = CustomProcessingConfig(\n    spectrogram_conversion=spectrogramm_conversion,\n    resizer=resizer,\n    melscale_conversion=melscale_conversion,\n    dbscale_conversion=dbscale_conversion,\n    normalize_spectrogram=True,\n    mean=-4.268,\n    std=4.569,\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### Spectrogram Transformation\n\nLike with the waveform transforms before the spectrogramms are also transformed to ensure more robust models.","metadata":{}},{"cell_type":"code","source":"spec_transforms = []","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.transforms import RandomApply\nfrom torchaudio.transforms import FrequencyMasking\n\nfrequency_masking = RandomApply(\n    p=0.5, transforms=[FrequencyMasking(freq_mask_param=100, iid_masks=True)]\n)\n\nspec_transforms.append(frequency_masking)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.transforms import RandomApply\nfrom torchaudio.transforms import TimeMasking\n\ntime_masking = RandomApply(\n    p=0.5, transforms=[TimeMasking(time_mask_param=100, iid_masks=True)]\n)\n\nspec_transforms.append(time_masking)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### On The Fly Transforming\n\nAs mentioned we want to apply the transforms on a batch only as it is getting processed. To achieve that behaviour we use the `set_transform` method of huggingface datatets. The test data does not get transformed using waveform and spectrogramm transformations. It only gets processed so that it can be used by the model.","metadata":{}},{"cell_type":"code","source":"transforms = BasicTransformsWrapper(\n    wav_transforms=wav_transforms,\n    spec_transforms=spec_transforms,\n    decoding=decoder,\n    feature_extractor=feature_extractor,\n    processing_config=processing_config,\n    nocall_sampler=nocall,\n)\n\ndataset[\"train\"].set_transform(transforms, output_all_columns=False)\ndataset[\"valid\"].set_transform(transforms, output_all_columns=False)\n\ntest_transforms = BasicTransformsWrapper(\n    wav_transforms=[],\n    spec_transforms=[],\n    decoding=decoder,\n    feature_extractor=feature_extractor,\n    processing_config=processing_config,\n    nocall_sampler=None,\n)\n\ndataset[\"test\"].set_transform(test_transforms, output_all_columns=False)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Wrapping the Data in a Datamodule\n\nTo use hte previously configured dataset in a trainer it needs to be wrapped in a datamodule. The datamodule provides dataloaders that are needed by the trainer.","metadata":{}},{"cell_type":"code","source":"datamodule = CustomDatamodule(\n    dataset=dataset, batch_size=32, num_workers=8, num_classes=21, task=\"multilabel\"\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"datamodule.dataset","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndl = datamodule.train_dataloader()\nsample = next(iter(dl))[\"input_values\"][0]\n\nplt.imshow(sample.squeeze().numpy())","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_loader = datamodule.test_dataloader()\nsample = next(iter(test_loader))\n\nplt.imshow(sample[\"input_values\"][2].squeeze().numpy())","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Setting up a Model","metadata":{}},{"cell_type":"markdown","source":"### Using BirdSet's Pretrained Models with BirdSet\n\nHere the ConvNeXT model pretrained on XCL by BirdSet is loaded. For more information regarding BirdSet's modules please refer to the `birdset-pipeline_tutorial` notebook.","metadata":{}},{"cell_type":"code","source":"from birdset.modules.models.convnext import ConvNextClassifier\nfrom birdset.modules.multilabel_module import MultilabelModule\nfrom birdset.configs import (\n    NetworkConfig,\n    LRSchedulerConfig,\n    MultilabelMetricsConfig,\n    LoggingParamsConfig,\n)\n\nnetwork = NetworkConfig(\n    model=ConvNextClassifier(\n        checkpoint=\"DBD-research-group/ConvNeXT-Base-BirdSet-XCL\",\n        num_classes=datamodule.num_classes,\n        num_channels=1,\n    ),\n    model_name=\"convnext\",\n    model_type=\"vision\",\n)\n\nmodel = MultilabelModule(\n    network=network,\n    num_epochs=5,\n    len_trainset=datamodule.len_trainset,\n    task=datamodule.task,\n    batch_size=datamodule.train_batch_size,\n)\n\nmodel","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loading Models through Hugging Face\n\nLoading a model through Hugging Face and then training it using a Lightning trainer is a bit more work as the trainer expects a `LightningModule`. This means that we need to wrap the model in such a module.","metadata":{}},{"cell_type":"code","source":"# ConvNextClassifierLightningModule is a custom class defined in the \"Helper\" section\n# HSN contains 21 classes\nmodel = ConvNextClassifierLightningModule(21, num_epochs=5)\nmodel","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Starting Fine-Tuning with a Lightning Trainer\n\nNow that the model and dataset are loaded the trainer needs to be configured. As before the trainer here is configured similar to how it would be in a BirdSet training although it is simplified for the use in this notebook. The `ModelCheckpoint` callback allows the easy selection of models after the evaluation.\n\nPlease note that the trainer is configured to only two epochs of training and validation. If you get errors regarding the usage of hardware resources you can try to lower the batch size or the numver of workers in the datamodules.","metadata":{}},{"cell_type":"code","source":"import lightning as L\nfrom lightning.pytorch.callbacks import ModelCheckpoint\nfrom lightning.pytorch.callbacks import RichModelSummary\n\nmodel_checkpoint = ModelCheckpoint(\n    dirpath=\"/mnt/stud/work/rantjuschin/callback_checkpoints\",\n    monitor=\"val/BCEWithLogitsLoss\",\n    verbose=False,\n    save_last=False,\n    save_top_k=2,\n    mode=\"min\",\n    auto_insert_metric_name=False,\n    save_weights_only=False,\n    every_n_train_steps=None,\n    train_time_interval=None,\n    every_n_epochs=1,\n    save_on_train_epoch_end=None,\n)\n\nrich_model_summary = RichModelSummary(max_depth=1)\n\ntrainer = L.Trainer(\n    min_epochs=1,\n    max_epochs=2,\n    gradient_clip_val=0.5,\n    precision=16,\n    accumulate_grad_batches=1,\n    callbacks=[model_checkpoint, rich_model_summary],\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now that the trainer is configured the model can be fine-tuned.","metadata":{}},{"cell_type":"code","source":"trainer.fit(datamodule=datamodule, model=model)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.callback_metrics","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ckpt_path = trainer.checkpoint_callback.best_model_path\nckpt_path","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.test(datamodule=datamodule, model=model, ckpt_path=ckpt_path)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}