{"metadata":{"kaggle":{"accelerator":"none","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":8773627,"sourceType":"datasetVersion","datasetId":5037230},{"sourceId":185199807,"sourceType":"kernelVersion"}],"dockerImageVersionId":30698,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false},"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# !wget https://raw.githubusercontent.com/LIHANG-HONG/birdclef2023-2nd-place-solution/main/modules/augmentations.py\n\n# !wget https://files.pythonhosted.org/packages/12/b9/ee24c41fb9fd12c72bf7175703bbcde4483a5f93fab36ca2c6022206809c/onnxruntime-1.18.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl\n# !pip install onnxruntime-1.18.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl --no-deps\n\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install omegaconf\n# !mkdir asset\n# !cp -r /kaggle/input/birdclef-2024-asset/* asset\n# !mkdir omegaconf\n# !cp -r /opt/conda/lib/python3.10/site-packages/omegaconf/* omegaconf\n# !mkdir antlr4\n# !cp -r /opt/conda/lib/python3.10/site-packages/antlr4/* antlr4\n# !mkdir ast_processor\n# !wget https://huggingface.co/MIT/ast-finetuned-audioset-10-10-0.4593/resolve/main/config.json?download=true -O ast_processor/config.json\n# !wget https://huggingface.co/MIT/ast-finetuned-audioset-10-10-0.4593/resolve/main/preprocessor_config.json?download=true -O ast_processor/preprocessor_config.json\n\n# !rm asset/2_2024_*\n# !cp /kaggle/input/birdclef-2024-asset/2_2024* ./asset\n\n# !ls asset ","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:03.252705Z","iopub.execute_input":"2024-06-24T11:55:03.253831Z","iopub.status.idle":"2024-06-24T11:55:12.387901Z","shell.execute_reply.started":"2024-06-24T11:55:03.253790Z","shell.execute_reply":"2024-06-24T11:55:12.385836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nif os.path.exists('/kaggle/input/birdclef-2024-inference-notebook') and os.path.isdir('/kaggle/input/birdclef-2024-inference-notebook'):\n    !cp -r /kaggle/input/birdclef-2024-inference-notebook/* ./\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:49:10.101775Z","iopub.execute_input":"2024-06-24T11:49:10.102603Z","iopub.status.idle":"2024-06-24T11:49:24.698721Z","shell.execute_reply.started":"2024-06-24T11:49:10.102560Z","shell.execute_reply":"2024-06-24T11:49:24.696998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install onnxruntime-1.18.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl --no-deps\n!pip install openvino-2024.1.0-15008-cp310-cp310-manylinux2014_x86_64.whl --no-deps\n!pip install openvino_telemetry-2024.1.0-py3-none-any.whl --no-deps\n!pip install openvino_dev-2024.1.0-15008-py3-none-any.whl --no-deps\n\n\n!export OMP_NUM_THREADS=N\n!export OMP_SCHEDULE=STATIC\n!export OMP_PROC_BIND=CLOSE\n!export GOMP_CPU_AFFINITY=\"N-M\"","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:49:24.701541Z","iopub.execute_input":"2024-06-24T11:49:24.701964Z","iopub.status.idle":"2024-06-24T11:51:03.539168Z","shell.execute_reply.started":"2024-06-24T11:49:24.701894Z","shell.execute_reply":"2024-06-24T11:51:03.537433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport os\n# os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':16:8'\n# os.environ['PYTHONHASHSEED'] = '123'\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\n# import multiprocessing\nfrom collections import deque\nimport io\nimport itertools\nfrom joblib.externals.loky.backend.context import get_context\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, ASTConfig\n# from transformers import ASTConfig, ASTModel\n\n\n# from augmentations import (\n#     CustomCompose,\n#     CustomOneOf,\n#     NoiseInjection,\n#     GaussianNoise,\n#     PinkNoise,\n#     AddGaussianNoise,\n#     AddGaussianSNR,\n# )\n# from audiomentations import Compose as amCompose\n# from audiomentations import OneOf as amOneOf\n# from audiomentations import AddBackgroundNoise, Gain, GainTransition, TimeStretch\n# from torch_audiomentations import Compose, PitchShift, Shift\n\n# from blocks import AttHead\n\n# from notebook_cache_outputs import output_cache\n\ntorch.set_flush_denormal(True)\ntorch.set_num_threads(torch.get_num_threads())\ntorchaudio.backend.set_audio_backend('soundfile')\n\nSAMPLING_RATE = 32000\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:51:03.541296Z","iopub.execute_input":"2024-06-24T11:51:03.541806Z","iopub.status.idle":"2024-06-24T11:51:16.668676Z","shell.execute_reply.started":"2024-06-24T11:51:03.541767Z","shell.execute_reply":"2024-06-24T11:51:16.667468Z"},"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","metadata":{"execution":{"iopub.status.busy":"2024-06-09T07:04:29.344902Z","iopub.execute_input":"2024-06-09T07:04:29.346072Z","iopub.status.idle":"2024-06-09T07:04:29.351554Z","shell.execute_reply.started":"2024-06-09T07:04:29.346006Z","shell.execute_reply":"2024-06-09T07:04:29.350340Z"},"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":"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 + 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        # 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 + 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        # 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 CacheDatasetWrapper(torch.utils.data.Dataset):\n    def __init__(self, dataset, max_items=None):\n        super().__init__()\n        self.dataset = dataset\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 __getitem__(self, index):\n        self.num_query.set(self.num_query.get() + 1)\n        if index in self.cached:\n            self.num_hit.set(self.num_hit.get() + 1)\n            return self.cached[index]\n        else:\n            items = self.dataset[index]\n            self.cached[index] = items\n            if self.max_items is not None:\n                self.queue.put(index)\n                while self.queue.qsize() > self.max_items:\n                    try:\n                        self.cached.pop(self.queue.get())\n                    except KeyError as err:\n                        print(err)\n            return items\n\n    def __len__(self):\n        return len(self.dataset)\n\n    def __del__(self):\n        self.manager.shutdown.cancel()\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\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-24T11:55:12.392480Z","iopub.execute_input":"2024-06-24T11:55:12.393103Z","iopub.status.idle":"2024-06-24T11:55:12.472848Z","shell.execute_reply.started":"2024-06-24T11:55:12.393052Z","shell.execute_reply":"2024-06-24T11:55:12.471361Z"},"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    resampling_weights = 1 / (reduced_duration.double() + smooth_factor * reduced_duration.double().max())\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\n# info = load_dataset_information('info.json')\n# info2024 = load_dataset_information('info2024.json')\n# taxonomy_maps = read_taxonomy_information(os.path.join('./Datasets/birdclef', 'birdclef-2024/eBird_Taxonomy_v2021.csv'))\n# ds20_list = load_additional_birdclef_datasets(taxonomy_maps, root='./Datasets/birdclef')\n# ds_list = [load_birdclef_dataset(year, taxonomy_maps, root='./Datasets/birdclef') for year in [2021, 2022, 2023, 2024]]\n\n\n# for names, (use_secondary, smooth_factor, use_multilabel) in zip(itertools.product(['', '_primary'], ['', '_s0.1', '_s1', '_uni'], ['_binary', '']), \n#                      itertools.product([True, False], [0, 0.1, 1, 1e7], [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-24T11:55:12.474989Z","iopub.execute_input":"2024-06-24T11:55:12.475488Z","iopub.status.idle":"2024-06-24T11:55:12.516526Z","shell.execute_reply.started":"2024-06-24T11:55:12.475441Z","shell.execute_reply":"2024-06-24T11:55:12.515016Z"},"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)\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\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    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","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:12.520167Z","iopub.execute_input":"2024-06-24T11:55:12.520659Z","iopub.status.idle":"2024-06-24T11:55:12.575941Z","shell.execute_reply.started":"2024-06-24T11:55:12.520617Z","shell.execute_reply":"2024-06-24T11:55:12.574239Z"},"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\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    \n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:12.577777Z","iopub.execute_input":"2024-06-24T11:55:12.578260Z","iopub.status.idle":"2024-06-24T11:55:12.606788Z","shell.execute_reply.started":"2024-06-24T11:55:12.578228Z","shell.execute_reply":"2024-06-24T11:55:12.605115Z"},"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, image_size=256):\n    processor_volodymyr = nn.Sequential(\n        torchaudio.transforms.MelSpectrogram(sample_rate=sampling_rate, \n                                             hop_length=audio_second*sampling_rate // (image_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-24T11:55:12.608737Z","iopub.execute_input":"2024-06-24T11:55:12.609314Z","iopub.status.idle":"2024-06-24T11:55:12.639382Z","shell.execute_reply.started":"2024-06-24T11:55:12.609279Z","shell.execute_reply":"2024-06-24T11:55:12.638132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class 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\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 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\nclass 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, x):\n        logits = self.module(x)\n        logits = torch.tensordot(logits, self.reduced_to_base_classes, dims=[[1], [1]])\n        return logits\n\n## only for nn.Linear\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 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    \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        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=False,\n            in_chans=1,\n        )\n        if decoder_type == 'fc':\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 = 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        \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-24T11:55:12.641308Z","iopub.execute_input":"2024-06-24T11:55:12.641805Z","iopub.status.idle":"2024-06-24T11:55:12.722399Z","shell.execute_reply.started":"2024-06-24T11:55:12.641771Z","shell.execute_reply":"2024-06-24T11:55:12.721212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss and Metric","metadata":{}},{"cell_type":"code","source":"class 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        i = torch.argmax(loss.abs())\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 SelectedSoftAccuracy(MeanMetric):\n    def __init__(self, indices):\n        if not isinstance(indices, (list, tuple)):\n            indices = [indices]\n        super().__init__()\n        self.indices = indices\n        self.sum_value = torch.zeros(len(self.indices))\n        self.count = torch.zeros(len(self.indices))\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[:, self.indices].sum(dim=0).cpu()\n        self.count += targets[:, self.indices].sum(dim=0).cpu()\n        \n    def compute(self):   \n        acc = self.sum_value / self.count\n        if len(acc) == 1:\n            acc = acc[0]\n        return acc\n    \n    def reset(self):\n        self.sum_value = torch.zeros(len(self.indices))\n        self.count = torch.zeros(len(self.indices))\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-24T11:55:12.723930Z","iopub.execute_input":"2024-06-24T11:55:12.724363Z","iopub.status.idle":"2024-06-24T11:55:12.774006Z","shell.execute_reply.started":"2024-06-24T11:55:12.724333Z","shell.execute_reply":"2024-06-24T11:55:12.771557Z"},"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":"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    y1, y2 = y[indices1], y[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        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    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        ### 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        ### 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            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-24T11:55:12.776268Z","iopub.execute_input":"2024-06-24T11:55:12.776710Z","iopub.status.idle":"2024-06-24T11:55:12.891651Z","shell.execute_reply.started":"2024-06-24T11:55:12.776663Z","shell.execute_reply":"2024-06-24T11:55:12.890357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## UDA Trainer","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"def 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.image_size = 256\n    cfg.dataset.use_secondary = True\n    cfg.dataset.root = '/kaggle/input/birdclef-2024'\n    cfg.dataset.roots = [('/kaggle/input/birdclef-2024', 0.5), ('/kaggle/input/birdclef-2024', 0.5)]\n    cfg.dataset.info_path = './info.json'\n    cfg.dataset.val_info_path = './info2024_primary.json'\n    cfg.dataset.kfold = 0\n    cfg.dataset.samples_per_epoch = None\n    cfg.dataset.batch_size = 8\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    \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-24T11:55:12.895821Z","iopub.execute_input":"2024-06-24T11:55:12.896269Z","iopub.status.idle":"2024-06-24T11:55:12.912740Z","shell.execute_reply.started":"2024-06-24T11:55:12.896238Z","shell.execute_reply":"2024-06-24T11:55:12.910923Z"},"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-24T11:55:12.914611Z","iopub.execute_input":"2024-06-24T11:55:12.915125Z","iopub.status.idle":"2024-06-24T11:55:12.939042Z","shell.execute_reply.started":"2024-06-24T11:55:12.915080Z","shell.execute_reply":"2024-06-24T11:55:12.937711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"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    experiment_name = f'{cfg.dataset.stage}_{_info_name}_{cfg.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    processor = get_spec_processor(cfg.dataset.sampling_rate, cfg.dataset.audio_second, cfg.dataset.image_size)[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_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(info['class_counts_after_resampling'])\n\n    if cfg.model.freeze_encoder:\n        optimizer = optim.Adam([{'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 = optim.Adam([{'params': model.encoder.parameters()}, \n                                {'params': model.decoder.parameters(), 'lr': cfg.train.learning_rate}],\n                               lr=cfg.train.learning_rate, weight_decay=cfg.train.weight_decay)\n        \n    # optimizer = optim.Adam(model.parameters(), lr=learning_rate, weight_decay=1e-6)\n    if cfg.model.optimizer_checkpoint is not None:\n        opt_state_dict = torch.load(cfg.model.optimizer_checkpoint)\n        opt_state_dict.pop('param_groups')\n        opt_state_dict['param_groups'] = optimizer.state_dict()['param_groups']\n        for i in opt_state_dict['state'].keys(): # reset step\n            opt_state_dict['state'][i]['step'] *= 0\n        optimizer.load_state_dict(opt_state_dict)\n        # change to cuda\n        for d in optimizer.state.values():\n            d['exp_avg'] = d['exp_avg'].cuda()\n            d['exp_avg_sq'] = d['exp_avg_sq'].cuda()\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    trainer = AudioTrainer(model, processor, transforms_dict['spec'], \n                           optimizer, criterion, info['reduced_class_counts'], 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    \n    \n    os.makedirs(cfg.output_dir, exist_ok=True)\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\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        train(cfg)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:12.941238Z","iopub.execute_input":"2024-06-24T11:55:12.941726Z","iopub.status.idle":"2024-06-24T11:55:12.973242Z","shell.execute_reply.started":"2024-06-24T11:55:12.941678Z","shell.execute_reply":"2024-06-24T11:55:12.972026Z"},"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 = 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.image_size)[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        model.load_state_dict(state_dict)\n        \n#     model.decoder = TestBirdCLEFWrapper(model.decoder, info, val_info)\n    model.decoder = TestBirdCLEFWrapperV2(model.decoder, info, val_info)\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    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    # os.makedirs(cfg.output_dir, exist_ok=True)\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   \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        # results.append((experiment_name, score))\n    # results = [f'{experiment_name}: {score}\\n' for experiment_name, score in results]\n    # results = ''.join(results)\n    # with open(dest, 'a') as fp:\n    #     fp.write(results)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:12.975082Z","iopub.execute_input":"2024-06-24T11:55:12.975542Z","iopub.status.idle":"2024-06-24T11:55:13.010409Z","shell.execute_reply.started":"2024-06-24T11:55:12.975500Z","shell.execute_reply":"2024-06-24T11:55:13.009005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# acc = (y * pred_y).mean(dim=0) / y.mean(dim=0)\n# plt.plot(acc.sort()[0])\n# # acc.max()\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:13.012273Z","iopub.execute_input":"2024-06-24T11:55:13.012783Z","iopub.status.idle":"2024-06-24T11:55:13.026463Z","shell.execute_reply.started":"2024-06-24T11:55:13.012747Z","shell.execute_reply":"2024-06-24T11:55:13.025181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"openvino_core = None","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:13.028129Z","iopub.execute_input":"2024-06-24T11:55:13.028944Z","iopub.status.idle":"2024-06-24T11:55:13.037251Z","shell.execute_reply.started":"2024-06-24T11:55:13.028879Z","shell.execute_reply":"2024-06-24T11:55:13.035976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"editable":true,"slideshow":{"slide_type":""},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef convert_test_time_model(model, decoder_type, info, val_info, to_prob=False, multilabel=False):\n    if decoder_type == 'fc':\n        if 'selected_class_names' in info:\n            model = TestTimeModelFinal(model, info, val_info, to_prob=to_prob, multilabel=multilabel)\n        else:\n            model = TestTimeModelFc(model, info, val_info, to_prob=to_prob, multilabel=multilabel)\n    elif decoder_type == 'well':\n        model = TestTimeModelWell(model, info, val_info, to_prob=to_prob, multilabel=multilabel)        \n    return model\n\nclass TestBirdCLEFSelectedClassWrapperV2(nn.Module):\n    def __init__(self, module, train_info, val_info):\n        super().__init__()\n        self.module = module\n        \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            if bc in train_info['selected_class_names']:\n                j = train_info['reduced_classes'].index(bc)\n                reduced_to_base_classes[i, j] = 1\n        # reduced_to_base_classes = reduced_to_base_classes / len(train_info['reduced_classes']) * len(train_info['selected_class_names'])\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 TestTimeModelFinal(nn.Module):\n    def __init__(self, model, info, val_info, to_prob=False, multilabel=False):\n        super().__init__()\n        self.to_prob = to_prob\n        self.multilabel = multilabel\n        self.encoder = model.encoder[0].module\n        self.decoder = TestBirdCLEFSelectedClassWrapperV2(model.decoder, info, val_info)\n        \n    def forward(self, x):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)[-1]\n        x = x.mean(dim=(2, 3))\n        x = self.decoder(x)\n        if self.to_prob:\n            if self.multilabel:\n                x = torch.sigmoid(x)\n            else:\n                x = torch.softmax(x, dim=1)\n        return x\n    \nclass TestTimeModelFc(nn.Module):\n    def __init__(self, model, info, val_info, to_prob=False, multilabel=False):\n        super().__init__()\n        self.to_prob = to_prob\n        self.multilabel = multilabel\n        self.encoder = model.encoder[0].module\n        self.decoder = TestBirdCLEFWrapperV2(model.decoder, info, val_info)\n        \n    def forward(self, x):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)[-1]\n        x = x.mean(dim=(2, 3))\n        x = self.decoder(x)\n        if self.to_prob:\n            if self.multilabel:\n                x = torch.sigmoid(x)\n            else:\n                x = torch.softmax(x, dim=1)\n        return x\n    \nclass TestTimeModelWell(nn.Module):\n    def __init__(self, model, info, val_info, to_prob=False, multilabel=False):\n        super().__init__()\n        self.to_prob = to_prob\n        self.multilabel = multilabel\n        self.encoder = model.encoder.module\n        self.decoder = TestBirdCLEFWrapper(model.decoder, info, val_info)\n        \n    def forward(self, x):\n        x = x.unsqueeze(1)\n        x = self.encoder(x)[-1]\n        x = self.decoder(x)\n        if self.to_prob:\n            if self.multilabel:\n                x = torch.sigmoid(x)\n            else:\n                x = torch.softmax(x, dim=1)\n        return x\n\n    \n# class Ensemble(nn.Module):\n#     def __init__(self, model_list, avg=False):\n#         super().__init__()\n#         self.model_list = nn.ModuleList(model_list)\n#         self.avg = avg\n\n#     def forward(self, x):\n#         outputs = [model(x) for model in self.model_list]\n#         if self.avg:\n#             outputs = sum(outputs) / len(outputs)\n#         return outputs\n\nclass FinalEnsemble:\n    def __init__(self, processor, model_list, hard_cls_path):\n        self.processor = processor\n        self.model1 = model_list[0]\n        self.model2 = model_list[1]\n        with open(hard_cls_path, 'r') as fp:\n            hard_cls_idx = json.load(fp)\n        self.hard_cls_idx = hard_cls_idx\n\n    def __call__(self, waveform, second=None):\n        waveform = waveform.squeeze(dim=1)\n        spec1 = self.processor(waveform, 0)\n        spec2 = self.processor(waveform, 1)\n        self.processor.end()\n        # assert softmax\n        prob1 = self.model1(spec1)\n        prob2 = self.model2(spec2)\n        \n        entropy2 = - torch.sum(prob2 * torch.log(prob2.clip(1e-7)), dim=1)\n        _alpha = 0.8 * (entropy2[:, None].clip(0.5, 2) - 0.5) / 1.5 + 0.1\n        alpha = torch.ones_like(prob1)\n        alpha[:, self.hard_cls_idx] *= _alpha\n        prob = alpha * prob1 + (1 - alpha) * prob2\n        prob = [prob]\n        return prob\n    \nclass InferenceModel:\n    def __init__(self, cfg, processor, model=None, model_list=None):\n        self.processor = processor\n        self.model_list = model_list\n        self.model = model\n        self.pre_prob = cfg.model.pre_prob\n        self.use_multilabel = cfg.model.use_multilabel\n        self.ensemble_weighs = None\n        if self.model_list is not None:\n            if cfg.model.ensemble_weighs is not None:\n                with open(cfg.model.ensemble_weighs, 'r') as fp:\n                    weights = json.load(fp)\n                weights = torch.FloatTensor(weights)\n                self.ensemble_weighs = weights.softmax(dim=0)\n            else:\n                self.ensemble_weighs = torch.ones(len(model_list)) / len(model_list)\n            self.second_to_indice = {sec: [i \n                                          for i, s in enumerate(cfg.dataset.audio_second)\n                                          if s == sec]\n                                    for sec in set(cfg.dataset.audio_second)}\n            self.second_to_model = {sec: [model \n                                          for model, s in zip(self.model_list, cfg.dataset.audio_second)\n                                          if s == sec]\n                                    for sec in set(cfg.dataset.audio_second)}\n            self.second_to_weight = {sec: [weights \n                                           for weights, s in zip(self.ensemble_weighs, cfg.dataset.audio_second)\n                                           if s == sec]\n                                     for sec in set(cfg.dataset.audio_second)}\n#             self.second_to_weight = {k: [w / sum(weights) for w in weights]\n#                                      for k, weights in self.second_to_weight.items()}\n    \n    def __call__(self, waveform, second=None):  \n        waveform = waveform.squeeze(dim=1)\n        if self.model is not None:\n            spec = self.processor(waveform)\n            logits = self.model(spec)\n        else:\n            if second is not None:\n                model_list = self.second_to_model[second]\n                weights = self.second_to_weight[second]\n                indice = self.second_to_indice[second]\n            else:\n                model_list = self.model_list\n                weights = self.ensemble_weighs\n                indice = range(len(self.model_list))\n            spec_list = [self.processor(waveform, i) for i in indice]\n            self.processor.end()\n            logits = [w[None] * model(spec) for model, w, spec in zip(model_list, weights, spec_list)]\n            logits = sum(logits)\n        \n        if not self.pre_prob:\n            if self.use_multilabel:\n                logits = torch.sigmoid(logits)\n            else:\n                logits = torch.softmax(logits, dim=1)\n                \n        logits = [logits]\n        return logits\n    \n# def get_ensemble_func(cfg, model_list):\n#     if cfg.model.ensemble_weighs is not None:\n#         with open(cfg.model.ensemble_weighs, 'r') as fp:\n#             weights = json.load(fp)\n#         weights = torch.FloatTensor(weights)\n#         weights = weights.softmax(dim=0)\n#         def ensemble_func(x):\n#             outputs = [w[None] * model(x) for model, w in zip(model_list, weights)]\n#             outputs = sum(outputs)\n#             return outputs\n#     else:\n#         def ensemble_func(x):\n#             outputs = [model(x) for model in model_list]\n#             outputs = sum(outputs) / len(outputs)\n#             return outputs\n# #         def ensemble_func(*args, **kwargs):\n# #             with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor:\n# #                 outputs = [executor.submit(model, *args, **kwargs) for model in model_list]\n# #             outputs = [o.result() for o in outputs]\n# #             outputs = sum(outputs) / len(outputs)\n# #             return outputs\n#     return ensemble_func\n\nclass EnsembleProcessor:\n    def __init__(self):\n        self._processor_dict = dict()\n        self._indice2key = []\n        self._processed = None\n        self._indice2processor = None\n        \n    def _get_key(self, second, image_size):\n        return f'{second}-{image_size}'\n    \n    def add(self, processor, second, image_size):\n        key = self._get_key(second, image_size)\n        self._processor_dict[key] = processor\n        self._indice2key.append(key)\n    \n    def finish_add(self):\n        self._indice2processor = [self._processor_dict[key] for key in self._indice2key]\n        self._processed = [None for _ in range(len(self._indice2processor))]\n        key2item = {k:[None] for k in self._processor_dict.keys()} # same key share the same list\n        for i, k in enumerate(self._indice2key):\n            self._processed[i] = key2item[k]\n        \n    def end(self):\n        for item in self._processed:\n                item[0] = None\n        \n    def __call__(self, x, i):\n        if self._processed[i][0] is None:\n            self._processed[i][0] = self._indice2processor[i](x)\n        return self._processed[i][0]\n        \n        \ndef get_inference_model(cfg):\n    if isinstance(cfg.model.name, (list, tuple, omegaconf.listconfig.ListConfig)):\n        if not isinstance(cfg.dataset.audio_second, (list, tuple, omegaconf.listconfig.ListConfig)):\n            cfg.dataset.audio_second = len(cfg.model.name) * [cfg.dataset.audio_second]\n        if not isinstance(cfg.dataset.image_size, (list, tuple, omegaconf.listconfig.ListConfig)):\n            cfg.dataset.image_size = len(cfg.model.name) * [cfg.dataset.image_size]\n        if not isinstance(cfg.dataset.info_path, (list, tuple, omegaconf.listconfig.ListConfig)):\n            cfg.dataset.info_path = len(cfg.model.name) * [cfg.dataset.info_path]\n        ensemble_processor = EnsembleProcessor()\n        model_list = []\n        model_names = cfg.model.name\n        decoder_types = cfg.model.decoder_type\n        checkpoints = cfg.model.checkpoint\n        audio_second = cfg.dataset.audio_second\n        image_size = cfg.dataset.image_size\n        info_path = cfg.dataset.info_path\n        for i, (name, d_type, ckpt, sec, im_size, path) in enumerate(zip(model_names, decoder_types, checkpoints, audio_second, image_size, info_path)):\n            cfg.model.name = name # hack\n            cfg.model.decoder_type = d_type # hack\n            cfg.model.checkpoint = ckpt # hack\n            cfg.dataset.audio_second = sec # hack\n            cfg.dataset.image_size = im_size # hack\n            cfg.dataset.info_path = path # hack\n            model = get_inference_model(cfg)\n            model_list.append(model.model)\n            ensemble_processor.add(model.processor, sec, im_size)\n        cfg.model.name = model_names\n        cfg.model.decoder_type = decoder_types\n        cfg.model.checkpoint = checkpoints\n        cfg.dataset.audio_second = audio_second\n        cfg.dataset.image_size = image_size\n        cfg.dataset.info_path = info_path\n        \n        ensemble_processor.finish_add()\n        if cfg.model.use_final_ensemble:\n            model = FinalEnsemble(ensemble_processor, model_list, cfg.model.hard_cls_path)\n        else:\n            model = InferenceModel(cfg, ensemble_processor, model_list=model_list)\n        return model\n    else:\n        is_auto = False\n        if cfg.model.format == 'auto':\n            is_auto = True\n            OmegaConf.update(cfg, 'model.format', \n                             {'efficientnetv2': 'openvino',\n                              'convnext_small': 'openvino',\n                              'convnextv2_tiny': 'openvino',\n                              'eca_nfnet': 'onnx',\n                              'pvt_v2_b2': 'torch',}[cfg.model.name])\n\n\n        if cfg.model.format == 'onnx':\n            model = build_onnx_model(cfg)\n        elif cfg.model.format == 'openvino':\n            model = build_openvino_model(cfg)\n        elif cfg.model.format == 'torch':\n            model = torch.jit.script(build_torch_model(cfg))\n        else:\n            raise ValueError(f'Invalid format: {cfg.model.format}')\n        if is_auto: # restore\n            OmegaConf.update(cfg, 'model.format', 'auto')\n            \n        # build processor\n        processor = get_spec_processor(cfg.dataset.sampling_rate, cfg.dataset.audio_second, cfg.dataset.image_size)[cfg.dataset.augmentation.processor]\n        model = InferenceModel(cfg, processor, model=model)\n        return model\n        \n\ndef clip_values(tensor):\n    tensor[tensor.abs() < 1e-7] = 0\n    return tensor\n\ndef export_onnx(model, example_input, path):\n    model.eval()\n\n    torch.onnx.export(model,               # model being run\n                  example_input,# model input (or a tuple for multiple inputs)\n                  path,   # where to save the model (can be a file or file-like object)\n#                   export_params=True,        # store the trained parameter weights inside the model file\n#                   opset_version=10,          # the ONNX version to export the model to\n                  do_constant_folding=True,  # whether to execute constant folding for optimization\n                  input_names = ['x'],   # the model's input names\n                  output_names = ['logit'], # the model's output names\n                  training=torch.onnx.TrainingMode.EVAL,\n                  dynamic_axes={'x' : {0 : 'batch_size'},    # variable length axes\n                                'logit' : {0 : 'batch_size'}})\n    \ndef extract_patches_v2(waveform, sr, size, stride=5):\n    center_indices = torch.arange(0, waveform.shape[1], stride*sr) + int(stride / 2 * sr)\n    start_indices = center_indices - int(sr * size * 0.5)\n    end_indices = start_indices + sr * size\n    end_indices += torch.relu(-start_indices)\n    start_indices = torch.relu(start_indices)\n\n    start_indices -= torch.relu(end_indices - waveform.shape[1])\n    end_indices = torch.clip(end_indices, 0, waveform.shape[1])\n\n    patches_waveform = [waveform[:, i:j] for i, j in zip(start_indices, end_indices)]\n    patches_waveform = torch.stack(patches_waveform, dim=0)\n    return patches_waveform    \n\ndef multi_outputs_patch_predict(waveforms, model_func,\n                                sampling_rate=32000,\n                                size = 10,\n                                stride = 5,\n                                batch_size=8):\n    ### ensure list ###\n    if not isinstance(waveforms, (tuple, list)):\n        waveforms = [waveforms]\n\n\n    ### ensure size of image is larger than size[0] and size[1] ###\n    for i, waveform in enumerate(waveforms):\n        if waveform.shape[1] < sampling_rate * size:\n            waveform = Replay(sampling_rate, size)(waveform)\n            waveforms[i] = waveform\n\n    # [n, channels, size]\n    patches_list = [extract_patches_v2(waveform, sampling_rate, size, stride) \n                    for waveform in waveforms]\n    num_patches = [patches.shape[0] for patches in patches_list]\n    patches = torch.cat(patches_list, dim=0)\n    \n    ### predict the segmentation mask ###\n    with torch.no_grad():\n        patches_outputs_list = list(zip(*[\n            [out.detach().cpu() for out in model_func(\n                patches[i: min(i+batch_size, patches.shape[0])],\n                size,\n            )]\n            for i in list(range(0, patches.shape[0], batch_size))\n        ]))\n        patches_outputs_list = [torch.cat(patches_outputs, dim=0) for patches_outputs in patches_outputs_list]\n\n    ### split the patches ##\n    flat_patch_outputs_list = [torch.split(patches_outputs, \n                                    num_patches, \n                                    dim=0)\n                                for patches_outputs in patches_outputs_list]\n\n    return flat_patch_outputs_list\n\ndef build_torch_model(cfg):\n    info = load_dataset_information(cfg.dataset.info_path)\n    val_info = load_dataset_information(cfg.dataset.val_info_path)\n    model = get_model(cfg.model.name, cfg.model.decoder_type, info)\n    if cfg.model.checkpoint is not None:\n        state_dict = torch.load(cfg.model.checkpoint, map_location=torch.device('cpu'))\n        state_dict = {k:clip_values(v) for k, v in state_dict.items()}\n        model.load_state_dict(state_dict)\n        \n#     model.decoder = TestBirdCLEFWrapper(model.decoder, info, val_info)\n#     model.decoder = TestBirdCLEFWrapperV2(model.decoder, info, val_info)\n    model = convert_test_time_model(model, cfg.model.decoder_type, info, val_info, \n                                    to_prob=cfg.model.pre_prob, multilabel=cfg.model.use_multilabel)\n    model.eval()\n    return model\n\ndef build_onnx_model(cfg):\n    import onnxruntime\n    if cfg.model.checkpoint is None:\n        fn = cfg.model.name\n    else:\n        fn = os.path.splitext(os.path.split(cfg.model.checkpoint)[1])[0]\n    if not os.path.exists(f'asset/{fn}.onnx'):\n        torch_model = build_torch_model(cfg)\n        processor = get_spec_processor(cfg.dataset.sampling_rate, cfg.dataset.audio_second, cfg.dataset.image_size)[cfg.dataset.augmentation.processor]\n        export_onnx(torch_model, \n                    processor(torch.randn(1, cfg.dataset.sampling_rate * cfg.dataset.audio_second)), \n                    f'asset/{fn}.onnx')\n        \n    model = onnxruntime.InferenceSession(f'asset/{fn}.onnx')\n    \n    def model_func(spec):\n        outs = model.run(\n            None,\n            {\"x\": spec.numpy()},\n        )[0]\n        outs = torch.from_numpy(outs)\n        return outs\n    return model_func\n\ndef build_openvino_model(cfg):\n    import onnxruntime\n    if cfg.model.checkpoint is None:\n        fn = cfg.model.name\n    else:\n        fn = os.path.splitext(os.path.split(cfg.model.checkpoint)[1])[0]\n    if not os.path.exists(f'asset/{fn}.onnx'):\n        torch_model = build_torch_model(cfg)\n        processor = get_spec_processor(cfg.dataset.sampling_rate, cfg.dataset.audio_second, cfg.dataset.image_size)[cfg.dataset.augmentation.processor]\n        export_onnx(torch_model,\n                    processor(torch.randn(1, cfg.dataset.sampling_rate * cfg.dataset.audio_second)), \n                    f'asset/{fn}.onnx')\n\n    from openvino.tools.mo import convert_model\n    import openvino.runtime\n    global openvino_core\n    if openvino_core is None:\n        openvino_core = openvino.runtime.Core()\n    model = convert_model(f'asset/{fn}.onnx', compress_to_fp16=True)\n    model = openvino_core.compile_model(model, device_name=\"CPU\")\n    model = model.create_infer_request()\n    \"\"\n    def model_func(spec):\n        outs = model.infer(\n            inputs=[spec.numpy()],\n        )['logit']\n        outs = torch.from_numpy(outs)\n        return outs\n    return model_func\n\ndef predict(cfg):\n    experiment_name = f'{cfg.model.name}_{cfg.train.loss_name}_{cfg.train.epochs}_{cfg.dataset.kfold}_{cfg.seed}'\n    \n    if cfg.seed is not None:\n        seed_everything(cfg.seed)\n    \n#     info = load_dataset_information(cfg.dataset.info_path)\n    val_info = load_dataset_information(cfg.dataset.val_info_path)\n    \n    model = get_inference_model(cfg)\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    \n    def _get_paths(d):\n        path_list = []\n        for fn in os.listdir(d):\n            path = os.path.join(d, fn)\n            if os.path.isdir(path):\n                path_list += _get_paths(path)\n            else:\n                path_list.append(path)\n        return path_list\n            \n    paths = _get_paths(cfg.dataset.root)\n    paths = [path for path in paths if path.endswith('.ogg')]\n    \n    def load_transform(path):\n        waveform, sr = torchaudio.load(path)\n        waveform = torch.mean(waveform, dim=0, keepdim=True)\n        waveform = torchaudio.transforms.Resample(sr, cfg.dataset.sampling_rate)(waveform)\n        return waveform\n    \n    dataset = AudioDatasetWrapper(paths, load_transform)\n    \n    \n    sampler = torch.utils.data.SequentialSampler(dataset)\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            dataset,\n            sampler=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=generator,\n            collate_fn=lambda data: data,\n    )\n\n#     def model_func(waveform):\n#         spec = processor(waveform.squeeze(dim=1))\n#         outs = model(spec)\n#         if not isinstance(outs, (tuple, list)):\n#             outs = [outs]\n#         if not cfg.model.pre_prob:\n#             if cfg.model.use_multilabel:\n#                 outs = [torch.sigmoid(o) for o in outs]\n#             else:\n#                 outs = [torch.softmax(o, dim=1) for o in outs]\n#         return outs\n         \n    loader_iter = iter(loader)\n    outputs_list = []\n    end_seconds = []\n    for waveforms in tqdm(loader_iter, total=len(loader)):\n        if isinstance(cfg.dataset.audio_second, (list, tuple, omegaconf.listconfig.ListConfig)):\n            outs_list = []\n            for sec in set(cfg.dataset.audio_second):\n                outs = multi_outputs_patch_predict(waveforms, model, \n                                                   sampling_rate=cfg.dataset.sampling_rate,\n                                                   size=sec,\n                                                   stride=5, \n                                                   batch_size=cfg.dataset.inner_batch_size)\n                outs_list.append(outs[0])\n            outs = [sum(o) for o in zip(*outs_list)]\n        else:\n            outs = multi_outputs_patch_predict(waveforms, model, \n                                               sampling_rate=cfg.dataset.sampling_rate,\n                                               size=cfg.dataset.audio_second,\n                                               stride=5, \n                                               batch_size=cfg.dataset.inner_batch_size)\n            outs = outs[0]\n        outputs_list.extend(outs)\n        end_seconds.extend([waveform.shape[-1] / cfg.dataset.sampling_rate for waveform in waveforms])\n    \n    return outputs_list, paths, end_seconds\n\ndef build_submission_df(info, pred_y, paths, end_seconds, audio_second=5):\n    row_ids = [os.path.splitext(os.path.split(path)[1])[0] for path in paths]\n    ends = [[min(audio_second * (1 + i), int(sec)) for i in range(len(p_y))] \n            for p_y, sec in zip(pred_y, end_seconds)]\n    row_ids = [row_id + \"_\" + str(sec) \n               for row_id, seconds in zip(row_ids, ends) \n               for sec in seconds]\n    if len(pred_y) == 0:\n        pred_y = torch.zeros(0, len(info['base_classes']))\n    else:\n        pred_y = torch.cat(pred_y, dim=0)\n    columns = [pd.Series(pred_y[:, i], dtype=pd.Float32Dtype(), name=name) \n               for i, name in enumerate(info['base_classes'])]\n    \n    columns = [pd.Series(row_ids, dtype=pd.StringDtype(), name='row_id'),\n               *columns]\n    return pd.concat(columns, axis=1)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:14.094927Z","iopub.execute_input":"2024-06-24T11:55:14.095352Z","iopub.status.idle":"2024-06-24T11:55:14.220341Z","shell.execute_reply.started":"2024-06-24T11:55:14.095322Z","shell.execute_reply":"2024-06-24T11:55:14.218801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = 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, 'model.format', 'auto') # ['auto', 'torch', 'onnx', 'openvino']\n\nOmegaConf.update(cfg, 'dataset.sampling_rate', 32000)\n\nOmegaConf.update(cfg, 'model.use_multilabel', False)\nOmegaConf.update(cfg, 'model.pre_prob', True)\n\n\nOmegaConf.update(cfg, 'seed', 1)\nOmegaConf.update(cfg, 'dataset.batch_size', 2)\nOmegaConf.update(cfg, 'dataset.inner_batch_size', 32)\nOmegaConf.update(cfg, 'dataset.num_workers', 1)\nOmegaConf.update(cfg, 'dataset.prefetch_factor', 8)\n# OmegaConf.update(cfg, 'model.checkpoint', None)\n# OmegaConf.update(cfg, 'dataset.info_path', './asset/info2024.json')\nOmegaConf.update(cfg, 'dataset.val_info_path', './asset/info2024_primary.json')\n\n\nOmegaConf.update(cfg, 'model.use_final_ensemble', False)\nOmegaConf.update(cfg, 'model.ensemble_weighs', None)\n\nOmegaConf.update(cfg, 'model.name', 'convnextv2_tiny')\nOmegaConf.update(cfg, 'dataset.image_size', 384)\nOmegaConf.update(cfg, 'dataset.audio_second', 5)\nOmegaConf.update(cfg, 'model.decoder_type', 'fc')\nOmegaConf.update(cfg, 'model.checkpoint', './asset/2_2024_p0.5_convnextv2_tiny_fc_ce_20_3_1_uda_weights.pt')\nOmegaConf.update(cfg, 'model.hard_cls_path', None)\nOmegaConf.update(cfg, 'dataset.info_path', './asset/info2024_p0.json')\n\n# 'pvt_v2_b2', 'convnext_small', 'convnextv2_tiny', 'eca_nfnet', 'efficientnetv2'\n# OmegaConf.update(cfg, 'model.name', ['eca_nfnet', 'efficientnetv2'])\n# OmegaConf.update(cfg, 'dataset.image_size', [384, 384])\n# OmegaConf.update(cfg, 'dataset.audio_second', [10, 10])\n# OmegaConf.update(cfg, 'model.decoder_type', ['fc', 'fc'])\n# OmegaConf.update(cfg, 'model.checkpoint', \n#                  ['./asset/2_2024_p0.5_eca_nfnet_fc_ce_15_None_1_weights.pt',\n#                   './asset/2_2024sel_p0.5_efficientnetv2_fc_ce_5_None_1_weights.pt',])\n# OmegaConf.update(cfg, 'model.hard_cls_path', './asset/hard_classes.json')\n# OmegaConf.update(cfg, 'dataset.info_path', ['./asset/info2024_p0.json',\n#                                             './asset/info2024sel_p0.5.json'])\n\n# OmegaConf.update(cfg, 'dataset.root', '/kaggle/input/birdclef-2024/train_audio/brfowl1')\nOmegaConf.update(cfg, 'dataset.root', '/kaggle/input/birdclef-2024/test_soundscapes')\n    \n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:40.826545Z","iopub.execute_input":"2024-06-24T11:55:40.827001Z","iopub.status.idle":"2024-06-24T11:55:40.868628Z","shell.execute_reply.started":"2024-06-24T11:55:40.826960Z","shell.execute_reply":"2024-06-24T11:55:40.867252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"outputs, paths, end_second = predict(cfg)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:41.003104Z","iopub.execute_input":"2024-06-24T11:55:41.003527Z","iopub.status.idle":"2024-06-24T11:55:49.072550Z","shell.execute_reply.started":"2024-06-24T11:55:41.003496Z","shell.execute_reply":"2024-06-24T11:55:49.070984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"info = load_dataset_information('./asset/info2024_primary.json')\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:49.075992Z","iopub.execute_input":"2024-06-24T11:55:49.076383Z","iopub.status.idle":"2024-06-24T11:55:49.132302Z","shell.execute_reply.started":"2024-06-24T11:55:49.076350Z","shell.execute_reply":"2024-06-24T11:55:49.131186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = build_submission_df(info, outputs, paths, end_second)\nsubmission_df.to_csv('submission.csv', index=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-24T11:55:49.133717Z","iopub.execute_input":"2024-06-24T11:55:49.134115Z","iopub.status.idle":"2024-06-24T11:55:49.209183Z","shell.execute_reply.started":"2024-06-24T11:55:49.134083Z","shell.execute_reply":"2024-06-24T11:55:49.207993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-06-10T15:16:44.297333Z","iopub.execute_input":"2024-06-10T15:16:44.297832Z","iopub.status.idle":"2024-06-10T15:16:46.583147Z","shell.execute_reply.started":"2024-06-10T15:16:44.297784Z","shell.execute_reply":"2024-06-10T15:16:46.581431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !rm asset/*.onnx","metadata":{"execution":{"iopub.status.busy":"2024-06-10T20:14:01.387752Z","iopub.execute_input":"2024-06-10T20:14:01.388176Z","iopub.status.idle":"2024-06-10T20:14:02.550606Z","shell.execute_reply.started":"2024-06-10T20:14:01.388143Z","shell.execute_reply":"2024-06-10T20:14:02.548862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !wget https://files.pythonhosted.org/packages/12/b9/ee24c41fb9fd12c72bf7175703bbcde4483a5f93fab36ca2c6022206809c/onnxruntime-1.18.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl\n# !wget https://files.pythonhosted.org/packages/c9/df/cbf6e024c62033c29aa34830531844292f89d814850f3ca1a6ef1d7ab772/openvino-2024.1.0-15008-cp310-cp310-manylinux2014_x86_64.whl\n# !wget https://files.pythonhosted.org/packages/2f/6c/89a38fd365446476be13811f565da44d52be4f184d7cb0a0abfc7970c0f5/openvino_telemetry-2024.1.0-py3-none-any.whl\n# !wget https://files.pythonhosted.org/packages/95/e0/26f07522ed343a727a0b5aa10222a5d30b629886c240da2a7fbc67676c4c/openvino_dev-2024.1.0-15008-py3-none-any.whl\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-06-03T17:26:39.180503Z","iopub.execute_input":"2024-06-03T17:26:39.181281Z","iopub.status.idle":"2024-06-03T17:26:39.666142Z","shell.execute_reply.started":"2024-06-03T17:26:39.181147Z","shell.execute_reply":"2024-06-03T17:26:39.664811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}