{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":25954,"databundleVersionId":2091745,"sourceType":"competition"},{"sourceId":33246,"databundleVersionId":3221581,"sourceType":"competition"},{"sourceId":44224,"databundleVersionId":5188730,"sourceType":"competition"},{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":5258767,"sourceType":"datasetVersion","datasetId":3060198},{"sourceId":5258926,"sourceType":"datasetVersion","datasetId":3060292},{"sourceId":5259354,"sourceType":"datasetVersion","datasetId":3060577},{"sourceId":5259641,"sourceType":"datasetVersion","datasetId":3060750},{"sourceId":5866811,"sourceType":"datasetVersion","datasetId":3373324},{"sourceId":8631469,"sourceType":"datasetVersion","datasetId":5167934},{"sourceId":8659985,"sourceType":"datasetVersion","datasetId":5037230},{"sourceId":183205886,"sourceType":"kernelVersion"}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"!pip install audiomentations\n!pip install torch-audiomentations\n!pip install torchlibrosa\n!pip install colorednoise\n!pip install -U librosa\n!pip install omegaconf\n!wget https://raw.githubusercontent.com/LIHANG-HONG/birdclef2023-2nd-place-solution/main/modules/augmentations.py\n\n!cp /kaggle/input/birdclef-2024-training/* ./","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:21:07.615715Z","iopub.execute_input":"2024-06-08T04:21:07.616059Z","iopub.status.idle":"2024-06-08T04:22:30.383386Z","shell.execute_reply.started":"2024-06-08T04:21:07.616030Z","shell.execute_reply":"2024-06-08T04:22:30.382386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-06-13T04:31:59.073908Z","iopub.execute_input":"2024-06-13T04:31:59.074717Z","iopub.status.idle":"2024-06-13T04:32:01.447190Z","shell.execute_reply.started":"2024-06-13T04:31:59.074671Z","shell.execute_reply":"2024-06-13T04:32:01.446055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport os\nos.environ['CUBLAS_WORKSPACE_CONFIG'] = ':16:8'\nos.environ['PYTHONHASHSEED'] = '123'\n# os.environ['CUDA_VISIBLE_DEVICES'] = '1'\n\nimport IPython.display\nimport glob\nimport json\nimport warnings\nfrom contextlib import nullcontext\nimport random\nimport math\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport omegaconf\nfrom omegaconf import OmegaConf\nimport multiprocessing\nfrom collections import deque\nimport io\nimport itertools\nfrom copy import deepcopy\n\nimport torch\nfrom torch import nn, optim\nfrom torch.cuda import amp\nimport torch.nn.functional as F\nimport torchaudio\nimport torchvision\nfrom torchmetrics import MeanMetric, Accuracy\nimport timm\nimport librosa\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import average_precision_score\nfrom sklearn.manifold import TSNE\n# from transformers import AutoProcessor, ASTModel\n\nfrom augmentations import (\n    CustomCompose,\n    CustomOneOf,\n    NoiseInjection,\n    GaussianNoise,\n    PinkNoise,\n    AddGaussianNoise,\n    AddGaussianSNR,\n)\nfrom audiomentations import Compose as amCompose\nfrom audiomentations import OneOf as amOneOf\nfrom audiomentations import AddBackgroundNoise, Gain, GainTransition, TimeStretch\nfrom torch_audiomentations import Compose#, PitchShift, Shift\n\n# from blocks import AttHead\n# from SDAT import ConditionalDomainAdversarialLoss, MinimumClassConfusionLoss, SAM, DomainDiscriminator\n\n\n# from notebook_cache_outputs import output_cache\n\ntorchaudio.backend.set_audio_backend('soundfile')\ntorch.set_flush_denormal(True)\n\nSAMPLING_RATE = 32000\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:32.560225Z","iopub.execute_input":"2024-06-08T04:22:32.560962Z","iopub.status.idle":"2024-06-08T04:22:42.370828Z","shell.execute_reply.started":"2024-06-08T04:22:32.560905Z","shell.execute_reply":"2024-06-08T04:22:42.369868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# taxonomy_df = pd.read_csv(os.path.join('/kaggle/input/birdclef-2024', 'eBird_Taxonomy_v2021.csv'))\n# submission_df = pd.read_csv(os.path.join('/kaggle/input/birdclef-2024', 'sample_submission.csv'))\n# train_metadata_df = pd.read_csv(os.path.join('/kaggle/input/birdclef-2024', 'train_metadata.csv'))\n\ndef plot_waveform(waveform, sampling_rate):\n    waveform = torchaudio.transforms.Resample(sampling_rate, 1600)(waveform)\n    plt.plot(torch.arange(0, waveform.shape[-1]) / 1600, waveform[0], linewidth=1)","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.373437Z","iopub.execute_input":"2024-06-08T04:22:42.374202Z","iopub.status.idle":"2024-06-08T04:22:42.379665Z","shell.execute_reply.started":"2024-06-08T04:22:42.374163Z","shell.execute_reply":"2024-06-08T04:22:42.378512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# audio, sr = torchaudio.load('/kaggle/input/birdclef-2021/train_short_audio/acafly/XC11209.ogg')\n# IPython.display.display(IPython.display.Audio(audio, rate=sr))\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.380794Z","iopub.execute_input":"2024-06-08T04:22:42.381161Z","iopub.status.idle":"2024-06-08T04:22:42.393123Z","shell.execute_reply.started":"2024-06-08T04:22:42.381129Z","shell.execute_reply":"2024-06-08T04:22:42.392249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"markdown","source":"## Classes","metadata":{}},{"cell_type":"code","source":"\ndef read_taxonomy_information(path):\n    taxonomy_df = pd.read_csv(path)\n    return create_taxonomy_maps(taxonomy_df)\n\ndef create_taxonomy_maps(taxonomy_df):\n    full_name_to_code_dict = dict()\n    code_to_order_dict = dict()\n    code_to_family_dict = dict()\n    for i in tqdm(range(len(taxonomy_df))):\n        full_name, code, order, family = taxonomy_df.iloc[i][['PRIMARY_COM_NAME', 'SPECIES_CODE', 'ORDER1', 'FAMILY']]\n        full_name_to_code_dict[full_name] = code\n        code_to_order_dict[code] = order\n        code_to_family_dict[code] = family\n    \n    return full_name_to_code_dict, code_to_order_dict, code_to_family_dict\n            \nclass BirdCLEFDatasetV2(torch.utils.data.Dataset):\n    def __init__(self, root, meta_df, taxonomy_maps, use_secondary=True, append_class_on_fn=False,\n                 primary_col_name='primary_label', decode_audio=True, roots=None, use_multilabel=False):\n        super().__init__()\n        full_name_to_code_dict, code_to_order_dict, code_to_family_dict = taxonomy_maps\n        self.root = root\n        self.roots = roots\n        self.meta_df = meta_df\n        self.use_secondary = use_secondary\n        self.append_class_on_fn = append_class_on_fn\n        self.primary_col_name = primary_col_name\n        self.decode_audio = decode_audio\n        self.use_multilabel = use_multilabel\n        self.order_dict, self.family_dict = code_to_order_dict, code_to_family_dict\n        self.full_name_to_code_dict = full_name_to_code_dict\n        self.classes = np.unique(self.meta_df[self.primary_col_name])\n        self.families = np.unique(list(self.family_dict.values()))\n        self.orders = np.unique(list(self.order_dict.values()))\n        self._class_dict = {c:i for i, c in enumerate(self.classes)}\n        self._family_dict = {c:i for i, c in enumerate(self.families)}\n        self._order_dict = {c:i for i, c in enumerate(self.orders)}\n        \n        \n    def __getitem__(self, index):\n        path, primary_label, secondary_labels, rating = \\\n                                            self.meta_df.iloc[index][['filename', \n                                                                      self.primary_col_name,\n                                                                      'secondary_labels',\n                                                                      'rating',]]\n        if self.append_class_on_fn:\n            path = os.path.join(primary_label, path)\n\n        labels = [primary_label]\n        if self.use_secondary:\n            # secondary_labels to list\n            if secondary_labels == '[]':\n                secondary_labels = []\n            else:\n                secondary_labels = secondary_labels.strip('[]').split(',')\n                secondary_labels = [label.strip(' \\'\"') for label in secondary_labels]\n            labels = labels + secondary_labels\n            \n        # weight of each label\n        if self.use_multilabel:\n            # label_weights = torch.ones(len(labels)) * 0.5 * (1 + (rating + 1) / 6)\n            label_weights = torch.FloatTensor(len(labels) * [0.5 * (1 + (rating + 3) / 8)])\n            label_smoothing = (5 - rating) / 5 * 0.1\n        else:\n            label_weights = torch.ones(len(labels)) * ((rating + 1) / 6 / len(labels))\n            label_smoothing = (1 - label_weights.sum()) / (len(self.classes) - len(labels))\n        \n        # labels to one-not encoded\n        onehot_c = label_smoothing * torch.ones(len(self.classes))\n        for w, lab in zip(label_weights, labels):\n            if lab not in self._class_dict:\n                continue\n            onehot_c[self._class_dict[lab]] = w\n\n        if self.roots is not None:\n            root = np.random.choice([r for r, _ in self.roots], 1, p=[p for _, p in self.roots])[0]\n        else:\n            root = self.root\n        path = os.path.join(root, path)\n        \n        if self.decode_audio:\n            audio, sr = torchaudio.load(path)\n        else:\n            with open(path, 'rb') as fp:\n                audio = fp.read()\n        return (audio, onehot_c)\n    \n    def __len__(self):\n        return len(self.meta_df)\n    \n\nclass BirdCLEFAdditionalDataset(BirdCLEFDatasetV2):\n    def __init__(self, roots, *args, **kwargs):\n        kwargs['primary_col_name'] = 'ebird_code'\n        super().__init__(None, *args, **kwargs)\n        self.roots = roots\n        self._lab2root_dict = dict([(dn, root) for root in self.roots for dn in os.listdir(root)])\n        \n    def __getitem__(self, index):\n        path, primary_label, secondary_labels, rating = \\\n                                            self.meta_df.iloc[index][['filename', \n                                                                      self.primary_col_name,\n                                                                      'secondary_labels',\n                                                                      'rating',]]\n        path = os.path.splitext(path)[0] + '.ogg'\n        if self.append_class_on_fn:\n            path = os.path.join(primary_label, path)\n            \n        labels = [primary_label]\n        if self.use_secondary:\n            # secondary_labels to list\n            if secondary_labels == '[]':\n                secondary_labels = []\n            else:\n                secondary_labels = secondary_labels.strip('[]').split(',')\n                secondary_labels = [label.strip(' \\'\"') for label in secondary_labels]\n            labels = labels + secondary_labels\n            \n        # weight of each label\n        if self.use_multilabel:\n            # label_weights = torch.ones(len(labels)) * 0.5 * (1 + (rating + 1) / 6)\n            label_weights = torch.FloatTensor(len(labels) * [0.5 * (1 + (rating + 3) / 8)])\n            label_smoothing = (5 - rating) / 5 * 0.1\n        else:\n            label_weights = torch.ones(len(labels)) * ((rating + 1) / 6 / len(labels))\n            label_smoothing = (1 - label_weights.sum()) / (len(self.classes) - len(labels))\n        \n        # labels to one-not encoded\n        onehot_c = label_smoothing * torch.ones(len(self.classes))\n        for w, lab in zip(label_weights, labels):\n            if lab not in self._class_dict:\n                continue\n            onehot_c[self._class_dict[lab]] = w\n        \n        path = os.path.join(self._lab2root_dict[primary_label], path)\n        if self.decode_audio:\n            audio, sr = torchaudio.load(path)\n        else:\n            with open(path, 'rb') as fp:\n                audio = fp.read()\n        return (audio, onehot_c)\n    \n    def __len__(self):\n        return len(self.meta_df)\n\nclass AudioDatasetWrapper(torch.utils.data.Dataset):\n    def __init__(self, dataset, transform):\n        super().__init__()\n        self.dataset = dataset\n        self.transform = transform\n        \n    def __getitem__(self, index):\n        items = self.dataset[index]\n        is_list = isinstance(items, (tuple, list))\n        \n        x = items[0] if is_list else items\n        x = self.transform(x)\n\n        if is_list:\n            return (x, *items[1:])\n        else:\n            return x\n    \n    def __len__(self):\n        return len(self.dataset)\n\nclass TestBirdCLEFDataset(torch.utils.data.Dataset):\n    def __init__(self, dataset, train_info, val_info):\n        super().__init__()\n        self.dataset = dataset\n        self.reduced_to_base_classes = torch.zeros((len(val_info['base_classes']), len(train_info['reduced_classes'])))\n        for i, bc in enumerate(val_info['base_classes']):\n            j = train_info['reduced_classes'].index(bc)\n            self.reduced_to_base_classes[i, j] = 1\n    \n    def __getitem__(self, index):\n        x, y = self.dataset[index]\n        s0 = y.sum()\n        y = torch.tensordot(self.reduced_to_base_classes, y, dims=[[1], [0]])\n        y = y / y.sum() * s0\n        return x, y\n\n    def __len__(self):\n        return len(self.dataset)\n\nclass RandomAccessDataset(torch.utils.data.Dataset):\n    def __init__(self, paths, roots):\n        super().__init__()\n        self.paths = paths\n        self.roots = roots\n\n    def __getitem__(self, index):\n        path = self.paths[index]\n        root = np.random.choice([r for r, _ in self.roots], 1, p=[p for _, p in self.roots])[0]\n        path = os.path.join(root, path)\n        audio, sr = torchaudio.load(path)\n        return audio\n\n    def __len__(self):\n        return len(self.paths)\n    \nclass UDADatasetWrapper(torch.utils.data.Dataset):\n    def __init__(self, dataset, weak_transform, strong_transform):\n        super().__init__()\n        self.dataset = dataset\n        self.weak_transform = weak_transform\n        self.strong_transform = strong_transform\n    \n    def __getitem__(self, index):\n        x = self.dataset[index]\n        x_w = self.weak_transform(x)\n        x_s = self.strong_transform(x)\n        return x_w, x_s\n\n    def __len__(self):\n        return len(self.dataset)\n\nclass ZipDatasetWrapper(torch.utils.data.Dataset):\n    def __init__(self, datasets, permute=False, trunc=True, flatten=False):\n        self.datasets = datasets\n        self.permute = permute\n        self.trunc = trunc\n        self.flatten = flatten\n\n    def __getitem__(self, index):\n        if self.permute:\n            data = []\n            for ds in self.datasets:\n                data.append(ds[index % len(ds)])\n                index = index // len(ds)\n        else:\n            if self.trunc:\n                data = [ds[index] for ds in self.datasets]\n            else:\n                data = [ds[index % len(ds)] for ds in self.datasets]\n\n        if self.flatten:\n            data = [d for dd in data for d in dd]\n        return data\n\n    def __len__(self):\n        if self.permute:\n            return int(np.prod([len(ds) for ds in self.datasets]))\n        else:\n            if self.trunc:\n                return min([len(ds) for ds in self.datasets])\n            else:\n                return max([len(ds) for ds in self.datasets])","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.394432Z","iopub.execute_input":"2024-06-08T04:22:42.394693Z","iopub.status.idle":"2024-06-08T04:22:42.441472Z","shell.execute_reply.started":"2024-06-08T04:22:42.394671Z","shell.execute_reply":"2024-06-08T04:22:42.440765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Info","metadata":{}},{"cell_type":"code","source":"def get_dataset_infomation(dataset_list, use_secondary=True, duration_list=None, smooth_factor=0, use_multilabel=False):\n    ## compute statistic information of classes\n    base_classes = list(set([c for ds in dataset_list for c in ds.classes]))\n    base_classes.sort()\n    base_class_to_index_dict = {k:i for i, k in enumerate(base_classes)}\n    base_label_list = []\n    labels_list = []\n    weights_list = []\n    if duration_list is None:\n        duration_given = False\n        duration_list = []\n    else:\n        duration_given = True\n    labels_counts_dict = dict()\n    labels_duration_dict = dict()\n    pbar = tqdm(total=sum([len(ds) for ds in dataset_list]))\n    j = 0\n    pbar.set_description(\"Classes\")\n    for ds in dataset_list:\n        meta_df = ds.meta_df\n        for i in range(len(meta_df)):\n            path, primary_label, secondary_labels, rating = meta_df.iloc[i][['filename',\n                                                                             ds.primary_col_name,\n                                                                             'secondary_labels',\n                                                                             'rating',]]\n            # set base class\n            base_label_list.append(base_class_to_index_dict[primary_label])\n            # Labels\n            labels = [primary_label]\n            if use_secondary:\n                # secondary_labels to list\n                if secondary_labels == '[]':\n                    secondary_labels = []\n                else:\n                    secondary_labels = secondary_labels.strip('[]').split(',')\n                    secondary_labels = [label.strip(' \\'\"') for label in secondary_labels]\n                if isinstance(ds, BirdCLEFAdditionalDataset):\n                    secondary_labels = [lab.split('_')[1] for lab in secondary_labels]\n                    secondary_labels = [ds.full_name_to_code_dict.get(lab, lab) for lab in secondary_labels]\n                \n                labels = labels + secondary_labels\n                \n            # weight of each label\n            if use_multilabel:\n                label_weights = torch.ones(len(labels)) * 0.5 * (1 + (rating + 1) / 6)\n                label_weights = torch.FloatTensor(len(labels) * [0.5 * (1 + (rating + 1) / 6)])\n                # label_smoothing = (1 - label_weights[0]) / (len(self.classes) - len(labels))\n            else:\n                label_weights = torch.ones(len(labels)) * ((rating + 1) / 6 / len(labels))\n                # label_smoothing = (1 - label_weights.sum()) / (len(self.classes) - len(labels))\n            \n            # duration\n            if not duration_given:\n                path = os.path.splitext(path)[0] + '.ogg'\n                if ds.append_class_on_fn:\n                    path = os.path.join(primary_label, path)\n                if isinstance(ds, BirdCLEFAdditionalDataset):\n                    path = os.path.join(ds._lab2root_dict[primary_label], path)\n                else:\n                    path = os.path.join(ds.root, path)\n                duration = librosa.get_duration(path=path)\n                duration_list.append(duration)\n            else:\n                duration = duration_list[j]\n            \n            # updates\n            for w, lab in zip(label_weights, labels):\n                labels_counts_dict[lab] = labels_counts_dict.get(lab, 0) + w\n                labels_duration_dict[lab] = labels_duration_dict.get(lab, 0) + w * duration\n            \n            labels_list.append(labels)\n            weights_list.append(label_weights)\n            pbar.update(1)\n            j += 1\n    pbar.close()\n    \n    \n    # to list\n    all_label = sorted(list(labels_counts_dict.keys()))\n    all_counts = torch.FloatTensor([labels_counts_dict[lab] for lab in all_label])\n    all_duration = torch.FloatTensor([labels_duration_dict[lab] for lab in all_label])\n    \n    # remove the classes that are too small\n    indices = np.argsort(all_counts.numpy())\n    _all_label = np.array(all_label)[indices]\n    index = np.argwhere([c in set(base_classes) for c in _all_label])[0, 0]\n    selected_indices = np.sort(indices[index:])\n    reduced_label = np.array(all_label)[selected_indices].tolist()\n    reduced_counts = all_counts[selected_indices]\n    reduced_duration = all_duration[selected_indices]\n    \n    if smooth_factor == 'square':\n        resampling_weights = (reduced_duration.double() / reduced_duration.double().sum())  ** (-0.5)\n    else:\n        # resampling_weights = 1 / (reduced_duration.double() + smooth_factor * reduced_duration.double().max())\n        resampling_weights = 1 / (reduced_duration.double()**smooth_factor)\n    resampling_weights /= resampling_weights.sum()\n    resampling_weights = resampling_weights.float()\n    \n    # duration of audio\n    duration_list = torch.FloatTensor(duration_list)\n    duration_weights = ((len(duration_list) / duration_list.sum()) * duration_list).to(torch.float32)\n    \n    # compute resampling weight for each sample\n    counts_after_resampling = torch.zeros_like(reduced_counts)\n    duration_after_resampling = torch.zeros_like(reduced_duration)\n    class_to_index_dict = {k:i for i, k in enumerate(reduced_label)}\n    resampling_weight_dict = {lab:w for lab, w in zip(reduced_label, resampling_weights)}\n    \n    i = 0\n    samples_weights = []\n    pbar = tqdm(total=sum([len(ds) for ds in dataset_list]))\n    pbar.set_description(\"Resampling\")\n    for ds in dataset_list:\n        meta_df = ds.meta_df\n        for _ in range(len(meta_df)):\n            labels = labels_list[i]\n            weights = weights_list[i]\n            weights, labels = list(zip(*[(w, lab) for w, lab in zip(weights, labels) if lab in resampling_weight_dict])) # to reduced classes\n            sw = 0\n            for w, lab in zip(weights, labels):\n                p = w * resampling_weight_dict[lab] * duration_weights[i]\n                counts_after_resampling[class_to_index_dict[lab]] += w * resampling_weight_dict[lab]\n                duration_after_resampling[class_to_index_dict[lab]] += w * resampling_weight_dict[lab] * duration_list[i]\n                sw += p\n            samples_weights.append(sw.item())\n            i += 1\n            pbar.update(1)\n    pbar.close()\n    return {'classes': all_label,\n            'class_counts': all_counts,\n            'class_weights': resampling_weights,\n            'class_duration': all_duration,\n            'reduced_classes': reduced_label,\n            'reduced_class_counts': reduced_counts,\n            'reduced_class_duration': reduced_duration,\n            'class_counts_after_resampling': counts_after_resampling,\n            'class_duration_after_resampling': duration_after_resampling,\n            'class_to_index_dict': class_to_index_dict,\n            'samples_weights': samples_weights,\n            'durations': duration_list,\n            'base_classes': base_classes,\n            'base_class_to_index_dict': base_class_to_index_dict,\n            'base_labels': base_label_list}\n\n# info = load_dataset_information('info.json')\n# info2024 = load_dataset_information('info2024.json')\n# taxonomy_maps = read_taxonomy_information(os.path.join('/kaggle/input', 'birdclef-2024/eBird_Taxonomy_v2021.csv'))\n# ds20_list = load_additional_birdclef_datasets(taxonomy_maps, root='/kaggle/input'')\n# ds_list = [load_birdclef_dataset(year, taxonomy_maps, root='/kaggle/input'') for year in [2021, 2022, 2023, 2024]]\n\n\n# for names, (use_secondary, smooth_factor, use_multilabel) in zip(itertools.product(['', '_primary'], ['_p0', '_p1', '_p2'], ['_binary', '']), \n#                      itertools.product([True, False], [0, 1, 2], [True, False])):\n#     name = ''.join(names)\n#     print(name)\n#     info_new = get_dataset_infomation([*ds20_list, *ds_list], \n#                                       use_secondary=use_secondary, \n#                                       duration_list=info['durations'], \n#                                       smooth_factor=smooth_factor,\n#                                       use_multilabel=use_multilabel)\n#     info2024_new = get_dataset_infomation([ds_list[-1]], \n#                                           use_secondary=use_secondary, \n#                                           duration_list=info2024['durations'], \n#                                           smooth_factor=smooth_factor,\n#                                           use_multilabel=use_multilabel)\n#     save_dataset_information(info_new, f'info{name}.json')\n#     save_dataset_information(info2024_new, f'info2024{name}.json')\n\n\n# info = load_dataset_information('./info_s1.json')\n# plt.plot(info['class_weights'] * info['reduced_class_counts'])\n\n# sample_weights = (\n#     all_primary_labels.value_counts() / \n#     all_primary_labels.value_counts().sum()\n# )  ** (-0.5)\n# info['class_duration_after_resampling'].max(), info['class_duration_after_resampling'].min()","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.442505Z","iopub.execute_input":"2024-06-08T04:22:42.442951Z","iopub.status.idle":"2024-06-08T04:22:42.473076Z","shell.execute_reply.started":"2024-06-08T04:22:42.442882Z","shell.execute_reply":"2024-06-08T04:22:42.472166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"def get_kfold_subset(dataset, fold, y, seed=None, n_splits=5):\n    skf = StratifiedKFold(n_splits=n_splits, random_state=seed, shuffle=True)\n    indices = np.arange(len(dataset))\n    indices = list(skf.split(np.arange(len(dataset)), y))\n    train_indices = indices[fold][0]\n    val_indices = indices[fold][1]\n    train_ds = torch.utils.data.Subset(dataset, train_indices)\n    val_ds = torch.utils.data.Subset(dataset, val_indices)\n    return train_ds, val_ds, train_indices, val_indices\n\ndef load_birdclef_dataset(year, taxonomy_maps, root='/kaggle/input', **kwargs):\n    roots = kwargs.pop('roots') if 'roots' in kwargs else None\n    if year == 2021:\n        if roots is not None:\n            roots = [(os.path.join(r, f'birdclef-{year}', 'train_short_audio'), p) for r, p in roots]\n        taxonomy_df = pd.read_csv(os.path.join(root, f'birdclef-2024', 'eBird_Taxonomy_v2021.csv'))\n        train_metadata_df = pd.read_csv(os.path.join(root, f'birdclef-{year}', 'train_metadata.csv'))\n        ds = BirdCLEFDatasetV2(os.path.join(root, f'birdclef-{year}', 'train_short_audio'), \n                               train_metadata_df, taxonomy_maps, append_class_on_fn=True, \n                               roots=roots, **kwargs)\n    else:\n        if roots is not None:\n            roots = [(os.path.join(r, f'birdclef-{year}'), 'train_audio', p) for r, p in roots]\n        train_metadata_df = pd.read_csv(os.path.join(root, f'birdclef-{year}', 'train_metadata.csv'))\n        ds = BirdCLEFDatasetV2(os.path.join(root, f'birdclef-{year}', 'train_audio'), train_metadata_df, taxonomy_maps, **kwargs)\n    return ds\n\ndef load_additional_birdclef_datasets(taxonomy_maps, root='/kaggle/input', **kwargs):\n    kwargs.pop('roots') if 'roots' in kwargs else None\n    train_df_bs = pd.read_csv(os.path.join(root, 'cornell-birdsong-recognition-2020-a-m-32khz-ogg', 'train.csv'))\n    train_df_xc = pd.read_csv(os.path.join(root, 'xeno-canto-bird-recordings-extended-a-m-32khz-ogg', 'train_extended.csv'))\n    ds_bs = BirdCLEFAdditionalDataset([os.path.join(root, 'cornell-birdsong-recognition-2020-a-m-32khz-ogg', 'train_audio'),\n                                       os.path.join(root, 'cornell-birdsong-recognition-2020-n-z-32khz-ogg1', 'train_audio')], \n                                       train_df_bs, taxonomy_maps, append_class_on_fn=True, **kwargs)\n    ds_xc = BirdCLEFAdditionalDataset([os.path.join(root, 'xeno-canto-bird-recordings-extended-a-m-32khz-ogg', 'A-M'),\n                                       os.path.join(root, 'xeno-canto-bird-recordings-extended-n-z-32khz-ogg', 'N-Z')], \n                                       train_df_xc, taxonomy_maps, append_class_on_fn=True, **kwargs)\n\n    return ds_bs, ds_xc\n\n\ndef unified_classes_of_datasets(dataset_list, info=None):\n    # dataset in dataset_list are instances of BirdCLEFDatasetV2\n    # Note that families and orders are already unified.\n    if info is None:\n        info = get_dataset_infomation(dataset_list)\n        \n    for ds in dataset_list:\n        ds.classes = info['reduced_classes']\n        ds._class_dict = info['class_to_index_dict']\n    all_ds = torch.utils.data.ConcatDataset(dataset_list)\n    \n    return all_ds, info\n    \ndef save_dataset_information(info, path='info.json'):\n    info = info.copy()\n    info = {k:v.tolist() if torch.is_tensor(v) else v for k,v in info.items()}\n    with open(path, 'w') as fp:\n        json.dump(info, fp)\n        \ndef load_dataset_information(path='info.json'):\n    with open(path, 'r') as fp:\n        info = json.load(fp)\n    info = {k:torch.FloatTensor(v) if isinstance(v, (tuple, list)) and len(v) and isinstance(v[0], (int, float)) else v \n            for k,v in info.items()}\n    return info\n\ndef read_all_datasets(root='/kaggle/input', info_path='info.json', **kwargs):\n    taxonomy_maps = read_taxonomy_information(os.path.join(root, 'birdclef-2024/eBird_Taxonomy_v2021.csv'))\n    ds20_list = load_additional_birdclef_datasets(taxonomy_maps, root=root, **kwargs)\n    ds_list = [load_birdclef_dataset(year, taxonomy_maps, root=root, **kwargs) for year in [2021, 2022, 2023, 2024]]\n    info = None\n    if os.path.isfile(info_path):\n        info = load_dataset_information(info_path)\n    all_ds, info = unified_classes_of_datasets([*ds20_list, *ds_list], info)\n    return all_ds, info \n\ndef read_stage2_dataset(root='/kaggle/input', info_path='info-2024.json', **kwargs):\n    taxonomy_maps = read_taxonomy_information(os.path.join(root, 'birdclef-2024/eBird_Taxonomy_v2021.csv'))\n    ds_list = [load_birdclef_dataset(2024, taxonomy_maps, root=root, **kwargs)]\n    info = None\n    if os.path.isfile(info_path):\n        info = load_dataset_information(info_path)\n    all_ds, info = unified_classes_of_datasets(ds_list, info)\n    return all_ds, info \n\ndef get_dataloaders(cfg):\n    transforms_dict = get_augmentation_transforms(cfg.dataset.orig_sampling_rate, cfg.dataset.sampling_rate,\n                                                  audio_second=cfg.dataset.audio_second, \n                                                  background_root=os.path.join(cfg.dataset.root, 'background-noise'),\n                                                  use_cache=cfg.dataset.use_cache, strong=cfg.dataset.stage==1)\n    \n    \n    train_transform = torchvision.transforms.Compose([\n        (lambda x: torchaudio.load(io.BytesIO(x))[0]) if cfg.dataset.use_cache else (lambda x: x),\n        transforms_dict['align'],\n        Squeeze(0),\n        ToNumpy(),\n        transforms_dict['np'],\n        transforms_dict['am'],\n        ToTorch(),\n        # lambda x: x[None, None],\n        # transforms_dict['torch'],\n        # lambda x: x[0, 0],\n    ])\n\n    val_transform = torchvision.transforms.Compose([\n        (lambda x: torchaudio.load(io.BytesIO(x))[0]) if cfg.dataset.use_cache else (lambda x: x),\n        transforms_dict['align'],\n        Squeeze(0),\n    ])\n\n    if cfg.dataset.stage == 1:\n        ds, info = read_all_datasets(cfg.dataset.root, cfg.dataset.info_path, use_secondary=cfg.dataset.use_secondary,\n                                     decode_audio=not cfg.dataset.use_cache, roots=cfg.dataset.roots)\n    elif cfg.dataset.stage == 2:\n        ds, info = read_stage2_dataset(cfg.dataset.root, cfg.dataset.info_path, use_secondary=cfg.dataset.use_secondary,\n                                       decode_audio=not cfg.dataset.use_cache, roots=cfg.dataset.roots)\n\n    if cfg.dataset.use_cache:\n        ds = CacheDatasetWrapper(ds, cfg.dataset.max_cache)\n\n    if cfg.dataset.kfold is None:\n        train_ds, val_ds, train_indices, val_indices = get_kfold_subset(ds, 0, info['base_labels'], seed=cfg.seed)\n        train_ds = ds\n        train_indices = np.arange(len(ds))\n    else:\n        train_ds, val_ds, train_indices, val_indices = get_kfold_subset(ds, cfg.dataset.kfold, info['base_labels'], seed=cfg.seed)\n    train_samples_weights = info['samples_weights'][train_indices]\n    val_samples_weights = info['samples_weights'][val_indices]\n\n    train_ds = AudioDatasetWrapper(train_ds, train_transform)\n    val_ds = AudioDatasetWrapper(val_ds, val_transform)\n\n\n    samples_per_epoch = len(train_ds) if cfg.dataset.samples_per_epoch is None else cfg.dataset.samples_per_epoch\n\n\n    train_sampler = torch.utils.data.WeightedRandomSampler(train_samples_weights, replacement=True, num_samples=samples_per_epoch)\n    # train_sampler = torch.utils.data.RandomSampler(train_ds, replacement=True, num_samples=samples_per_epoch)\n\n    if cfg.dataset.resample_val:\n        val_sampler = torch.utils.data.WeightedRandomSampler(val_samples_weights, replacement=True, num_samples=len(val_ds))\n    else:\n        val_sampler = torch.utils.data.SequentialSampler(val_ds)\n\n    def seed_worker(worker_id):\n        worker_seed = torch.initial_seed() % 2**32\n        np.random.seed(worker_seed)\n        random.seed(worker_seed)\n    \n    train_generator = torch.Generator()\n    test_generator = torch.Generator()\n    if cfg.seed is not None:\n        train_generator.manual_seed(cfg.seed); test_generator.manual_seed(cfg.seed)\n\n    train_loader = torch.utils.data.DataLoader(\n            train_ds,\n            sampler=train_sampler,\n            batch_size=cfg.dataset.batch_size,\n            num_workers=cfg.dataset.num_workers,\n            drop_last=True,\n            pin_memory=True,\n            prefetch_factor=cfg.dataset.prefetch_factor,\n            worker_init_fn=seed_worker,\n            generator=train_generator\n    )\n    val_loader = torch.utils.data.DataLoader(\n            val_ds,\n            sampler=val_sampler,\n            batch_size=cfg.dataset.batch_size,\n            num_workers=cfg.dataset.num_workers,\n            drop_last=False,\n            pin_memory=True,\n            prefetch_factor=cfg.dataset.prefetch_factor,\n            worker_init_fn=seed_worker,\n            generator=test_generator,\n    )\n    return train_ds, val_ds, train_loader, val_loader, transforms_dict, info\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.474447Z","iopub.execute_input":"2024-06-08T04:22:42.474766Z","iopub.status.idle":"2024-06-08T04:22:42.513089Z","shell.execute_reply.started":"2024-06-08T04:22:42.474736Z","shell.execute_reply":"2024-06-08T04:22:42.512344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## UDA","metadata":{}},{"cell_type":"code","source":"def get_uda_dataloaders(cfg):\n    transforms_dict = get_augmentation_transforms(cfg.dataset.orig_sampling_rate, cfg.dataset.sampling_rate,\n                                                  audio_second=cfg.dataset.audio_second, \n                                                  background_root=os.path.join(cfg.dataset.root, 'background-noise'),\n                                                  use_cache=cfg.dataset.use_cache)\n    \n    strong_transform = torchvision.transforms.Compose([\n        (lambda x: torchaudio.load(io.BytesIO(x))[0]) if cfg.dataset.use_cache else (lambda x: x),\n        transforms_dict['align'],\n        Squeeze(0),\n        ToNumpy(),\n        transforms_dict['np'],\n        transforms_dict['am'],\n        ToTorch(),\n    ])\n\n    weak_transform = torchvision.transforms.Compose([\n        (lambda x: torchaudio.load(io.BytesIO(x))[0]) if cfg.dataset.use_cache else (lambda x: x),\n        transforms_dict['align'],\n        Squeeze(0),\n    ])\n\n    unlabeled_root = os.path.join(cfg.dataset.root, 'birdclef2024-unlabeled-patches')\n    if cfg.dataset.roots is None:\n        paths = [os.path.join(unlabeled_root, fn) for fn in os.listdir(unlabeled_root)]\n        ds = AudioDatasetWrapper(paths, lambda x: torchaudio.load(x)[0])\n    else:\n        paths = [os.path.join('birdclef2024-unlabeled-patches', fn) for fn in os.listdir(unlabeled_root)]\n        ds = RandomAccessDataset(paths, cfg.dataset.roots)\n\n    if cfg.dataset.use_cache:\n        ds = CacheDatasetWrapper(ds, cfg.dataset.max_cache)\n\n    ds = UDADatasetWrapper(ds, weak_transform, strong_transform)\n    \n    samples_per_epoch = len(ds) if cfg.dataset.samples_per_epoch is None else cfg.dataset.samples_per_epoch\n\n    sampler = torch.utils.data.RandomSampler(ds, replacement=True, num_samples=samples_per_epoch)\n\n    def seed_worker(worker_id):\n        worker_seed = torch.initial_seed() % 2**32\n        np.random.seed(worker_seed)\n        random.seed(worker_seed)\n    \n    generator = torch.Generator()\n    if cfg.seed is not None:\n        generator.manual_seed(cfg.seed)\n\n    loader = torch.utils.data.DataLoader(\n            ds,\n            sampler=sampler,\n            batch_size=cfg.dataset.batch_size,\n            num_workers=cfg.dataset.num_workers,\n            drop_last=True,\n            pin_memory=True,\n            prefetch_factor=cfg.dataset.prefetch_factor,\n            worker_init_fn=seed_worker,\n            generator=generator\n    )\n\n    return ds, loader, transforms_dict\n    ","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.514269Z","iopub.execute_input":"2024-06-08T04:22:42.514538Z","iopub.status.idle":"2024-06-08T04:22:42.527662Z","shell.execute_reply.started":"2024-06-08T04:22:42.514516Z","shell.execute_reply":"2024-06-08T04:22:42.526822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Augmentation","metadata":{}},{"cell_type":"code","source":"class RandomCrop:\n    def __init__(self, sampling_rate, crop_second):\n        self.sampling_rate = sampling_rate\n        self.crop_second = crop_second\n\n    def __call__(self, x):\n        win_size = self.sampling_rate * self.crop_second\n        if x.shape[-1] > win_size:\n            offset = random.randint(0, x.shape[-1] - win_size - 1)\n            x = x[..., offset:offset + win_size]\n        return x\n    \nclass Replay:\n    def __init__(self, sampling_rate, output_second):\n        self.sampling_rate = sampling_rate\n        self.output_second = output_second\n\n    def __call__(self, x):\n        output_size = int(self.sampling_rate * self.output_second)\n        if x.shape[-1] < output_size:\n            n_replay =  math.ceil(output_size / x.shape[-1])\n            x = x.repeat(1, n_replay)\n            x = x[..., :output_size]\n        return x\n\nclass Squeeze:\n    def __init__(self, dim):\n        self.dim = dim\n    def __call__(self, x):\n        return x.squeeze(self.dim)\n    \nclass Unsqueeze:\n    def __init__(self, dim):\n        self.dim = dim\n    def __call__(self, x):\n        return x.unsqueeze(self.dim)\n\nclass ToNumpy:\n    def __call__(self, x):\n        return x.numpy()\n\nclass ToTorch:\n    def __call__(self, x):\n        return torch.from_numpy(x)\n    \nclass NormalizeMelSpec(nn.Module):\n    def __init__(self, eps=1e-6, exportable=False):\n        super().__init__()\n        self.eps = eps\n        self.exportable = exportable\n\n    def forward(self, X):\n        mean = X.mean((1, 2), keepdim=True)\n        std = X.std((1, 2), keepdim=True)\n        Xstd = (X - mean) / (std + self.eps)\n        if self.exportable:\n            norm_max = torch.amax(Xstd, dim=(1, 2), keepdim=True)\n            norm_min = torch.amin(Xstd, dim=(1, 2), keepdim=True)\n            return (Xstd - norm_min) / (norm_max - norm_min + self.eps)\n        else:\n            norm_min, norm_max = (\n                Xstd.min(-1)[0].min(-1)[0],\n                Xstd.max(-1)[0].max(-1)[0],\n            )\n            fix_ind = (norm_max - norm_min) > self.eps * torch.ones_like(\n                (norm_max - norm_min)\n            )\n            V = torch.zeros_like(Xstd)\n            if fix_ind.sum():\n                V_fix = Xstd[fix_ind]\n                norm_max_fix = norm_max[fix_ind, None, None]\n                norm_min_fix = norm_min[fix_ind, None, None]\n                V_fix = torch.max(\n                    torch.min(V_fix, norm_max_fix),\n                    norm_min_fix,\n                )\n                V_fix = (V_fix - norm_min_fix) / (norm_max_fix - norm_min_fix)\n                V[fix_ind] = V_fix\n            return V\n\nclass LowerUpperFrequency(nn.Module):\n    def forward(self, images):\n        r = torch.randint(images.shape[-2] // 2, images.shape[-2], size=(1,))[0].item()\n        x = (torch.rand(size=(1,))[0] / 2).item()\n        pink_noise = (\n            torch.from_numpy(np.array([\n                np.concatenate([\n                    1 - np.arange(r) * x / r,\n                    np.zeros(128 - r) - x + 1,\n                ])\n            ])).t().float().to(images.device)\n        )\n        images = images * pink_noise\n        return images\n\nclass AddBackgroundNoiseV2(AddBackgroundNoise):\n    def __init__(self, *args, max_items=None, **kwargs):\n        super().__init__(*args, **kwargs)\n        self.max_items = max_items\n        self.manager = multiprocessing.Manager()\n        self.cached = self.manager.dict()\n        self.queue = self.manager.JoinableQueue()\n        self.num_query = self.manager.Value(int, 0)\n        self.num_hit = self.manager.Value(int, 0)\n        \n        def _load_sound(file_path, sample_rate):\n            # load raw\n            self.num_query.set(self.num_query.get() + 1)\n            if file_path in self.cached:\n                self.num_hit.set(self.num_hit.get() + 1)\n                b = self.cached[file_path]\n            else:\n                b = self._load_raw_audio(file_path)\n                self.cached[file_path] = b\n                if self.max_items is not None:\n                    self.queue.put_nowait(file_path)\n                    if self.queue.qsize() > self.max_items:\n                        self.cached.pop(self.queue.get_nowait())\n            waveform, sr = self._decode_raw_audio(b)\n            if sr != sample_rate:\n                waveform = torchaudio.functional.resample(waveform, sr, sample_rate)\n            waveform = waveform.mean(dim=0)\n            return waveform.numpy(), sample_rate\n            \n        # load_sound = self._load_sound\n        # def _load_sound(file_path, sample_rate):\n        #     x = load_sound(file_path, sample_rate)\n        #     print(x[0].shape)\n        #     return x\n        self._load_sound = _load_sound\n\n    def _decode_raw_audio(self, b):\n        return torchaudio.load(io.BytesIO(b))\n        \n    def _load_raw_audio(self, file_path):\n        with open(file_path, 'rb') as fp:\n            b = fp.read()\n        return b\n\n    def __del__(self):\n        self.manager.shutdown.cancel()\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.531513Z","iopub.execute_input":"2024-06-08T04:22:42.531830Z","iopub.status.idle":"2024-06-08T04:22:42.560169Z","shell.execute_reply.started":"2024-06-08T04:22:42.531807Z","shell.execute_reply":"2024-06-08T04:22:42.559352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_augmentation_transforms(orig_sampling_rate, sampling_rate, audio_second=5, background_root='/kaggle/input/background-noise',\n                                use_cache=False, max_cache=None, strong=True):\n    align_transforms = torchvision.transforms.Compose([\n        lambda x: torch.mean(x, dim=0, keepdim=True),\n        Replay(orig_sampling_rate, audio_second),\n        RandomCrop(orig_sampling_rate, audio_second),\n        torchaudio.transforms.Resample(orig_sampling_rate, sampling_rate),\n    ])\n    paths_list = []\n    if strong:\n        paths_list.append(glob.glob(os.path.join(background_root, 'birdclef2020_nocall/*')) +\\\n                          glob.glob(os.path.join(background_root, 'birdclef2021_nocall/*')))\n        paths_list.append(glob.glob(os.path.join(background_root, 'freefield/*')) +\\\n                          glob.glob(os.path.join(background_root, 'warblrb/*')) +\\\n                          glob.glob(os.path.join(background_root, 'birdvox/*')))\n        paths_list.append(glob.glob(os.path.join(background_root, 'rainforest/*')) +\\\n                          glob.glob(os.path.join(background_root, 'environment/*')))\n        paths_prob_list = [0.6, 0.3, 0.4]\n    else:\n        paths_list.append(glob.glob(os.path.join(background_root, 'freefield/*')) +\\\n                          glob.glob(os.path.join(background_root, 'warblrb/*')) +\\\n                          glob.glob(os.path.join(background_root, 'birdvox/*')))\n        paths_list.append(glob.glob(os.path.join(background_root, 'environment/*')))\n        paths_prob_list = [0.4, 0.6]\n   \n    noise_transform_func = (lambda *args, **kwargs: AddBackgroundNoiseV2(*args, max_items=max_cache, **kwargs)) if use_cache else AddBackgroundNoise\n    am_audio_transforms = amCompose(\n        [noise_transform_func(paths, min_snr_in_db=0, max_snr_in_db=3, p=p) \n         for paths, p in zip(paths_list, paths_prob_list)] + \\\n        [\n            amOneOf(\n                [\n                    Gain(min_gain_in_db=-15, max_gain_in_db=15, p=0.8),\n                    GainTransition(min_gain_in_db=-15, max_gain_in_db=15, p=0.8),\n                ],\n            ),\n        ]\n    )\n    def am_audio_transforms_func(x):\n        with warnings.catch_warnings():\n            warnings.simplefilter(\"ignore\")\n            x = am_audio_transforms(x, sample_rate=sampling_rate)\n        return x\n\n\n    np_audio_transforms = CustomCompose(\n        [\n            CustomOneOf(\n                [\n                    NoiseInjection(p=1, max_noise_level=0.04),\n                    GaussianNoise(p=1, min_snr=5, max_snr=20),\n                    PinkNoise(p=1, min_snr=5, max_snr=20),\n                    AddGaussianNoise(min_amplitude=0.0001, max_amplitude=0.03, p=0.5),\n                    AddGaussianSNR(min_snr_in_db=5, max_snr_in_db=15, p=0.5),\n                ],\n                p=0.3,\n            ),\n        ]\n    )\n\n#     torch_audio_transforms = Compose(\n#         [\n#             PitchShift(\n#                 min_transpose_semitones=-4,\n#                 max_transpose_semitones=4,\n#                 sample_rate=sampling_rate,\n#                 p=0.4,\n#             ),\n#             Shift(min_shift=-0.5, max_shift=0.5, p=0.4),\n#         ]\n#     )\n#     def torch_audio_transforms_func(x):\n#         x = torch_audio_transforms(x)\n#         return x\n\n    spec_augmentations = Compose([\n        torchaudio.transforms.TimeMasking(\n            time_mask_param=60, iid_masks=True, p=0.5\n        ),\n        torchvision.transforms.RandomApply(\n            nn.ModuleList([\n                torchaudio.transforms.FrequencyMasking(freq_mask_param=24, iid_masks=True),\n            ]), p=0.5\n        ),\n        torchvision.transforms.RandomApply(\n            nn.ModuleList([\n                LowerUpperFrequency(),\n            ]), p=0.5\n        ),\n    ])\n    \n    \n    \n    return {'align': align_transforms,\n            'am': am_audio_transforms_func, \n            'np': np_audio_transforms, \n#             'torch': torch_audio_transforms_func, \n            'spec': spec_augmentations}\n\ndef get_spec_processor(sampling_rate, audio_second, time_size=256):\n    processor_volodymyr = nn.Sequential(\n        # torchaudio.transforms.MelSpectrogram(sample_rate=sampling_rate, \n        #                                      n_mels=128, \n        #                                      f_min=20,\n        #                                      n_fft=2048,\n        #                                      hop_length=512,\n        #                                      normalized=True),\n        torchaudio.transforms.MelSpectrogram(sample_rate=sampling_rate, \n                                             hop_length=audio_second*sampling_rate // (time_size-1),\n                                             n_mels=128, f_min=0, f_max=sampling_rate//2, n_fft=2048, \n                                             center=True, pad_mode='constant',\n                                             norm='slaney',onesided=True,mel_scale='slaney'),\n        torchaudio.transforms.AmplitudeToDB(top_db=80),\n        NormalizeMelSpec(exportable=False),\n        # LambdaLayer(lambda x: F.interpolate(x[:, None], (224, 224), mode='bilinear')[:, 0]),\n    )\n    \n\n    # ast_processor = AutoProcessor.from_pretrained(\"MIT/ast-finetuned-audioset-10-10-0.4593\")\n#     ast_processor = AutoProcessor.from_pretrained(\"/home/lab/.cache/huggingface/hub/models--MIT--ast-finetuned-audioset-10-10-0.4593/snapshots/f826b80d28226b62986cc218e5cec390b1096902/\")\n\n#     ast_processor_func = lambda x: ast_processor(x, sampling_rate=sampling_rate, return_tensors=\"pt\")['input_values'][0]\n#     def ast_processor_batch_func(x):\n#         device = x.device\n#         x = x.cpu()\n#         return torch.stack([ast_processor_func(audio) for audio in x], dim=0).to(device)\n\n    return {\n#         'ast': ast_processor_batch_func,\n        'volodymyr': processor_volodymyr,\n    }\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.561424Z","iopub.execute_input":"2024-06-08T04:22:42.561881Z","iopub.status.idle":"2024-06-08T04:22:42.583084Z","shell.execute_reply.started":"2024-06-08T04:22:42.561841Z","shell.execute_reply":"2024-06-08T04:22:42.582251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"markdown","source":"## Decoders","metadata":{}},{"cell_type":"code","source":"\nclass AttHead(nn.Module):\n    def __init__(\n        self,\n        in_chans,\n        hidden_chans=512,\n        p=0.5,\n        num_class=397,\n        **kwargs,\n    ):\n        super().__init__()\n       \n        self.dense_layers = nn.Sequential(\n            nn.Dropout(p / 2),\n            nn.Linear(in_chans, hidden_chans),\n            nn.ReLU(),\n            nn.Dropout(p),\n        )\n        self.attention = nn.Conv1d(\n            in_channels=hidden_chans,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n        self.fix_scale = nn.Conv1d(\n            in_channels=hidden_chans,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n\n    def forward(self, feat):\n        feat = feat.squeeze(2)\n        feat = self.dense_layers(feat.permute(0, 2, 1)).permute(0, 2, 1)  # (bs, 512, time)\n\n        time_att = torch.tanh(self.attention(feat))\n        feat_v = self.fix_scale(feat)\n\n        clipwise_logits_long = torch.sum(\n            feat_v * torch.softmax(time_att, dim=-1),\n            dim=-1,\n        )\n\n        framewise_logits_long = feat_v.permute(0, 2, 1)\n        return {\n            'clipwise_logits_long': clipwise_logits_long,\n            'framewise_logits_long': framewise_logits_long,\n        }\n        \n\nclass WellHead(nn.Module):\n    def __init__(\n        self,\n        in_chans,\n        num_class,\n        hidden_chans=512,\n        p=0.2,\n    ):\n        super().__init__()\n        \n        self.time_dense_layers = nn.Sequential(\n            nn.Dropout(p),\n            nn.Linear(in_chans, hidden_chans),\n            nn.GELU(),\n            nn.Dropout(p),\n        )\n        self.freq_dense_layers = nn.Sequential(\n            nn.Dropout(p),\n            nn.Linear(in_chans, hidden_chans),\n            nn.GELU(),\n            nn.Dropout(p),\n        )\n        self.fea_dense_layers = nn.Sequential(\n            nn.Dropout(p),\n            nn.Linear(in_chans, hidden_chans),\n            nn.GELU(),\n            nn.Dropout(p),\n        )\n        self.time_attention = nn.Conv1d(\n            in_channels=hidden_chans,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n        self.freq_attention = nn.Conv1d(\n            in_channels=hidden_chans,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n        self.fix_scale = nn.Conv2d(\n            in_channels=hidden_chans,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n\n    def forward(self, fea):\n        # shape: [bs, ch, freq, time]\n        fea_time = torch.mean(fea.clamp(min=1e-6)**3, dim=2) ** (1.0 / 3)\n        fea_freq = torch.mean(fea.clamp(min=1e-6)**3, dim=3) ** (1.0 / 3)\n        \n        fea_time = fea_time.permute(0, 2, 1)  # (bs, time, ch)\n        fea_freq = fea_freq.permute(0, 2, 1)  # (bs, freq, ch)\n        fea_time = self.time_dense_layers(fea_time).permute(0, 2, 1) # (bs, 512, time)\n        fea_freq = self.freq_dense_layers(fea_freq).permute(0, 2, 1) # (bs, 512, freq)\n\n        time_att = self.time_attention(fea_time).softmax(dim=2) # [bs, cls, time]\n        freq_att = self.freq_attention(fea_freq).softmax(dim=2) # [bs, cls, freq]\n        att = time_att[..., None, :] * freq_att[..., :, None] # [bs, cls, freq, time]\n        \n        fea = self.fea_dense_layers(fea.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) # [bs, 512, freq, time]\n        fea_v = self.fix_scale(fea) # [bs, cls, freq, time]\n        \n        clipwise_logits_long = torch.sum(fea_v * att, dim=(2, 3))\n        return clipwise_logits_long\n        \n\nclass CosineClassifier(nn.Module):\n    def __init__(self, in_channels, num_classes, tau=0.2, bias=True):\n        super().__init__()\n        self.in_channels = in_channels\n        self.num_classes = num_classes\n        self.weight = nn.Parameter(torch.randn(num_classes, in_channels), requires_grad=True)\n        self.bias = nn.Parameter(torch.zeros(num_classes), requires_grad=bias)\n        self.tau = tau\n        nn.init.xavier_normal_(self.weight)\n\n    def forward(self, x):\n        x = F.normalize(x, dim=-1)\n        weight = F.normalize(self.weight, dim=1)\n        x = (x / self.tau**0.5) @ (weight.t() / self.tau**0.5) + self.bias\n        return x\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.584312Z","iopub.execute_input":"2024-06-08T04:22:42.585009Z","iopub.status.idle":"2024-06-08T04:22:42.607866Z","shell.execute_reply.started":"2024-06-08T04:22:42.584983Z","shell.execute_reply":"2024-06-08T04:22:42.606968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"class EncoderDecoder(nn.Module):\n    def __init__(self, encoder, decoder):\n        super().__init__()\n        self.encoder = encoder\n        self.decoder = decoder\n    \n    def forward(self, x):\n        x = self.encoder(x)\n        x = self.decoder(x)\n        return x\n    \nclass LambdaLayer(nn.Module):\n    def __init__(self, func, module=None):\n        super().__init__()\n        self.module = module\n        self.func = func\n    \n    def forward(self, *args, **kwargs):\n        return self.func(*args, **kwargs)\n\n\nclass TestBirdCLEFWrapper(nn.Module):\n    def __init__(self, module, train_info, val_info):\n        super().__init__()\n        self.module = module\n        reduced_to_base_classes = torch.zeros((len(val_info['base_classes']), len(train_info['reduced_classes'])))\n        for i, bc in enumerate(val_info['base_classes']):\n            j = train_info['reduced_classes'].index(bc)\n            reduced_to_base_classes[i, j] = 1\n        self.register_buffer('reduced_to_base_classes', reduced_to_base_classes)\n    \n    def forward(self, *args, **kwargs):\n        logits = self.module(*args, **kwargs)\n        is_list = isinstance(logits, (list, tuple))\n        if not is_list:\n            logits = [logits]\n        logits = [torch.tensordot(l, self.reduced_to_base_classes, dims=[[1], [1]]) for l in logits]\n        logits = [l.permute(0, -1, *list(range(1, len(l.shape) - 1))) for l in logits]\n        if not is_list:\n            logits = logits[0]\n        return logits\n\nclass TestBirdCLEFWrapperV2(nn.Module):\n    def __init__(self, module, train_info, val_info):\n        super().__init__()\n        self.module = module\n        reduced_to_base_classes = torch.zeros((len(val_info['base_classes']), len(train_info['reduced_classes'])))\n        for i, bc in enumerate(val_info['base_classes']):\n            j = train_info['reduced_classes'].index(bc)\n            reduced_to_base_classes[i, j] = 1\n\n        self.weight = nn.Parameter(torch.randn(len(val_info['base_classes']), self.module.weight.shape[1]))\n        self.bias = nn.Parameter(torch.zeros(len(val_info['base_classes']))) if module.bias is not None else None\n        with torch.no_grad():\n            self.weight.copy_(reduced_to_base_classes @ module.weight)\n            if self.bias is not None:\n                self.bias.copy_((reduced_to_base_classes @ module.bias.unsqueeze(1)).squeeze(1))\n            \n    def forward(self, fea):\n        logits = fea @ self.weight.t()\n        if self.bias is not None:\n            logits += self.bias.unsqueeze(0)\n        return logits\n\nclass Ensemble(nn.ModuleList):\n    def forward(self, x):\n        outputs = [model(x) for model in self]\n        outputs = sum(outputs) / len(outputs)\n        return outputs\n\nclass MultiModels(nn.Module):\n    def __init__(self, model_list):\n        super().__init__()\n        self.model_list = nn.ModuleList(model_list)\n\n    def forward(self, x):\n        outputs = [model(x) for model in self.model_list]\n        return outputs\n\nclass MultiDecoders(nn.Module):\n    def __init__(self, model_list):\n        super().__init__()\n        self.model_list = nn.ModuleList(model_list)\n\n    def forward(self, inputs):\n        outputs = [model(x) for model, x in zip(self.model_list, inputs)]\n        return outputs\n\n\ndef get_att_decoder(in_channels, num_classes):\n    head = AttHead(in_channels, p=0.5, num_class=num_classes,\n                   exportable=True, train_period= 10, infer_period= 10,)\n    def func(x):\n        outputs = head(x)\n        # outputs = [outputs['clipwise_logits_long'], outputs['framewise_logits_long'].max(dim=1)[0]]\n        outputs = outputs['clipwise_logits_long']\n        return outputs\n    return LambdaLayer(func, head)\n\nclass TimmLambda1(nn.Module):\n    def __init__(self, module):\n        super().__init__()\n        self.module = module\n    def forward(self, x):\n        return self.module(x[:, None])[-1]\nclass TimmLambda2(nn.Module):\n    def forward(self, x):\n        return x[:, :, 0, 0]\n\n    \ndef get_model(name, decoder_type, info, split_models=True):\n    if isinstance(name, (list, tuple, omegaconf.listconfig.ListConfig)):\n        model_list = [get_model(m, decoder_type, info) for m in name]\n        if split_models:\n            encoder_list, decoder_list = list(zip(*[(model.encoder, model.decoder) for model in model_list]))\n            encoder, decoder = MultiModels(encoder_list), MultiDecoders(decoder_list)\n            model =  EncoderDecoder(encoder, decoder)\n        else:\n            model = MultiModels(model_list)\n        return model\n    timm_model_names = {\n        'efficientnetv2': 'tf_efficientnetv2_s_in21k',\n        'convnext_small': 'convnext_small.fb_in22k_ft_in1k_384',\n        'convnextv2_tiny': 'convnextv2_tiny.fcmae_ft_in22k_in1k_384',\n        'eca_nfnet': 'eca_nfnet_l0',\n        'pvt_v2_b2': 'pvt_v2_b2',\n        'pvt_v2_b1': 'pvt_v2_b1',\n        'pvt_v2_b0': 'pvt_v2_b0'\n    }\n\n    if name == 'ast':\n        # ast_encoder = ASTModel.from_pretrained(\"MIT/ast-finetuned-audioset-10-10-0.4593\")\n        ast_encoder = ASTModel.from_pretrained(\"/home/lab/.cache/huggingface/hub/models--MIT--ast-finetuned-audioset-10-10-0.4593/snapshots/f826b80d28226b62986cc218e5cec390b1096902/\")\n\n        if decoder_type == 'fc':\n            model = EncoderDecoder(LambdaLayer(lambda x: ast_encoder(x)['pooler_output'], ast_encoder), \n                                   nn.Linear(768, len(info['reduced_classes']), bias=True))\n        elif decoder_type == 'att':\n            att_head = get_att_decoder(768, len(info['reduced_classes']))\n            model = EncoderDecoder(LambdaLayer(lambda x: ast_encoder(x)['last_hidden_state'], ast_encoder), \n                                   LambdaLayer(lambda x: att_head(x.permute(0, 2, 1).unsqueeze(2)), att_head))\n            \n        elif decoder_type == 'ocr':\n            ocr_head = OCRHead(768, len(info['reduced_classes']))\n            model = EncoderDecoder(LambdaLayer(lambda x: ast_encoder(x)['last_hidden_state'], ast_encoder), \n                                   LambdaLayer(lambda x: ocr_head(x.permute(0, 2, 1)), ocr_head))\n        \n    elif name in timm_model_names:\n        _backbone = timm.create_model(\n            timm_model_names[name],\n            features_only=True,\n            pretrained=True,\n            in_chans=1,\n        )\n        if decoder_type == 'fc':\n            backbone = nn.Sequential(TimmLambda1(_backbone), \n                                      nn.AdaptiveAvgPool2d((1, 1)), \n                                      TimmLambda2())\n            classifier = nn.Linear(_backbone.feature_info.channels()[-1], \n                                   len(info['reduced_classes']))\n            \n        elif decoder_type == 'well':\n            backbone = LambdaLayer(lambda x: _backbone(x[:, None])[-1], _backbone)\n            classifier = WellHead(_backbone.feature_info.channels()[-1], len(info['reduced_classes']))\n            \n        elif decoder_type == 'att':\n            backbone = LambdaLayer(lambda x: _backbone(x[:, None])[-1], _backbone)\n            classifier = get_att_decoder(_backbone.feature_info.channels()[-1], len(info['reduced_classes']))\n\n        elif decoder_type == 'cosine':\n            backbone = nn.Sequential(LambdaLayer(lambda x: _backbone(x[:, None])[-1], _backbone), \n                                              nn.AdaptiveAvgPool2d((1, 1)), \n                                              LambdaLayer(lambda x: x[:, :, 0, 0]),)\n            classifier = CosineClassifier(_backbone.feature_info.channels()[-1], \n                                   len(info['reduced_classes']), bias=False)\n\n        model = EncoderDecoder(backbone, \n                               classifier)\n        \n    return model\n    \ndef split_multi_models(name_list, decoder_type, checkpoint, info_path):\n    full_name = '_'.join(name_list)\n    model = get_model(name_list, decoder_type, load_dataset_information(info_path))\n    state_dict = torch.load(checkpoint)   \n    model.load_state_dict(state_dict)\n    model_list = [EncoderDecoder(encoder, decoder) for encoder, decoder in zip(model.encoder.model_list, model.decoder.model_list)]\n    for model, name in zip(model_list, name_list):\n        torch.save(model.state_dict(), checkpoint.replace(full_name, name))\n# model = get_model('pvt_v2', info)\n# x, y, *_ = next(train_iter)\n# spec = processor_volodymyr(x)\n# spec = ast_processor_batch_func(x)\n# model(torch.randn(1, 128, 128)).shape\n# print(\n#     sum([param.element_size() * param.numel() for param in model.parameters() if param.requires_grad]) / 1024**2 / 4, {k:sum([param.element_size() * param.numel() for param in v.parameters() if param.requires_grad]) / 1024**2 / 4 for k, v in model._modules.items()} \n# )\n\n# ast_encoder\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.609357Z","iopub.execute_input":"2024-06-08T04:22:42.609876Z","iopub.status.idle":"2024-06-08T04:22:42.650618Z","shell.execute_reply.started":"2024-06-08T04:22:42.609846Z","shell.execute_reply":"2024-06-08T04:22:42.649813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss and Metric","metadata":{}},{"cell_type":"code","source":"class SoftLALoss(nn.Module):\n    def __init__(self, c, lambda_c):\n        super().__init__()\n        self.lambda_c = lambda_c\n        c = torch.FloatTensor(c).clip(1e-7)\n        c = c / c.max()\n        self.register_buffer('log_c', torch.log(c))\n        \n    def forward(self, logits, targets):\n        z = torch.log_softmax(logits + self.lambda_c * self.log_c[None, :], dim=1)\n        loss = torch.sum(targets * z, dim=1)\n        return - loss.mean()\n        \nclass SoftBinaryFocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n    def forward(self, logits, targets):\n        p = torch.sigmoid(logits)\n        ce_loss = F.binary_cross_entropy_with_logits(logits, targets, reduction=\"none\")\n        p_t = p * targets + (1 - p) * (1 - targets)\n        loss = ce_loss * ((1 - p_t) ** self.gamma)\n        return loss.mean()\n\n# class FocalLossBCE(nn.Module):\n#     def __init__(\n#             self,\n#             alpha: float = 0.25,\n#             gamma: float = 2,\n#             reduction: str = \"mean\",\n#             bce_weight: float = 1.0,\n#             focal_weight: float = 1.0,\n#     ):\n#         super().__init__()\n#         self.alpha = alpha\n#         self.gamma = gamma\n#         self.reduction = reduction\n#         self.bce = torch.nn.BCEWithLogitsLoss(reduction=reduction)\n#         self.bce_weight = bce_weight\n#         self.focal_weight = focal_weight\n\n#     def forward(self, logits, targets):\n#         focall_loss = torchvision.ops.focal_loss.sigmoid_focal_loss(\n#             inputs=logits,\n#             targets=targets,\n#             alpha=self.alpha,\n#             gamma=self.gamma,\n#             reduction=self.reduction,\n#         )\n#         bce_loss = self.bce(logits, targets)\n#         return self.bce_weight * bce_loss + self.focal_weight * focall_loss\n        \nclass FocalLossBCE(nn.Module):\n    def __init__(\n        self,\n        alpha: float = 0.25,\n        gamma: float = 2,\n        reduction: str = \"mean\",\n    ):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, inputs, targets):\n        return torchvision.ops.focal_loss.sigmoid_focal_loss(\n            inputs=inputs,\n            targets=targets,\n            alpha=self.alpha,\n            gamma=self.gamma,\n            reduction=self.reduction,\n        )\n        \nclass SoftFocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n    def forward(self, logits, targets):\n        ce = - torch.sum(targets * torch.log_softmax(logits, dim=1), dim=1)\n        p = torch.exp(-ce)\n        loss = (1 - p) ** self.gamma * ce\n        return loss.mean()\n\nclass SoftRecallLoss(nn.Module):\n    def __init__(self, num_classes, momentum=0.1):\n        super().__init__()\n        self.num_classes = num_classes\n        self.err = nn.Parameter(torch.ones(num_classes) / num_classes, requires_grad=False)\n        self.momentum = momentum\n\n    def forward(self, logits, targets):\n        dims = (1, 2, 3)\n\n        with torch.no_grad():\n            pred_prob = torch.softmax(logits, dim=1) # b, c, h, w\n            cm = torch.sum(targets.unsqueeze(2) * pred_prob.unsqueeze(1), dim=(3, 4))\n            tp = cm.diagonal(dim1=1, dim2=2)\n            fn = cm.sum(dim=1) - tp\n            err = fn / (tp + fn).clip_(1e-6)\n\n            err = (err + 1e-6) / (err.sum(dim=1, keepdim=True) + 1e-6)\n            err = err.mean(dim=0)\n            \n            self.err.copy_((1 - self.momentum) * self.err + self.momentum * err)\n\n        losses = torch.sum(self.err.unsqueeze(0) * targets * torch.log_softmax(logits, dim=1), dim=1)\n        loss = - losses.mean()\n        return loss\n        \nclass SoftCrossEntropy(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n    def forward(self, logits, targets):\n        losses = torch.sum(targets * torch.log_softmax(logits, dim=1), dim=1)\n        loss = - losses.mean()\n        return loss\n\nclass SoftARBLoss(nn.Module):\n    def __init__(self, c):\n        super().__init__()\n        c = torch.FloatTensor(c).clip(1e-7)\n        c = c / c.max()\n        self.register_buffer('log_c', torch.log(c))\n        \n    def forward(self, logits, targets):\n        z = torch.log_softmax(logits + self.log_c[None, :], dim=1)\n        loss = torch.sum(targets * z, dim=1)\n        return - loss.mean()\n    \nclass SoftBinaryARBLoss(nn.Module):\n    def __init__(self, c):\n        super().__init__()\n        c = torch.FloatTensor(c).clip(1e-7)\n        c = c / c.max()\n        c = torch.stack([c.sum() - c, c], dim=0)\n        c = c / c.max(dim=0, keepdim=True)[0]\n        self.register_buffer('log_c', torch.log(c))\n        \n    def forward(self, logits, targets):  \n        targets = torch.stack([1 - targets, targets], dim=1)\n        logits = torch.stack([-logits, logits], dim=1)\n        logits = logits + self.log_c[None]\n        z = torch.log_softmax(logits, dim=1)\n        loss = torch.sum(targets * z, dim=1)\n        return - loss.mean()\n\nclass SoftAccuracy(MeanMetric):\n    def __init__(self, weights=None):\n        super().__init__()\n        self.weights = weights\n        self.sum_value = torch.zeros([])\n        self.count = 0\n    \n    def update(self, logits, targets):\n        logits = logits.float()\n        targets = targets.float()\n        probs = torch.softmax(logits, dim=1)\n        acc = probs * targets\n        if self.weights is not None:\n            weights = torch.FloatTensor(self.weights).to(logits.dtype).to(logits.device)\n            weights = weights / weights.sum()\n            acc *= weights[None, :]\n        acc = acc.sum(dim=1)\n        self.sum_value += acc.sum().cpu()\n        self.count += logits.shape[0]\n        \n    def compute(self):   \n        return self.sum_value / self.count\n    \n    def reset(self):\n        self.sum_value = torch.zeros([])\n        self.count = 0\n        \nclass RankAtSoftAccuracy(MeanMetric):\n    def __init__(self, ranks, num_classes):\n        if not isinstance(ranks, (list, tuple)):\n            ranks = [ranks]\n        super().__init__()\n        self.ranks = ranks\n        self.num_classes = num_classes\n        self.sum_value = torch.zeros(self.num_classes)\n        self.count = torch.zeros(self.num_classes)\n        \n    def update(self, logits, targets):\n        logits = logits.float()\n        targets = targets.float()\n        probs = torch.softmax(logits, dim=1)\n        acc = probs * targets    \n        self.sum_value += acc.sum(dim=0).cpu()\n        self.count += targets.sum(dim=0).cpu()\n        \n    def compute(self):   \n        acc = self.sum_value / self.count.clip(1e-7)\n        sort_acc, _ = torch.sort(acc)\n        sort_acc = sort_acc[self.ranks]\n        if len(sort_acc) == 1:\n            sort_acc = sort_acc[0]\n        return sort_acc\n    \n    def reset(self):\n        self.sum_value = torch.zeros(self.num_classes)\n        self.count = torch.zeros(self.num_classes)\n\nclass MacroSoftAccuracy(MeanMetric):\n    def __init__(self, num_classes):\n        super().__init__()\n        self.num_classes = num_classes\n        self.sum_value = torch.zeros(self.num_classes)\n        self.count = torch.zeros(self.num_classes)\n        \n    def update(self, logits, targets):\n        logits = logits.float()\n        targets = targets.float()\n        probs = torch.softmax(logits, dim=1)\n        acc = probs * targets\n        self.sum_value += acc.sum(dim=0).cpu()\n        self.count += targets.sum(dim=0).cpu()\n        \n    def compute(self):   \n        acc = self.sum_value / self.count.clip(1e-7)\n        acc = acc.mean()\n        return acc\n    \n    def reset(self):\n        self.sum_value = torch.zeros(self.num_classes)\n        self.count = torch.zeros(self.num_classes)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.651798Z","iopub.execute_input":"2024-06-08T04:22:42.652075Z","iopub.status.idle":"2024-06-08T04:22:42.692137Z","shell.execute_reply.started":"2024-06-08T04:22:42.652053Z","shell.execute_reply":"2024-06-08T04:22:42.691197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Trainer","metadata":{}},{"cell_type":"markdown","source":"## Audio Trainer","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def unitwise_norm(x, norm_type=2.0):\n    if x.ndim <= 1:\n        return x.norm(norm_type)\n    else:\n        # works for nn.ConvNd and nn,Linear where output dim is first in the kernel/weight tensor\n        # might need special cases for other weights (possibly MHA) where this may not be true\n        return x.norm(norm_type, dim=tuple(range(1, x.ndim)), keepdim=True)\n\n\ndef adaptive_clip_grad(parameters, clip_factor=0.01, eps=1e-3, norm_type=2.0):\n    if isinstance(parameters, torch.Tensor):\n        parameters = [parameters]\n    for p in parameters:\n        if p.grad is None:\n            continue\n        p_data = p.detach()\n        g_data = p.grad.detach()\n        max_norm = unitwise_norm(p_data, norm_type=norm_type).clamp_(min=eps).mul_(clip_factor)\n        grad_norm = unitwise_norm(g_data, norm_type=norm_type)\n        clipped_grad = g_data * (max_norm / grad_norm.clamp(min=1e-6))\n        new_grads = torch.where(grad_norm < max_norm, g_data, clipped_grad)\n        p.grad.detach().copy_(new_grads)\n        \ndef get_mixup_pair_permutation(batch_size, p=0.5):\n    p0 = 1 - 1 / batch_size\n    indices1 = torch.randperm(batch_size)\n    if p > p0:\n        if  random.random() < batch_size * (p - p0):\n            indices2 = torch.roll(indices1, 1)\n        else:\n            indices2 = torch.randperm(batch_size)\n    else:\n        if random.random() > 1 / p0 * (p0 - p):\n            indices2 = torch.randperm(batch_size)\n        else:\n            indices2 = indices1\n    return indices1, indices2\n\ndef get_mixup_pair_samples(x, y=None, alpha=None, p=0.5):\n    indices1, indices2 = get_mixup_pair_permutation(x.shape[0], p=p)\n    x1, x2 = x[indices1], x[indices2]\n    alpha = alpha.view(-1, *(len(x.shape) - 1) * [1])\n    x = alpha * x1 + (1 - alpha) * x2\n    if y is not None:\n        y1, y2 = y[indices1], y[indices2]\n        alpha = alpha.view(-1, *(len(y.shape) - 1) * [1])\n        y = alpha * y1 + (1 - alpha) * y2\n        return x, y\n    return x\n\ndef compute_loss(logits, targets, criterion, loss_coefs=None, return_targets=True, return_logits=True):\n    if not isinstance(logits, (list, tuple)):\n        logits = [logits]\n    if not isinstance(targets, (list, tuple)):\n        targets = len(logits) * [targets]\n    logits = list(logits)\n    \n    for i in range(len(logits)):\n        if len(logits[i].shape) > 2:\n            targets[i] = targets[i].repeat(np.product(logits[i].shape[2:]), 1)\n            logits[i] = logits[i].permute(0, *list(range(2, len(logits[i].shape))), 1).view(-1, logits[i].shape[1])\n   \n    if loss_coefs is None:\n        loss = sum([criterion(l, y) for l, y in zip(logits, targets)])\n    else:\n        loss = sum([coef * criterion(l, y) for l, y, coef in zip(logits, targets, loss_coefs)])\n\n    outputs = [loss]\n    if return_targets:\n        outputs.append(targets)\n    if return_logits:\n        outputs.append(logits)\n    if len(outputs) == 1:\n        outputs = outputs[0]\n    return outputs\n\nclass AudioTrainer:\n    def __init__(self, model, processor, spec_augmentations, \n                 optimizer, criterion, class_counts, num_accum_steps=1, \n                 freeze_encoder=False, loss_coefs=None,\n                 lr_scheduler=None, use_cuda=None, use_amp=True, use_ema=False):\n        if torch.cuda.device_count() > 1:\n            model.encoder = nn.DataParallel(model.encoder)\n            model.decoder = nn.DataParallel(model.decoder)\n        ### set models ###\n        self.model = model\n        self.use_ema = use_ema\n        self.avg_model = ModelEMA(self.model) if self.use_ema else model\n        ##################\n        \n        self.processor = processor\n        self.spec_augmentations = spec_augmentations\n        self.optimizer = optimizer\n        self.criterion = criterion\n        self.lr_scheduler = lr_scheduler\n        self.num_accum_steps = num_accum_steps\n        self.class_counts = torch.FloatTensor(class_counts)\n        self.class_counts = self.class_counts / self.class_counts.sum()\n        self.freeze_encoder = freeze_encoder\n        self.loss_coefs = loss_coefs\n        self._train_count = 0\n\n        self._initialize_cuda(use_cuda, use_amp)\n        self.grad_scaler = amp.GradScaler(enabled=self.use_amp)\n        \n        ### define metrics ###\n        self.metrics = {\n            'loss': MeanMetric(),\n            'acc': SoftAccuracy(),\n            'wacc': MacroSoftAccuracy(len(class_counts)),\n            'hacc': RankAtSoftAccuracy(-1, len(class_counts)),\n            'lacc': RankAtSoftAccuracy(0, len(class_counts)),\n        }\n        ######################\n    \n    def fit(self, train_loader, val_loader=None, epochs=1, verbose=2):\n        history = {name: [] for name in self.metrics.keys()}\n        if val_loader is not None:  # return history for validation data\n            history = {**history,\n                       **{'val_' + name: [] for name in self.metrics.keys()}}\n            \n        epoch_iter = range(1, epochs + 1)\n        if verbose == 1:\n            epoch_iter = tqdm(epoch_iter, total=epochs, ncols=25*len(self.metrics))\n            epoch_iter.set_description(f'Epochs: ')\n        for epoch in epoch_iter:\n            ### model setting ###\n            if self.use_cuda:\n                if isinstance(self.processor, nn.Module):\n                    self.processor.cuda()\n                self.spec_augmentations.cuda()\n                self.model.cuda()\n                self.avg_model.cuda()\n                self.criterion.cuda()\n            if self.freeze_encoder:\n                self.model.encoder.eval()\n                self.model.decoder.train()\n            else:\n                self.model.train()\n            ######################\n\n            # train_loader.sampler.set_epoch(epoch)\n            train_iter = iter(train_loader)\n            if verbose == 2:\n                train_iter = tqdm(train_iter, total=len(train_loader), ncols=sum([len(k) + 6 for k in self.metrics.keys()]) + 100)\n                train_iter.set_description(f'Epoch {epoch}/{epochs}')\n            self._reset_metrics()\n\n            for step, data in enumerate(train_iter):\n                self.train_step(data)\n                if verbose == 1:\n                    epoch_iter.set_postfix({'batch': f'{step+1}/{len(train_loader)}', **self._compute_metrics()})\n                elif verbose == 2:  # show metrics on train data\n                    train_iter.set_postfix(self._compute_metrics())\n            self._update_history(history)\n\n            if val_loader is not None:\n                self.evaluate(val_loader, verbose=verbose)\n                self._update_history(history, prefix='val_')\n        return history\n\n    def evaluate(self, val_loader, verbose=2):\n        ### model setting ###\n        if self.use_cuda:\n            if isinstance(self.processor, nn.Module):\n                self.processor.cuda()\n            self.spec_augmentations.cuda()\n            self.model.cuda()\n            self.avg_model.cuda()\n            self.criterion.cuda()\n        self.model.eval()\n        self.avg_model.eval()\n        #####################\n\n        val_iter = iter(val_loader)\n        if verbose == 2:\n            val_iter = tqdm(val_iter, total=len(val_loader), ncols=25*len(self.metrics))\n            val_iter.set_description(f'Eval')\n\n        self._reset_metrics()\n\n        for step, data in enumerate(val_iter):\n            self.test_step(data)\n            if verbose == 2:\n                val_iter.set_postfix(self._compute_metrics())  # show metrics on val data\n        return self._compute_metrics()\n\n    def predict(self, data_loader, verbose=True, **kwargs):\n        ### model setting ###\n        if self.use_cuda:\n            if isinstance(self.processor, nn.Module):\n                self.processor.cuda()\n            self.spec_augmentations.cuda()\n            self.model.cuda()\n            self.avg_model.cuda()\n        self.model.eval()\n        self.avg_model.eval()\n        #####################\n\n        data_iter = iter(data_loader)\n        if verbose:\n            data_iter = tqdm(data_iter, total=len(data_loader))\n\n        results = []\n        for step, data in enumerate(data_iter):\n            outputs = self.predict_step(data, **kwargs)\n            results.append(outputs)\n\n        if isinstance(outputs, (tuple, list)):  # multi outputs\n            results = list(zip(*results))\n            results = [torch.cat(tensor, dim=0) for tensor in results]\n        else:\n            results = torch.cat(results, dim=0)\n        return results\n\n    def train_step(self, data):\n        global x\n        ### config input data ###\n        data = self._convert_cuda_data(data)\n        x, y = data\n        mix_up_on_freq = random.random() > 0.5\n        if not mix_up_on_freq:\n            alpha = torch.rand((x.shape[0],)).to(x.device)\n            x, y = get_mixup_pair_samples(x, y, alpha=alpha, p=0.5)\n        x = self.processor(x)\n        if mix_up_on_freq:\n            alpha = torch.distributions.Beta(2, 2).rsample((x.shape[0],)).to(x.device)\n            x, y = get_mixup_pair_samples(x, y, alpha, p=0.5)\n        x = self.spec_augmentations(x)\n        #########################\n        \n        #############################\n        with self.autocast():\n            ### forward pass ###\n            with (torch.no_grad() if self.freeze_encoder else nullcontext()):\n                fea = self.model.encoder(x)\n            logits = self.model.decoder(fea)\n            loss, y, logits = compute_loss(logits, y, self.criterion, loss_coefs=self.loss_coefs, \n                                           return_targets=True, return_logits=True)\n            \n            ####################\n        \n        ### update model ###              \n        self.grad_scaler.scale(loss).backward()\n        if (self._train_count + 1) % self.num_accum_steps == 0:\n            self.grad_scaler.unscale_(self.optimizer)\n            # adaptive_clip_grad(self.model.parameters())\n            torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1e9)\n            self.grad_scaler.step(self.optimizer)\n            self.grad_scaler.update()\n            self.optimizer.zero_grad()\n            if self.use_ema:\n                self.avg_model.update_parameters(self.model)\n            if self.lr_scheduler is not None:\n                self.lr_scheduler.step()\n        self._train_count += 1\n        #####################\n\n        ### update metrics ###\n        logits = logits[0]\n        y = y[0]\n        \n        loss, y, logits = self._convert_not_training_data([loss, y, logits])\n        \n        self.metrics['loss'].update(loss)\n        self.metrics['acc'].update(logits, y)\n        self.metrics['wacc'].update(logits, y)\n        self.metrics['lacc'].update(logits, y)\n        self.metrics['hacc'].update(logits, y)\n\n    def test_step(self, data):\n        ### config input data ###\n        data = self._convert_cuda_data(data)\n        x, y = data\n        x = self.processor(x)\n        #########################\n\n        with self.autocast():\n            with torch.no_grad():\n                ### forward pass ###\n                fea = self.model.encoder(x)\n                logits = self.model.decoder(fea)\n                loss, y, logits = compute_loss(logits, y, self.criterion, loss_coefs=self.loss_coefs, \n                                               return_targets=True, return_logits=True)\n                ####################\n\n        ### update metrics ###\n        logits = logits[0]\n        y = y[0]\n        loss, y, logits = self._convert_not_training_data([loss, y, logits])\n\n        self.metrics['loss'].update(loss)\n        self.metrics['acc'].update(logits, y)\n        self.metrics['wacc'].update(logits, y)\n        self.metrics['lacc'].update(logits, y)\n        self.metrics['hacc'].update(logits, y)\n\n        \n        #######################\n        return\n\n    def predict_step(self, data, return_fea=False, return_gt=False):\n        ### config input data ###\n        data = self._convert_cuda_data(data)\n        x, y = data\n        x = self.processor(x)\n        #########################\n\n        with self.autocast():\n            with torch.no_grad():\n                ### forward pass ###\n                fea = self.model.encoder(x)\n                logits = self.model.decoder(fea)\n                ####################\n        if isinstance(logits, (list, tuple)):\n            logits = logits[0]\n            \n        outputs = [logits.detach().cpu()]\n        if return_fea:\n            outputs.append(fea.detach().cpu())\n        if return_gt:\n            outputs.append(y.detach().cpu())\n        return outputs\n\n    def _initialize_cuda(self, use_cuda=None, use_amp=True):\n        if use_cuda is None:\n            use_cuda = torch.cuda.is_available()\n        self.use_cuda = use_cuda\n        self.use_amp = use_amp and self.use_cuda\n        self.autocast = amp.autocast if self.use_amp else nullcontext\n\n    def _convert_not_training_data(self, data, detach=True, cpu=True, numpy=False):\n        if isinstance(data, (tuple, list)):\n            return [self._convert_not_training_data(e) for e in data]\n        elif isinstance(data, dict):\n            return {k: self._convert_not_training_data(v) for k, v in data.items()}\n        else:\n            if detach:\n                data = data.detach()\n            if cpu:\n                data = data.cpu()\n            if numpy:\n                data = data.numpy()\n            return data\n\n    def _convert_cuda_data(self, data):\n        if not self.use_cuda:\n            return data\n        if isinstance(data, (tuple, list)):\n            return [self._convert_cuda_data(e) for e in data]\n        elif isinstance(data, dict):\n            return {k: self._convert_cuda_data(v) for k, v in data.items()}\n        else:\n            return data.cuda()\n\n    def _reset_metrics(self):\n        for metric in self.metrics.values():\n            metric.reset()\n\n    def _compute_metrics(self):\n        return {name: float(metric.compute()) for name, metric in self.metrics.items()}\n\n    def _update_history(self, history, prefix=''):\n        metrics = self._compute_metrics()\n        for name, value in metrics.items():\n            history[prefix + name].append(value)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.693585Z","iopub.execute_input":"2024-06-08T04:22:42.694142Z","iopub.status.idle":"2024-06-08T04:22:42.762877Z","shell.execute_reply.started":"2024-06-08T04:22:42.694106Z","shell.execute_reply":"2024-06-08T04:22:42.761793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## UDA Trainer","metadata":{}},{"cell_type":"code","source":"\ndef get_mix_pair_samples(x1, y1, x2, y2, alpha):\n    alpha = alpha.view(-1, *(len(x1.shape) - 1) * [1])\n    x = alpha * x1 + (1 - alpha) * x2\n    alpha = alpha.view(-1, *(len(y1.shape) - 1) * [1])\n    y = alpha * y1 + (1 - alpha) * y2\n    return x, y\n\nclass ModelEMA(nn.Module):\n    def __init__(self, model, decay=0.9999):\n        super().__init__()\n        self.module = deepcopy(model)\n        self.module.eval()\n        self.decay = decay\n\n    def forward(self, input):\n        return self.module(input)\n\n    def _update(self, model, update_fn):\n        with torch.no_grad():\n            for ema_v, model_v in zip(self.module.parameters(), model.parameters()):\n                ema_v.copy_(update_fn(ema_v, model_v))\n            for ema_v, model_v in zip(self.module.buffers(), model.buffers()):\n                ema_v.copy_(model_v)\n\n    def update_parameters(self, model):\n        self._update(model, update_fn=lambda e, m: self.decay * e + (1. - self.decay) * m)\n\n    def state_dict(self):\n        return self.module.state_dict()\n\n    def load_state_dict(self, state_dict):\n        self.module.load_state_dict(state_dict)\n        \nclass UDAAudioTrainer(AudioTrainer):\n    def __init__(self, model, processor, spec_augmentations, \n                 optimizer, criterion, class_counts, num_accum_steps=1, \n                 freeze_encoder=False, loss_coefs=None, multilabel=False, \n                 uda_steps=1000, uda_coef=4,\n                 lr_scheduler=None, use_cuda=None, use_amp=True, use_ema=False):\n        ### set models ###\n        if torch.cuda.device_count() > 1:\n            model.encoder = nn.DataParallel(model.encoder)\n            model.decoder = nn.DataParallel(model.decoder)\n        self.model = model\n        self.use_ema = use_ema\n        self.avg_model = ModelEMA(self.model) if self.use_ema else model\n        ##################\n        \n        self.processor = processor\n        self.spec_augmentations = spec_augmentations\n        self.optimizer = optimizer\n        self.criterion = criterion\n        self.lr_scheduler = lr_scheduler\n        self.num_accum_steps = num_accum_steps\n        self.class_counts = torch.FloatTensor(class_counts)\n        self.class_counts = self.class_counts / self.class_counts.sum()\n        self.freeze_encoder = freeze_encoder\n        self.loss_coefs = loss_coefs\n        self.multilabel = multilabel\n        self.uda_steps = uda_steps\n        self.uda_coef = uda_coef\n        self._train_count = 0\n\n        self._initialize_cuda(use_cuda, use_amp)\n        self.grad_scaler = amp.GradScaler(enabled=self.use_amp)\n        \n        ### define metrics ###\n        self.metrics = {\n            'loss': MeanMetric(),\n            'loss_l': MeanMetric(),\n            'loss_u': MeanMetric(),\n            'acc': SoftAccuracy(),\n            'wacc': MacroSoftAccuracy(len(class_counts)),\n            'hacc': RankAtSoftAccuracy(-1, len(class_counts)),\n            'lacc': RankAtSoftAccuracy(0, len(class_counts)),\n        }\n        ######################\n\n    def fit(self, train_loader, uda_loader, val_loader=None, epochs=1, verbose=2):\n        history = {name: [] for name in self.metrics.keys()}\n        if val_loader is not None:  # return history for validation data\n            history = {**history,\n                       **{'val_' + name: [] for name in self.metrics.keys()}}\n            \n        epoch_iter = range(1, epochs + 1)\n        if verbose == 1:\n            epoch_iter = tqdm(epoch_iter, total=epochs, ncols=25*len(self.metrics))\n            epoch_iter.set_description(f'Epochs: ')\n        for epoch in epoch_iter:\n            ### model setting ###\n            if self.use_cuda:\n                if isinstance(self.processor, nn.Module):\n                    self.processor.cuda()\n                self.spec_augmentations.cuda()\n                self.model.cuda()\n                self.avg_model.cuda()\n                self.criterion.cuda()\n            if self.freeze_encoder:\n                self.model.encoder.eval()\n                self.model.decoder.train()\n            else:\n                self.model.train()\n            ######################\n\n            # train_loader.sampler.set_epoch(epoch)\n            train_iter = zip(iter(train_loader), iter(uda_loader))\n            if verbose == 2:\n                train_iter = tqdm(train_iter, total=len(train_loader), ncols=sum([len(k) + 6 for k in self.metrics.keys()]) + 100)\n                train_iter.set_description(f'Epoch {epoch}/{epochs}')\n            self._reset_metrics()\n\n            for step, data in enumerate(train_iter):\n                self.train_step(data)\n                if verbose == 1:\n                    epoch_iter.set_postfix({'batch': f'{step+1}/{len(train_loader)}', **self._compute_metrics()})\n                elif verbose == 2:  # show metrics on train data\n                    train_iter.set_postfix(self._compute_metrics())\n            self._update_history(history)\n\n            if val_loader is not None:\n                self.evaluate(val_loader, verbose=verbose)\n                self._update_history(history, prefix='val_')\n        return history\n\n    def _uda_preprocessing(self, x_l, y_l, x_uw, x_us):\n        x_u = torch.cat([x_uw, x_us], dim=0)\n        y_u = torch.zeros(x_uw.shape[0], y_l.shape[1], dtype=y_l.dtype, device=y_l.device)\n        \n        mix_up_on_freq = random.random() > 0.5\n        if not mix_up_on_freq:\n            alpha_l = torch.rand((x_l.shape[0],)).to(x_l.device)\n            alpha_u = torch.rand((x_u.shape[0],)).to(x_u.device)\n            x_l, y_l = get_mixup_pair_samples(x_l, y_l, alpha=alpha_l, p=0.5)\n            x_u = get_mixup_pair_samples(x_u, None, alpha=alpha_u, p=0.5)\n\n            alpha_lu = torch.rand((x_l.shape[0],)).to(x_l.device)\n            _x_l = torch.cat([x_l, x_l], dim=0)\n            _x_l = _x_l / _x_l.std(dim=-1, keepdim=True) * x_u.std(dim=-1, keepdim=True) # to x_u's rms\n            x_u, y_u = get_mix_pair_samples(x_u, torch.cat([y_u, y_u], dim=0), \n                                            _x_l, torch.cat([y_l, y_l], dim=0),\n                                            torch.cat([alpha_lu, alpha_lu], dim=0))\n            y_u = y_u[:x_uw.shape[0]]\n        \n        x_l = self.processor(x_l)\n        x_u = self.processor(x_u)\n        \n        if mix_up_on_freq:\n            alpha_l = torch.distributions.Beta(2, 2).rsample((x_l.shape[0],)).to(x_l.device)\n            alpha_u = torch.distributions.Beta(2, 2).rsample((x_u.shape[0],)).to(x_u.device)\n            x_l, y_l = get_mixup_pair_samples(x_l, y_l, alpha=alpha_l, p=0.5)\n            x_u = get_mixup_pair_samples(x_u, alpha=alpha_u, p=0.5)\n            \n            alpha_lu = torch.distributions.Beta(2, 2).rsample((x_l.shape[0],)).to(x_l.device)\n            x_u, y_u = get_mix_pair_samples(x_u, torch.cat([y_u, y_u], dim=0), \n                                            torch.cat([x_l, x_l], dim=0), torch.cat([y_l, y_l], dim=0),\n                                            torch.cat([alpha_lu, alpha_lu], dim=0))\n            y_u = y_u[:x_uw.shape[0]]\n                    \n        x_uw, x_us = torch.split(x_u, [x_uw.shape[0], x_us.shape[0]], dim=0)\n\n        x_l = self.spec_augmentations(x_l)\n        x_us = self.spec_augmentations(x_us)\n        return x_l, y_l, x_uw, x_us, y_u, alpha_lu\n        \n    def train_step(self, data):\n        ### config input data ###\n        uda_coef = self.uda_coef * min(1., (self._train_count + 1) / self.uda_steps)\n        \n        data = self._convert_cuda_data(data)\n        (x_l, y_l), (x_uw, x_us) = data\n        x_l, y_l, x_uw, x_us, y_u, alpha_lu = self._uda_preprocessing(x_l, y_l, x_uw, x_us)\n        \n        #########################\n        \n        #############################\n        with self.autocast():\n            ### forward pass ###\n            x = torch.cat([x_l, x_us], dim=0)\n            with (torch.no_grad() if self.freeze_encoder else nullcontext()):\n                fea = self.model.encoder(x)\n            logits = self.model.decoder(fea)\n            if isinstance(logits, (list, tuple)):\n                logits_l, logits_us = list(zip(*[l.split([x_l.shape[0], x_us.shape[0]], dim=0) for l in logits]))\n            else:\n                logits_l, logits_us = logits.split([x_l.shape[0], x_us.shape[0]], dim=0)\n            #     fea_l = self.model.encoder(x_l)\n            #     fea_us = self.model.encoder(x_us)\n            # logits_l = self.model.decoder(fea_l)\n            # logits_us = self.model.decoder(fea_us)\n            \n            self.model.eval()\n            with torch.no_grad():\n                logits_uw = self.model(x_uw)\n                if isinstance(logits_uw, (tuple, list)):\n                    y_u_pseudo = [torch.sigmoid(l) if self.multilabel else torch.softmax(l, dim=1) \n                                  for l in logits_uw]\n                    y_u = [y_u + alpha_lu[:, None] * y_p for y_p in y_u_pseudo]\n                else:\n                    y_u_pseudo = torch.sigmoid(logits_uw) if self.multilabel else torch.softmax(logits_uw, dim=1)\n                    y_u += alpha_lu[:, None] * y_u_pseudo\n            self.model.train()\n            \n            \n            \n            loss_l, y_l, logits_l = compute_loss(logits_l, y_l, self.criterion, loss_coefs=self.loss_coefs, \n                                                 return_targets=True, return_logits=True)\n            loss_u, y_u, logits_us = compute_loss(logits_us, y_u, self.criterion, loss_coefs=self.loss_coefs, \n                                                  return_targets=True, return_logits=True)\n\n            loss = loss_l + uda_coef * loss_u\n            ####################\n        \n        ### update model ###              \n        self.grad_scaler.scale(loss).backward()\n        if (self._train_count + 1) % self.num_accum_steps == 0:\n            self.grad_scaler.unscale_(self.optimizer)\n            # adaptive_clip_grad(self.model.parameters())\n            torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1e9)\n            self.grad_scaler.step(self.optimizer)\n            self.grad_scaler.update()\n            self.optimizer.zero_grad()\n            if self.use_ema:\n                self.avg_model.update_parameters(self.model)\n            if self.lr_scheduler is not None:\n                self.lr_scheduler.step()\n        self._train_count += 1\n        #####################\n\n        ### update metrics ###\n        logits = logits_l[0]\n        y = y_l[0]\n        \n        loss, y, logits, loss_l, loss_u = self._convert_not_training_data([loss, y, logits, loss_l, loss_u])\n        \n        self.metrics['loss'].update(loss)\n        self.metrics['loss_l'].update(loss_l)\n        self.metrics['loss_u'].update(loss_u)\n        self.metrics['acc'].update(logits, y)\n        self.metrics['wacc'].update(logits, y)\n        self.metrics['lacc'].update(logits, y)\n        self.metrics['hacc'].update(logits, y)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.764418Z","iopub.execute_input":"2024-06-08T04:22:42.764705Z","iopub.status.idle":"2024-06-08T04:22:42.823209Z","shell.execute_reply.started":"2024-06-08T04:22:42.764681Z","shell.execute_reply":"2024-06-08T04:22:42.822052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## SDAT Trainer","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Metric Learning Trainer","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Experts Trainer","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ds.meta_df.iloc[0]['filename']\n# ds.meta_df.iloc[0]['primary_label']\n# ds.meta_df.at[0, 'filename'] = ds.meta_df.at[0, 'primary_label'] + '/' + ds.meta_df.at[0, 'filename']\n# ds.meta_df\n# import multiprocessing\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.824496Z","iopub.execute_input":"2024-06-08T04:22:42.824934Z","iopub.status.idle":"2024-06-08T04:22:42.837675Z","shell.execute_reply.started":"2024-06-08T04:22:42.824867Z","shell.execute_reply":"2024-06-08T04:22:42.836800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"\ndef get_default_config():\n    cfg = OmegaConf.create()\n    cfg.dataset = {}\n    cfg.dataset.augmentation = {}\n    cfg.model = {}\n    cfg.train = {}\n    \n    cfg.seed = 1\n    \n    cfg.dataset.orig_sampling_rate = 32000\n    cfg.dataset.sampling_rate = 32000\n    cfg.dataset.audio_second = 5\n    cfg.dataset.use_secondary = True\n    cfg.dataset.root = '/kaggle/input'\n    cfg.dataset.roots = None\n    cfg.dataset.info_path = './info.json'\n    cfg.dataset.val_info_path = './info2024_primary.json'\n    cfg.dataset.kfold = 0\n    cfg.dataset.resample_val = False\n    cfg.dataset.samples_per_epoch = None\n    cfg.dataset.batch_size = 32\n    cfg.dataset.num_workers = 8\n    cfg.dataset.stage = 1\n    cfg.dataset.prefetch_factor = 64\n    cfg.dataset.use_cache = False\n    cfg.dataset.max_cache = 60000\n    \n    cfg.dataset.augmentation.processor = 'volodymyr' # volodymyr, ast\n    \n    cfg.model.name = 'ast' # 'ast', 'convnext_small', 'convnextv2_tiny', 'eca_nfnet', 'pvt_v2'\n    cfg.model.decoder_type = 'fc' # 'fc', 'att'\n    cfg.model.checkpoint = None\n    cfg.model.optimizer_checkpoint = None\n    cfg.model.freeze_encoder = False\n    cfg.model.re_init_decoder = False\n    cfg.model.use_multilabel = False\n    \n    cfg.train.learning_rate = 1e-4\n    cfg.train.weight_decay = 1e-6\n    cfg.train.warmup_epochs = 1\n    cfg.train.epochs = 10\n    cfg.train.finetune_epochs = 1     \n    cfg.train.loss_name = 'ce' # 'arb', 'ce', 'bce', 'binary-arb'\n    cfg.train.loss_coefs = None\n    cfg.train.uda = False\n    cfg.train.uda_coef = 2\n    cfg.train.uda_epochs = 2\n    cfg.train.sdat = False\n    cfg.train.use_mcc = True\n    cfg.train.use_adv = True\n    cfg.train.use_sam = True\n    cfg.train.metric_learning = False\n    cfg.train.multi_experts = False\n    cfg.train.optimizer = 'adam'\n    \n    cfg.output_dir = './outputs'\n    \n    return cfg\n\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    \n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    torch.use_deterministic_algorithms(True, warn_only=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.839215Z","iopub.execute_input":"2024-06-08T04:22:42.839791Z","iopub.status.idle":"2024-06-08T04:22:42.854799Z","shell.execute_reply.started":"2024-06-08T04:22:42.839755Z","shell.execute_reply":"2024-06-08T04:22:42.853787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG_DICT = {\n    'stage1': {\n        'dataset': {\n            'stage': 1,\n            'info_path': './info.json',\n        },\n        'train': {\n            'warmup_epochs': 0.5,\n            'epochs': 10,\n        },\n    },\n    'stage2': {\n        'dataset': {\n            'stage': 2,\n            'info_path': './info-2024.json',\n        },\n        'train': {\n            'warmup_epochs': 0.1,\n            'epochs': 2,\n        },\n        'model': {\n            'checkpoint': True,\n        }\n    },\n    'ast': {\n        'dataset': {\n            'sampling_rate': 16000,\n            'augmentation': {\n                'processor': 'ast',\n            },\n        }, \n        'model': {\n            'name': 'ast',\n        },\n        'train': {\n            'warmup_epochs': 0.5,\n            'epochs': 8,\n        },\n    },\n    'volodymyr': {\n        'dataset': {\n            'sampling_rate': 32000,\n            'augmentation': {\n                'processor': 'volodymyr',\n            },\n        }, \n    },\n    'binary-arb': {\n        'train': {\n            'loss_name': 'binary-arb',\n        }, \n    },\n    'use_cache': {\n        'dataset': {\n            'prefetch_factor': 8,\n            'use_cache': True,\n            'max_cache': 30000,\n        }, \n    },\n    'not_use_cache': {\n        'dataset': {\n            'prefetch_factor': 64,\n            'use_cache': False,\n        }, \n    },\n}\nCONFIG_DICT = {k:OmegaConf.create(v) for k,v in CONFIG_DICT.items()}","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.856142Z","iopub.execute_input":"2024-06-08T04:22:42.856444Z","iopub.status.idle":"2024-06-08T04:22:42.875620Z","shell.execute_reply.started":"2024-06-08T04:22:42.856419Z","shell.execute_reply":"2024-06-08T04:22:42.874770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg_list = [\n    [('model.name', ['efficientnetv2', 'convnextv2_tiny', 'eca_nfnet']), \n     ('model.checkpoint', 'outputs/1__sq_efficientnetv2_convnextv2_tiny_eca_nfnet_fc_ce_10_1_1_weights.pt'), \n     ('dataset.info_path', './info2024_p0.json'), ('dataset.stage', 2), ('train.warmup_epochs', 1), ('train.epochs', 20), ('dataset.kfold', 1), ('train.uda', True),],\n    [('model.name', ['efficientnetv2', 'convnextv2_tiny', 'eca_nfnet']), \n     ('model.checkpoint', 'outputs/1__sq_efficientnetv2_convnextv2_tiny_eca_nfnet_fc_ce_10_1_1_weights.pt'), \n     ('dataset.info_path', './info2024_p1.json'), ('dataset.stage', 2), ('train.warmup_epochs', 1), ('train.epochs', 20), ('dataset.kfold', 1), ('train.uda', True),],\n    [('model.name', ['efficientnetv2', 'convnextv2_tiny', 'eca_nfnet']), \n     ('model.checkpoint', 'outputs/1__sq_efficientnetv2_convnextv2_tiny_eca_nfnet_fc_ce_10_1_1_weights.pt'), \n     ('dataset.info_path', './info2024_p2.json'), ('dataset.stage', 2), ('train.warmup_epochs', 1), ('train.epochs', 20), ('dataset.kfold', 1), ('train.uda', True),],\n]\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.876722Z","iopub.execute_input":"2024-06-08T04:22:42.877056Z","iopub.status.idle":"2024-06-08T04:22:42.889474Z","shell.execute_reply.started":"2024-06-08T04:22:42.877022Z","shell.execute_reply":"2024-06-08T04:22:42.888642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"\ndef get_experiment_name(cfg):\n    _info_name = os.path.splitext(os.path.split(cfg.dataset.info_path)[1])[0][4:]\n    _model_name = cfg.model.name\n    if isinstance(_model_name, (tuple, list, omegaconf.listconfig.ListConfig)):\n        _model_name = '_'.join(_model_name)\n    experiment_name = f'{cfg.dataset.stage}_{_info_name}_{_model_name}_{cfg.model.decoder_type}_{cfg.train.loss_name}_{cfg.train.epochs}_{cfg.dataset.kfold}_{cfg.seed}'\n    return experiment_name\n    \ndef train(cfg):\n    global trainer, train_ds, val_ds\n    experiment_name = get_experiment_name(cfg)\n    print('Experiment name:', experiment_name)\n    \n    if cfg.seed is not None:\n        seed_everything(cfg.seed)\n    \n    train_ds, val_ds, train_loader, val_loader, transforms_dict, info = get_dataloaders(cfg)\n    uda_ds, uda_loader, uda_transforms_dict = get_uda_dataloaders(cfg)\n    \n    processor = get_spec_processor(cfg.dataset.sampling_rate, cfg.dataset.audio_second)[cfg.dataset.augmentation.processor]\n    \n    model = get_model(cfg.model.name, cfg.model.decoder_type, info)\n    print(\n        sum([param.element_size() * param.numel() for param in model.parameters() if param.requires_grad]) / 1024**2 / 4, {k:sum([param.element_size() * param.numel() for param in v.parameters() if param.requires_grad]) / 1024**2 / 4 for k, v in model._modules.items()} \n    )\n\n    if cfg.model.checkpoint == True:\n        cfg.model.checkpoint = os.path.join(cfg.output_dir, experiment_name + '_weights.pt')\n    if cfg.model.checkpoint is not None:\n        state_dict = torch.load(cfg.model.checkpoint)\n        if cfg.model.re_init_decoder: # retrain the decoder\n            # only load weights of encoder\n            decoder = model.decoder\n            del model.decoder\n            model.load_state_dict(state_dict, strict=False)\n            model.decoder = decoder\n        try:\n            model.load_state_dict(state_dict)\n        except:\n            # only load weights of encoder\n            decoder = model.decoder\n            del model.decoder\n            model.load_state_dict(state_dict, strict=False)\n            model.decoder = decoder\n    \n    warmup_steps = int(cfg.train.warmup_epochs * len(train_loader))\n\n    if cfg.train.loss_name == 'ce':\n        criterion = SoftCrossEntropy()\n    elif cfg.train.loss_name == 'arb':\n        criterion = SoftARBLoss(info['class_duration_after_resampling'])\n    elif cfg.train.loss_name == 'bce':\n        criterion = nn.BCEWithLogitsLoss()\n    elif cfg.train.loss_name == 'binary-arb':\n        criterion = SoftBinaryARBLoss(info['class_duration_after_resampling'])\n    elif cfg.train.loss_name == 'focal-bce':\n        # criterion = SoftBinaryFocalLoss()\n        criterion = FocalLossBCE()\n    elif cfg.train.loss_name == 'focal-ce':\n        criterion = SoftFocalLoss()\n    elif cfg.train.loss_name == 'recall-ce':\n        criterion = SoftRecallLoss()\n\n    Optimizer = None\n    if cfg.train.optimizer == 'adam':\n        Optimizer = optim.Adam\n    elif cfg.train.optimizer == 'sgd':\n        Optimizer = lambda *args, **kwargs: optim.SGD(*args, nesterov=True, momentum=0.99, **kwargs)\n        \n    if cfg.model.freeze_encoder:\n        optimizer = Optimizer([{'params': model.decoder.parameters(), 'lr': cfg.train.learning_rate}],\n                               lr=cfg.train.learning_rate, weight_decay=cfg.train.weight_decay)\n    else:\n        optimizer = Optimizer(model.parameters(),\n                               lr=cfg.train.learning_rate, weight_decay=cfg.train.weight_decay)\n    \n\n    lr_scheduler = optim.lr_scheduler.SequentialLR(optimizer,\n                                                   [optim.lr_scheduler.LinearLR(optimizer, 1e-3, 1, warmup_steps),\n                                                    optim.lr_scheduler.CosineAnnealingLR(optimizer, cfg.train.epochs * len(train_loader) - warmup_steps, 1e-6),\n                                                    optim.lr_scheduler.ConstantLR(optimizer, 1e-3, int(cfg.train.finetune_epochs * len(train_loader)))],\n                                                   [warmup_steps, cfg.train.epochs * len(train_loader)])\n    \n    if cfg.train.uda:\n        uda_steps = cfg.train.uda_epochs * min(len(train_loader), len(uda_loader))\n        trainer = UDAAudioTrainer(model, processor, transforms_dict['spec'], \n                                  optimizer, criterion, info['class_duration_after_resampling'], multilabel=cfg.model.use_multilabel,\n                                  uda_steps=uda_steps, uda_coef=cfg.train.uda_coef,\n                                  loss_coefs=cfg.train.loss_coefs, lr_scheduler=lr_scheduler)\n        history = trainer.fit(train_loader, uda_loader, val_loader=val_loader, epochs=cfg.train.epochs)\n        # trainer.evaluate(val_loader)\n    else:\n        \n        trainer = AudioTrainer(model, processor, transforms_dict['spec'], \n                               optimizer, criterion, info['class_duration_after_resampling'], loss_coefs=cfg.train.loss_coefs, lr_scheduler=lr_scheduler)\n    \n        history = trainer.fit(train_loader, val_loader=val_loader, epochs=cfg.train.epochs)\n        # trainer.evaluate(val_loader)\n    \n    os.makedirs(cfg.output_dir, exist_ok=True)\n    \n    if isinstance(trainer.model.encoder, nn.DataParallel):\n        trainer.model.encoder = trainer.model.encoder.module\n    if isinstance(trainer.model.decoder, nn.DataParallel):\n        trainer.model.decoder = trainer.model.decoder.module\n            \n    torch.save(trainer.model.state_dict(), os.path.join(cfg.output_dir, experiment_name + '_weights.pt'))\n    with open(os.path.join(cfg.output_dir, experiment_name + f'_history.json'), 'w') as f:\n        json.dump(history, f)\n\n    torch.cuda.empty_cache()\n    return trainer\n\n\ndef train_at_once(cfg_list):\n    for new_cfg in cfg_list:\n        cfg = get_default_config()\n        for k, v in new_cfg:\n            OmegaConf.update(cfg, k, v)\n\n        # OmegaConf.update(cfg, 'model.checkpoint', 'outputs/1__s0.1_ast_fc_ce_12_None_1_weights.pt')\n        OmegaConf.update(cfg, 'model.re_init_decoder', True)\n        # OmegaConf.update(cfg, 'model.freeze_encoder', True)\n        train(cfg)\n\n    ","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.890629Z","iopub.execute_input":"2024-06-08T04:22:42.890992Z","iopub.status.idle":"2024-06-08T04:22:42.921607Z","shell.execute_reply.started":"2024-06-08T04:22:42.890961Z","shell.execute_reply":"2024-06-08T04:22:42.920713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%output_cache\n\n# train_at_once(cfg_list)\n# # test_at_once(cfg_list)","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.923146Z","iopub.execute_input":"2024-06-08T04:22:42.924060Z","iopub.status.idle":"2024-06-08T04:22:42.934584Z","shell.execute_reply.started":"2024-06-08T04:22:42.924023Z","shell.execute_reply":"2024-06-08T04:22:42.933586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ncfg = get_default_config()\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['stage1'])\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['volodymyr'])\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['not_use_cache'])\n\n# OmegaConf.update(cfg, 'model.name', 'convnextv2_tiny')\nOmegaConf.update(cfg, 'model.name', ['efficientnetv2', 'convnextv2_tiny', 'eca_nfnet'])\nOmegaConf.update(cfg, 'model.decoder_type', 'fc')\nOmegaConf.update(cfg, 'dataset.sampling_rate', 32000)\nOmegaConf.update(cfg, 'dataset.audio_second', 10)\n# OmegaConf.update(cfg, 'train.loss_coefs', [0.5, 0.5])\n\nOmegaConf.update(cfg, 'train.epochs', 2)\nOmegaConf.update(cfg, 'train.warmup_epochs', 1)\nOmegaConf.update(cfg, 'seed', 1)\nOmegaConf.update(cfg, 'dataset.batch_size', 32)\nOmegaConf.update(cfg, 'dataset.kfold', None)\n\nOmegaConf.update(cfg, 'model.checkpoint', None)\n# OmegaConf.update(cfg, 'model.checkpoint', 'outputs/1__s0.1_pvt_v2_b2_pvt_v2_b0_pvt_v2_b1_convnext_small_convnextv2_tiny_eca_nfnet_fc_ce_10_None_1_weights.pt')\n\nOmegaConf.update(cfg, 'model.use_multilabel', False)\nOmegaConf.update(cfg, 'train.loss_name', 'ce')\nOmegaConf.update(cfg, 'dataset.info_path', './info_p0.5.json')\nOmegaConf.update(cfg, 'model.re_init_decoder', True)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.935890Z","iopub.execute_input":"2024-06-08T04:22:42.936322Z","iopub.status.idle":"2024-06-08T04:22:42.971064Z","shell.execute_reply.started":"2024-06-08T04:22:42.936290Z","shell.execute_reply":"2024-06-08T04:22:42.970065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%output_cache\n\ntrain(cfg)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:22:42.972043Z","iopub.execute_input":"2024-06-08T04:22:42.972314Z","iopub.status.idle":"2024-06-08T04:26:00.630347Z","shell.execute_reply.started":"2024-06-08T04:22:42.972291Z","shell.execute_reply":"2024-06-08T04:26:00.628863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ncfg = get_default_config()\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['stage2'])\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['volodymyr'])\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['not_use_cache'])\n\nOmegaConf.update(cfg, 'train.epochs', 12)\nOmegaConf.update(cfg, 'train.warmup_epochs', 1)\nOmegaConf.update(cfg, 'dataset.sampling_rate', 32000)\nOmegaConf.update(cfg, 'dataset.audio_second', 5)\n\nOmegaConf.update(cfg, 'model.name', ['efficientnetv2', 'convnextv2_tiny', 'eca_nfnet'])\n# OmegaConf.update(cfg, 'model.name', 'pvt_v2_b2')\nOmegaConf.update(cfg, 'model.decoder_type', 'fc')\n\nOmegaConf.update(cfg, 'seed', 1)\nOmegaConf.update(cfg, 'dataset.kfold', 1)\n\nOmegaConf.update(cfg, 'dataset.batch_size', 32)\nOmegaConf.update(cfg, 'model.checkpoint', 'outputs/1__sq_efficientnetv2_convnextv2_tiny_eca_nfnet_fc_ce_10_1_1_weights.pt')\n# OmegaConf.update(cfg, 'model.checkpoint', None)\n\nOmegaConf.update(cfg, 'model.use_multilabel', False)\nOmegaConf.update(cfg, 'train.loss_name', 'ce')\n\nOmegaConf.update(cfg, 'dataset.info_path', './info2024_sq.json')\n# OmegaConf.update(cfg, 'dataset.info_path', './info2024_s0.1.json')\n\n\n# OmegaConf.update(cfg, 'model.freeze_encoder', True)\nOmegaConf.update(cfg, 'model.re_init_decoder', True)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:03.137397Z","iopub.execute_input":"2024-06-08T04:26:03.138106Z","iopub.status.idle":"2024-06-08T04:26:03.168110Z","shell.execute_reply.started":"2024-06-08T04:26:03.138076Z","shell.execute_reply":"2024-06-08T04:26:03.167209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%output_cache\n# # torch.autograd.set_detect_anomaly(True)\n# train(cfg)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:04.149599Z","iopub.execute_input":"2024-06-08T04:26:04.150325Z","iopub.status.idle":"2024-06-08T04:26:04.153887Z","shell.execute_reply.started":"2024-06-08T04:26:04.150294Z","shell.execute_reply":"2024-06-08T04:26:04.152959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ncfg = get_default_config()\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['stage2'])\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['volodymyr'])\ncfg = OmegaConf.merge(cfg, CONFIG_DICT['not_use_cache'])\n\nOmegaConf.update(cfg, 'train.learning_rate', 1e-4)\nOmegaConf.update(cfg, 'train.epochs', 10)\nOmegaConf.update(cfg, 'train.warmup_epochs', 1)\nOmegaConf.update(cfg, 'dataset.sampling_rate', 32000)\nOmegaConf.update(cfg, 'dataset.audio_second', 5)\n\nOmegaConf.update(cfg, 'model.name', ['efficientnetv2', 'convnextv2_tiny', 'eca_nfnet'])\n# OmegaConf.update(cfg, 'model.name', 'efficientnetv2')\n\nOmegaConf.update(cfg, 'model.decoder_type', 'fc')\n\nOmegaConf.update(cfg, 'seed', 1)\nOmegaConf.update(cfg, 'dataset.kfold', None)\n\nOmegaConf.update(cfg, 'dataset.batch_size', 32)\n# OmegaConf.update(cfg, 'model.checkpoint',\n#                  'outputs/1__sq_pvt_v2_b2_efficientnetv2_convnext_small_convnextv2_tiny_eca_nfnet_fc_ce_10_None_1_5sec_weights.pt')\nOmegaConf.update(cfg, 'model.checkpoint', 'outputs/2_2024_sq_efficientnetv2_convnextv2_tiny_eca_nfnet_fc_ce_8_1_1_weights.pt')\n# OmegaConf.update(cfg, 'model.checkpoint', None)\n\nOmegaConf.update(cfg, 'train.loss_name', 'focal-bce')\nOmegaConf.update(cfg, 'model.use_multilabel', True)\n\nOmegaConf.update(cfg, 'dataset.info_path', './info2024_sq.json')\n# OmegaConf.update(cfg, 'train.sdat', False)\n# OmegaConf.update(cfg, 'train.use_mcc', False)\n# OmegaConf.update(cfg, 'train.use_adv', False)\n# OmegaConf.update(cfg, 'train.use_sam', False)\nOmegaConf.update(cfg, 'train.uda', True)\nOmegaConf.update(cfg, 'train.uda_epochs', 2)\nOmegaConf.update(cfg, 'train.uda_coef', 2)\n# OmegaConf.update(cfg, 'train.multi_experts', True)\n\n\nOmegaConf.update(cfg, 'train.optimizer', 'adam')\n\n# OmegaConf.update(cfg, 'model.freeze_encoder', True)\nOmegaConf.update(cfg, 'model.re_init_decoder', True)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:04.553322Z","iopub.execute_input":"2024-06-08T04:26:04.553682Z","iopub.status.idle":"2024-06-08T04:26:04.587530Z","shell.execute_reply.started":"2024-06-08T04:26:04.553654Z","shell.execute_reply":"2024-06-08T04:26:04.586575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%output_cache\n# # torch.autograd.set_detect_anomaly(True)\n# train(cfg)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:04.838635Z","iopub.execute_input":"2024-06-08T04:26:04.839308Z","iopub.status.idle":"2024-06-08T04:26:04.842774Z","shell.execute_reply.started":"2024-06-08T04:26:04.839277Z","shell.execute_reply":"2024-06-08T04:26:04.841942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test","metadata":{}},{"cell_type":"code","source":"def test(cfg):\n    global trainer, train_ds, val_ds\n    experiment_name = get_experiment_name(cfg)\n    print('Experiment name:', experiment_name)\n    \n    if cfg.seed is not None:\n        seed_everything(cfg.seed)\n\n    train_info_path = cfg.dataset.info_path\n    val_info_path = cfg.dataset.val_info_path\n    cfg.dataset.info_path = val_info_path\n    train_ds, val_ds, train_loader, val_loader, transforms_dict, val_info = get_dataloaders(cfg)\n    info = load_dataset_information(train_info_path)\n    cfg.dataset.info_path = train_info_path\n\n    train_ds.transform = torchvision.transforms.Compose([\n        train_ds.transform.transforms[0],\n        train_ds.transform.transforms[1],\n        train_ds.transform.transforms[2],\n    ])\n    \n    # train_ds = TestBirdCLEFDataset(train_ds, info, val_info)\n    # val_ds = TestBirdCLEFDataset(val_ds, info, val_info)\n\n    def seed_worker(worker_id):\n        worker_seed = torch.initial_seed() % 2**32\n        np.random.seed(worker_seed)\n        random.seed(worker_seed)\n    \n    test_generator = torch.Generator()\n    if cfg.seed is not None:\n        test_generator.manual_seed(cfg.seed)\n\n    val_loader = torch.utils.data.DataLoader(\n            val_ds,\n            sampler=torch.utils.data.SequentialSampler(val_ds),\n            batch_size=cfg.dataset.batch_size,\n            num_workers=cfg.dataset.num_workers,\n            drop_last=False,\n            pin_memory=True,\n            prefetch_factor=cfg.dataset.prefetch_factor,\n            worker_init_fn=seed_worker,\n            generator=test_generator,\n    )\n    \n    processor = get_spec_processor(cfg.dataset.sampling_rate, cfg.dataset.audio_second)[cfg.dataset.augmentation.processor]\n    \n    model = get_model(cfg.model.name, cfg.model.decoder_type, info)\n    print(\n        sum([param.element_size() * param.numel() for param in model.parameters() if param.requires_grad]) / 1024**2 / 4, {k:sum([param.element_size() * param.numel() for param in v.parameters() if param.requires_grad]) / 1024**2 / 4 for k, v in model._modules.items()} \n    )\n\n    if cfg.train.multi_experts:\n        model.cuda()\n        model = Ensemble([deepcopy(model), deepcopy(model), deepcopy(model)])\n        \n    if cfg.model.checkpoint == True:\n        cfg.model.checkpoint = os.path.join(cfg.output_dir, experiment_name + '_weights.pt')\n    if cfg.model.checkpoint is not None:\n        state_dict = torch.load(cfg.model.checkpoint)   \n        model.load_state_dict(state_dict)\n\n    if cfg.train.multi_experts:\n        for m in model:\n            m.decoder = TestBirdCLEFWrapper(m.decoder, info, val_info)\n    else:\n        model.decoder = TestBirdCLEFWrapper(model.decoder, info, val_info)\n    \n    \n        \n    if cfg.train.loss_name == 'ce':\n        criterion = SoftCrossEntropy()\n    elif cfg.train.loss_name == 'arb':\n        criterion = SoftARBLoss(val_info['class_counts_after_resampling'])\n    elif cfg.train.loss_name == 'bce':\n        criterion = nn.BCEWithLogitsLoss()\n    elif cfg.train.loss_name == 'binary-arb':\n        criterion = SoftBinaryARBLoss(val_info['class_counts_after_resampling'])\n\n    optimizer = optim.Adam(model.parameters(), lr=cfg.train.learning_rate, weight_decay=cfg.train.weight_decay)\n\n    if cfg.train.multi_experts:\n        trainer = ExpertsTrainer(model, processor, transforms_dict['spec'], \n                                 optimizer, criterion, val_info['reduced_class_counts'])\n    else:\n        trainer = AudioTrainer(model, processor, transforms_dict['spec'], \n                               optimizer, criterion, val_info['reduced_class_counts'])\n    \n    reture_score = True\n    if reture_score:\n        logits, y = trainer.predict(val_loader, return_fea=False, return_gt=True)\n        y = ((y - y.max(dim=1, keepdim=True)[0]) >= 0).to(torch.float32) # binarize\n        if cfg.model.use_multilabel:\n            pred_y = torch.sigmoid(logits.float())\n        else:\n            pred_y = torch.softmax(logits.float(), dim=1)\n        if len(pred_y.shape) > 2:\n            pred_y = pred_y.mean(dim=list(range(2, len(pred_y.shape))))\n            # u, c = torch.unique(pred_y[i].argmax(dim=1), return_counts=True)\n            # u[torch.argsort(c)[-5:]], torch.argmax(pred_y[i].mean(dim=0)), torch.argmax(y[i])\n            \n        y = torch.cat([y, torch.ones(5, y.shape[1])], dim=0)\n        pred_y = torch.cat([pred_y, torch.ones(5, y.shape[1])], dim=0)\n        score = average_precision_score(\n            y.numpy(),\n            pred_y.numpy(),\n            average='macro',\n        )\n        torch.cuda.empty_cache()\n        return score, y, pred_y\n    else:\n        metrics = trainer.evaluate(val_loader)\n        torch.cuda.empty_cache()\n        return metrics\n\n   \ndef test_at_once(cfg_list, dest='./outputs/results.txt'):\n    # results = []\n    for new_cfg in cfg_list:\n        cfg = get_default_config()\n        for k, v in new_cfg:\n            OmegaConf.update(cfg, k, v)\n        if cfg.train.loss_name in['bce', 'binary_arb']: \n            OmegaConf.update(cfg, 'model.use_multilabel', True)\n        OmegaConf.update(cfg, 'dataset.use_secondary', False)\n        OmegaConf.update(cfg, 'dataset.val_info_path', './info2024_primary.json')\n        experiment_name = get_experiment_name(cfg)\n        OmegaConf.update(cfg, 'model.checkpoint', f'outputs/{experiment_name}_weights.pt')\n        score, _, _ = test(cfg)\n        \n        with open(dest, 'a') as fp:\n            fp.write(f'{experiment_name}: {score}\\n')\n\n\n# split_multi_models(['efficientnetv2', 'convnextv2_tiny', 'eca_nfnet'], 'fc', \n#                     'outputs/2_2024_p2_efficientnetv2_convnextv2_tiny_eca_nfnet_fc_ce_20_None_1_weights.pt', './info2024_sq.json')\n# split_multi_models(['pvt_v2_b2', 'efficientnetv2', 'convnext_small', 'convnextv2_tiny', 'eca_nfnet'], 'fc', \n#                     'outputs/1__sq_pvt_v2_b2_efficientnetv2_convnext_small_convnextv2_tiny_eca_nfnet_fc_ce_10_None_1_10sec_weights.pt', './info_sq.json')\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:06.013455Z","iopub.execute_input":"2024-06-08T04:26:06.013828Z","iopub.status.idle":"2024-06-08T04:26:06.041447Z","shell.execute_reply.started":"2024-06-08T04:26:06.013785Z","shell.execute_reply":"2024-06-08T04:26:06.040422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# with open('./outputs/2_2024_sq_efficientnetv2_convnextv2_tiny_eca_nfnet_fc_ce_12_1_1_history.json', 'r') as fp:\n#     obj = json.load(fp)\n\n# plt.plot(obj['val_loss'])","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:07.073274Z","iopub.execute_input":"2024-06-08T04:26:07.073990Z","iopub.status.idle":"2024-06-08T04:26:07.078173Z","shell.execute_reply.started":"2024-06-08T04:26:07.073957Z","shell.execute_reply":"2024-06-08T04:26:07.077209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# cfg = get_default_config()\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['stage2'])\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['volodymyr'])\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['not_use_cache'])\n\n# # OmegaConf.update(cfg, 'model.name', 'convnextv2_tiny')\n# OmegaConf.update(cfg, 'model.name', 'convnextv2_tiny')\n# OmegaConf.update(cfg, 'model.decoder_type', 'fc')\n# OmegaConf.update(cfg, 'dataset.sampling_rate', 32000)\n# OmegaConf.update(cfg, 'dataset.audio_second', 5)\n\n# OmegaConf.update(cfg, 'seed', 1)\n# OmegaConf.update(cfg, 'dataset.batch_size', 32)\n# OmegaConf.update(cfg, 'dataset.kfold', 0)\n# OmegaConf.update(cfg, 'model.checkpoint', 'outputs/2_2024_p2_efficientnetv2_convnextv2_tiny_eca_nfnet_fc_ce_20_None_1_weights.pt')\n# OmegaConf.update(cfg, 'model.use_multilabel', False)\n# OmegaConf.update(cfg, 'dataset.use_secondary', False)\n# # OmegaConf.update(cfg, 'train.multi_experts', True)\n# # OmegaConf.update(cfg, 'dataset.samples_per_epoch', 32 * 150)\n# OmegaConf.update(cfg, 'dataset.resample_val', True)\n\n\n# OmegaConf.update(cfg, 'dataset.val_info_path', './info2024_primary_p1.json')\n# # OmegaConf.update(cfg, 'dataset.val_info_path', './info2024_primary_binary.json')\n\n# # OmegaConf.update(cfg, 'dataset.info_path', './info.json')\n# OmegaConf.update(cfg, 'dataset.info_path', './info2024.json')\n# # OmegaConf.update(cfg, 'dataset.info_path', './info2024_binary.json')\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:07.273154Z","iopub.execute_input":"2024-06-08T04:26:07.273770Z","iopub.status.idle":"2024-06-08T04:26:07.278516Z","shell.execute_reply.started":"2024-06-08T04:26:07.273737Z","shell.execute_reply":"2024-06-08T04:26:07.277640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# score, y, pred_y = test(cfg)\n# print(score)\n# # metrics = test(cfg)\n\n# # plt.plot((trainer.metrics['wacc'].sum_value / trainer.metrics['wacc'].count).sort()[0])\n# acc = (pred_y * y).sum(dim=0)  / y.sum(dim=0)\n# plt.plot(pred_y[:, np.argsort(acc)].argmax(dim=1).unique(return_counts=True)[1])","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:07.481431Z","iopub.execute_input":"2024-06-08T04:26:07.482124Z","iopub.status.idle":"2024-06-08T04:26:07.486070Z","shell.execute_reply.started":"2024-06-08T04:26:07.482091Z","shell.execute_reply":"2024-06-08T04:26:07.485052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Analysis","metadata":{}},{"cell_type":"markdown","source":"## Pred feat","metadata":{}},{"cell_type":"code","source":"# def predict(self, data_loader, verbose=True, num_iterations=None, **kwargs):\n#     ### model setting ###\n#     if self.use_cuda:\n#         if isinstance(self.processor, nn.Module):\n#             self.processor.cuda()\n#         self.spec_augmentations.cuda()\n#         self.model.cuda()\n#         self.avg_model.cuda()\n#     self.model.eval()\n#     self.avg_model.eval()\n#     #####################\n#     if num_iterations is None:\n#         num_iterations = len(data_loader)\n#     data_iter = iter(data_loader)\n#     if verbose:\n#         data_iter = tqdm(data_iter, total=num_iterations)\n\n#     results = []\n#     for step, data in enumerate(data_iter):\n#         if step >= num_iterations:\n#             break\n#         outputs = predict_step(self, data, **kwargs)\n#         results.append(outputs)\n\n#     if isinstance(outputs, (tuple, list)):  # multi outputs\n#         results = list(zip(*results))\n#         results = [torch.cat(tensor, dim=0) for tensor in results]\n#     else:\n#         results = torch.cat(results, dim=0)\n#     if len(results) == 1:\n#         results = results[0]\n#     return results\n    \n# def predict_step(self, data, x_idx=0, y_idx=1, return_pred=True, return_fea=False, return_gt=False):\n#     ### config input data ###\n#     data = self._convert_cuda_data(data)\n#     x = data[x_idx]\n#     x = self.processor(x)\n#     #########################\n\n#     with self.autocast():\n#         with torch.no_grad():\n#             ### forward pass ###\n#             fea = self.model.encoder(x)\n#             logits = self.model.decoder(fea)\n#             ####################\n#     if isinstance(logits, (list, tuple)):\n#         logits = logits[0]\n        \n#     outputs = []\n#     if return_pred:\n#         outputs.append(logits.detach().cpu())\n#     if return_fea:\n#         outputs.append(fea.detach().cpu())\n#     if return_gt:\n#         outputs.append(data[y_idx].detach().cpu())\n#     return outputs\n\n# cfg = get_default_config()\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['stage2'])\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['volodymyr'])\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['not_use_cache'])\n# OmegaConf.update(cfg, 'model.name', 'convnextv2_tiny')\n# OmegaConf.update(cfg, 'model.decoder_type', 'fc')\n# OmegaConf.update(cfg, 'seed', 1)\n# OmegaConf.update(cfg, 'dataset.batch_size', 8)\n# OmegaConf.update(cfg, 'dataset.kfold', 0)\n# OmegaConf.update(cfg, 'dataset.sampling_rate', 32000)\n# OmegaConf.update(cfg, 'dataset.audio_second', 5)\n# OmegaConf.update(cfg, 'model.use_multilabel', False)\n# OmegaConf.update(cfg, 'model.checkpoint', 'outputs/2_2024_sq_convnextv2_tiny_fc_ce_20_None_1_weights.pt')\n# OmegaConf.update(cfg, 'dataset.val_info_path', './info2024_primary.json')\n# OmegaConf.update(cfg, 'dataset.info_path', './info2024.json')\n\n# train_info_path = cfg.dataset.info_path\n# val_info_path = cfg.dataset.val_info_path\n# cfg.dataset.info_path = val_info_path\n# train_ds, val_ds, train_loader, val_loader, transforms_dict, val_info = get_dataloaders(cfg)\n# info = load_dataset_information(train_info_path)\n# cfg.dataset.info_path = train_info_path\n\n# uda_ds, uda_loader, uda_transforms_dict = get_uda_dataloaders(cfg)\n\n# processor = get_spec_processor(cfg.dataset.sampling_rate, cfg.dataset.audio_second)[cfg.dataset.augmentation.processor]\n\n# model = get_model(cfg.model.name, cfg.model.decoder_type, info)\n\n\n# state_dict = torch.load(cfg.model.checkpoint)\n# model.load_state_dict(state_dict)\n\n\n# criterion = SoftCrossEntropy()\n# optimizer = optim.Adam(model.parameters(), lr=cfg.train.learning_rate, weight_decay=cfg.train.weight_decay)\n\n# trainer = AudioTrainer(model, processor, transforms_dict['spec'], \n#                        optimizer, criterion, info['reduced_class_counts'], loss_coefs=cfg.train.loss_coefs)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:08.764825Z","iopub.execute_input":"2024-06-08T04:26:08.765760Z","iopub.status.idle":"2024-06-08T04:26:08.773306Z","shell.execute_reply.started":"2024-06-08T04:26:08.765715Z","shell.execute_reply":"2024-06-08T04:26:08.772322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# reduced_to_base_classes = TestBirdCLEFWrapper(model.decoder, info, load_dataset_information(cfg.dataset.val_info_path)).reduced_to_base_classes\n# weight = reduced_to_base_classes @ model.decoder.weight.detach().cpu()\n# bias = reduced_to_base_classes @ model.decoder.bias.detach().cpu()\n\n# fea_tra, gt_tra = predict(trainer, train_loader, num_iterations=800, x_idx=0, return_pred=False, return_fea=True, return_gt=True)\n# fea_sup, gt_sup = predict(trainer, val_loader, num_iterations=800, x_idx=0, return_pred=False, return_fea=True, return_gt=True)\n# fea_uda = predict(trainer, uda_loader, num_iterations=400, x_idx=0, return_pred=False, return_fea=True, return_gt=False)\n# fea_center = weight + bias[:, None]\n\n# # gt_sup = ((gt_sup - gt_sup.max(dim=1, keepdim=True)[0]) >= 0).to(torch.float32)\n# # gt_tra = ((gt_tra - gt_tra.max(dim=1, keepdim=True)[0]) >= 0).to(torch.float32)\n\n# # TSNE\n# fea = torch.cat([fea_sup, fea_uda, fea_tra, fea_center], dim=0)\n\n# fea2 = TSNE(perplexity=50, metric='cosine').fit_transform(fea)\n\n# fea2_sup, fea2_uda, fea2_tra, fea2_center = torch.from_numpy(fea2).split([fea_sup.shape[0], fea_uda.shape[0], fea_tra.shape[0], fea_center.shape[0]], dim=0)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:08.954673Z","iopub.execute_input":"2024-06-08T04:26:08.955397Z","iopub.status.idle":"2024-06-08T04:26:08.959691Z","shell.execute_reply.started":"2024-06-08T04:26:08.955369Z","shell.execute_reply":"2024-06-08T04:26:08.958828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# def random_scatter(points, labels):\n#     from matplotlib import colormaps\n#     hsv = colormaps.get_cmap('hsv')\n\n#     indices = np.arange(len(points))\n#     np.random.shuffle(indices)\n#     points = points[indices]\n#     labels = labels[indices]\n#     labels = np.array(labels).astype(np.int64)\n    \n#     all_cls = np.unique(labels)\n#     for i, c in enumerate(all_cls):\n#         labels[labels == c] == i\n        \n#     colors = np.array([hsv(int(i) % 256) for i in np.linspace(32, 256+32, len(all_cls), False)])\n#     colors[:, :-1] *= (1 + 1) / (colors[:, :-1].sum(axis=1, keepdims=True) + 1)\n    \n#     plt.scatter(points[:, 0], points[:, 1], c=colors[labels], s=2)\n#     plt.legend(handles=[plt.Line2D([],[],color=c, ls=\"\",marker=\"o\") for c in colors],\n#                         labels=all_cls.tolist())\n\n\n# random_scatter(fea2[:-fea_center.shape[0]], torch.cat([torch.zeros(fea_sup.shape[0]), torch.ones(fea_uda.shape[0]), 2*torch.ones(fea_tra.shape[0])], dim=0))\n\n# plt.scatter(fea2[-fea_center.shape[0]:, 0], fea2[-fea_center.shape[0]:, 1], c='k', s=5, marker='s')\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:09.617990Z","iopub.execute_input":"2024-06-08T04:26:09.618710Z","iopub.status.idle":"2024-06-08T04:26:09.623175Z","shell.execute_reply.started":"2024-06-08T04:26:09.618677Z","shell.execute_reply":"2024-06-08T04:26:09.622245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Ensemble","metadata":{}},{"cell_type":"code","source":"# def predict(self, data_loader, verbose=True, num_iterations=None, epochs=1, **kwargs):\n#     ### model setting ###\n#     if self.use_cuda:\n#         if isinstance(self.processor, nn.Module):\n#             self.processor.cuda()\n#         self.spec_augmentations.cuda()\n#         self.model.cuda()\n#         self.avg_model.cuda()\n#     self.model.eval()\n#     self.avg_model.eval()\n#     #####################\n#     if num_iterations is None:\n#         num_iterations = len(data_loader)\n#     data_iter = iter(data_loader)\n#     if verbose:\n#         data_iter = tqdm(data_iter, total=num_iterations)\n\n#     results = []\n#     for epoch in range(epochs):\n#         for step, data in enumerate(data_iter):\n#             if step >= num_iterations:\n#                 break\n#             outputs = predict_step(self, data, **kwargs)\n#             results.append(outputs)\n\n#     if isinstance(outputs, (tuple, list)):  # multi outputs\n#         results = list(zip(*results))\n#         results = [torch.cat(tensor, dim=0) for tensor in results]\n#     else:\n#         results = torch.cat(results, dim=0)\n#     if len(results) == 1:\n#         results = results[0]\n#     return results\n    \n# def predict_step(self, data, x_idx=0, y_idx=1, return_pred=True, return_fea=False, return_gt=False):\n#     ### config input data ###\n#     data = self._convert_cuda_data(data)\n#     x = data[x_idx]\n#     x = self.processor(x)\n#     #########################\n\n#     with self.autocast():\n#         with torch.no_grad():\n#             ### forward pass ###\n#             logits = self.model(x)\n#             ####################\n#     if isinstance(logits, (list, tuple)):\n#         logits = logits[0]\n        \n#     outputs = []\n#     if return_pred:\n#         outputs.append(logits.detach().cpu())\n#     if return_fea:\n#         outputs.append(fea.detach().cpu())\n#     if return_gt:\n#         outputs.append(data[y_idx].detach().cpu())\n#     return outputs\n\n# cfg = get_default_config()\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['stage2'])\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['volodymyr'])\n# cfg = OmegaConf.merge(cfg, CONFIG_DICT['not_use_cache'])\n# OmegaConf.update(cfg, 'model.name', 'convnextv2_tiny')\n# OmegaConf.update(cfg, 'model.decoder_type', 'fc')\n# OmegaConf.update(cfg, 'seed', 1)\n# OmegaConf.update(cfg, 'dataset.batch_size', 32)\n# OmegaConf.update(cfg, 'dataset.kfold', 0)\n# OmegaConf.update(cfg, 'dataset.sampling_rate', 32000)\n# OmegaConf.update(cfg, 'dataset.audio_second', 5)\n# OmegaConf.update(cfg, 'model.use_multilabel', False)\n# OmegaConf.update(cfg, 'model.checkpoint', 'outputs/2_2024_sq_convnextv2_tiny_fc_ce_20_None_1_weights.pt')\n# OmegaConf.update(cfg, 'dataset.val_info_path', './info2024_primary.json')\n# OmegaConf.update(cfg, 'dataset.info_path', './info2024.json')\n\n# model_names = ['efficientnetv2', 'efficientnetv2', 'efficientnetv2']\n# ckpts = ['outputs/2_2024_p0_efficientnetv2_fc_ce_20_None_1_weights.pt', \n#          'outputs/2_2024_p1_efficientnetv2_fc_ce_20_None_1_weights.pt', \n#          'outputs/2_2024_p2_efficientnetv2_fc_ce_20_None_1_weights.pt']\n\n# train_info_path = './info2024.json'\n# val_info_path = './info2024_primary_p1.json'\n# cfg.dataset.info_path = val_info_path\n# train_ds, val_ds, train_loader, val_loader, transforms_dict, val_info = get_dataloaders(cfg)\n# info = load_dataset_information(train_info_path)\n\n# train_ds.transform = torchvision.transforms.Compose([\n#     train_ds.transform.transforms[0],\n#     train_ds.transform.transforms[1],\n#     train_ds.transform.transforms[2],\n# ])\n\n# uda_ds, uda_loader, uda_transforms_dict = get_uda_dataloaders(cfg)\n\n# processor = get_spec_processor(cfg.dataset.sampling_rate, cfg.dataset.audio_second)[cfg.dataset.augmentation.processor]\n\n# models = []\n# trainers = []\n# for m, c in zip(model_names, ckpts):\n#     cfg.dataset.info_path = train_info_path\n#     model = get_model(m, cfg.model.decoder_type, info)\n    \n#     state_dict = torch.load(c)\n#     model.load_state_dict(state_dict)\n#     model.decoder = TestBirdCLEFWrapper(model.decoder, info, val_info)\n#     model.eval()\n#     optimizer = optim.Adam(model.parameters(), lr=cfg.train.learning_rate, weight_decay=cfg.train.weight_decay)\n    \n#     trainer = AudioTrainer(model, processor, transforms_dict['spec'], \n#                            optimizer, SoftCrossEntropy(), info['reduced_class_counts'], loss_coefs=cfg.train.loss_coefs)\n#     models.append(model)\n#     trainers.append(trainer)\n\n# reduced_to_base_classes = TestBirdCLEFWrapper(None, info, load_dataset_information(cfg.dataset.val_info_path)).reduced_to_base_classes\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:10.426806Z","iopub.execute_input":"2024-06-08T04:26:10.427629Z","iopub.status.idle":"2024-06-08T04:26:10.435432Z","shell.execute_reply.started":"2024-06-08T04:26:10.427593Z","shell.execute_reply":"2024-06-08T04:26:10.434557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# predictions = []\n# groundtruths = []\n# for trainer in trainers:\n#     pred_tra, gt_tra = predict(trainer, train_loader, num_iterations=1600, x_idx=0, return_pred=True, return_fea=False, return_gt=True, epochs=4)\n#     predictions.append(pred_tra)\n#     groundtruths.append(gt_tra)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:10.637681Z","iopub.execute_input":"2024-06-08T04:26:10.638140Z","iopub.status.idle":"2024-06-08T04:26:10.642176Z","shell.execute_reply.started":"2024-06-08T04:26:10.638108Z","shell.execute_reply":"2024-06-08T04:26:10.641303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# acc_list = []\n# for pred, gt in zip(predictions, groundtruths):\n    \n#     acc = torch.sum((pred.float() @ reduced_to_base_classes.t()).softmax(dim=1) * gt.float(), dim=0)\n#     acc /= torch.sum(gt.float(), dim=0)\n#     acc_list.append(acc)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:11.080895Z","iopub.execute_input":"2024-06-08T04:26:11.081522Z","iopub.status.idle":"2024-06-08T04:26:11.085638Z","shell.execute_reply.started":"2024-06-08T04:26:11.081489Z","shell.execute_reply":"2024-06-08T04:26:11.084718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cls_num = reduced_to_base_classes @ info['reduced_class_duration']\n# indices = torch.argsort(cls_num)\n# for acc in acc_list:\n#     plt.plot(acc[indices])\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:11.516834Z","iopub.execute_input":"2024-06-08T04:26:11.517662Z","iopub.status.idle":"2024-06-08T04:26:11.522346Z","shell.execute_reply.started":"2024-06-08T04:26:11.517631Z","shell.execute_reply":"2024-06-08T04:26:11.521426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cls_weights = torch.stack(acc_list, dim=0)\n# cls_weights = cls_weights.softmax(dim=0)\n# for w in cls_weights:\n#     plt.plot(w[indices])\n# acc_list = [acc.tolist() for acc in acc_list]\n# with open('outputs/2_2024_efficientnetv2_fc_ce_20_None_1_ens.json', 'w') as fp:\n#     json.dump(acc_list, fp)","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:11.703422Z","iopub.execute_input":"2024-06-08T04:26:11.704096Z","iopub.status.idle":"2024-06-08T04:26:11.709080Z","shell.execute_reply.started":"2024-06-08T04:26:11.704062Z","shell.execute_reply":"2024-06-08T04:26:11.708198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Extract patches","metadata":{}},{"cell_type":"code","source":"def make_patch_dataset(dataset, dest_dir, dest_df_path, sr=32000, size=10, stride=5):\n    dataset_roots = dataset.roots\n    if dataset_roots is not None:\n        dataset_roots = [root[0] if isinstance(root, (list, tuple)) else root for root in dataset_roots]\n        \n    if dataset_roots is None:\n        cls_dirs = [os.path.join(dataset.root, d) for d in os.listdir(dataset.root)]\n    else:\n        cls_dirs = []\n        for root in dataset_roots:\n            cls_dirs += [os.path.join(root, d) for d in os.listdir(root)]\n        cls_dirs = list(set(cls_dirs))\n\n    os.makedirs(dest_dir, exist_ok=True)\n    for cls_d in cls_dirs:\n        cls_d = os.path.split(cls_d)[1]\n        cls_d = os.path.join(dest_dir, cls_d)\n        os.makedirs(cls_d, exist_ok=True)\n\n    loader = torch.utils.data.DataLoader(\n            dataset,\n            sampler=torch.utils.data.SequentialSampler(dataset),\n            batch_size=1,\n            num_workers=1,\n            drop_last=False,\n            prefetch_factor=4,\n            collate_fn=lambda data: list(zip(*data)),\n    )\n    new_df = []\n    data_iter = iter(loader)\n    for i, (audio, _) in tqdm(enumerate(data_iter), total=len(loader)):\n        audio = audio[0]\n        row = dataset.meta_df.iloc[i]\n        path = dataset._get_path(row)\n        fn, primary_label, secondary_labels, rating = row[['filename', \n                                                           dataset.primary_col_name,\n                                                           'secondary_labels',\n                                                           'rating',]]\n        if dataset_roots is None:\n            dest_path = path.replace(dataset.root, dest_dir)\n        else:\n            dest_path = path\n            for root in dataset_roots:\n                dest_path = dest_path.replace(root, dest_dir)\n\n        if audio.shape[-1] <= sr * size:\n            torchaudio.save(dest_path, patch, sr)\n            new_df.append((fn, primary_label, secondary_labels, rating))\n            continue\n            # print(fn)\n            # print(dest_path)\n            \n        patches, remains = extract_patches(audio.unsqueeze(-1), (size*sr, 1), (stride*sr, 1))\n        patches = patches[0, ..., 0]\n        for i in range(len(patches)):\n            patch = patches[i]\n            sec = min(size + stride * i, math.ceil(audio.shape[-1] / sr))\n            torchaudio.save(os.path.splitext(dest_path)[0] + '_' + str(sec) + '.ogg', patch, sr)\n\n            ori_path, new_fn = os.path.split(fn)\n            row_id, ext = os.path.splitext(new_fn)\n            new_fn = os.path.join(ori_path, row_id + '_' + str(sec) + '.ogg')\n            new_df.append((new_fn, primary_label, secondary_labels, rating))\n            # print(new_fn)\n            # print(os.path.splitext(dest_path)[0] + '_' + str(sec) + '.ogg')\n        \n    new_df = pd.DataFrame(new_df, columns=['filename', dataset.primary_col_name, 'secondary_labels', 'rating'])\n    new_df.to_csv(dest_df_path, index=False)\n        \n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-08T04:26:14.317215Z","iopub.execute_input":"2024-06-08T04:26:14.317822Z","iopub.status.idle":"2024-06-08T04:26:14.334630Z","shell.execute_reply.started":"2024-06-08T04:26:14.317792Z","shell.execute_reply":"2024-06-08T04:26:14.333704Z"},"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":[]}]}