{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":2315253,"sourceType":"datasetVersion","datasetId":1389941},{"sourceId":11925292,"sourceType":"datasetVersion","datasetId":6847135},{"sourceId":12145524,"sourceType":"datasetVersion","datasetId":7649448},{"sourceId":12145632,"sourceType":"datasetVersion","datasetId":7649526},{"sourceId":12145693,"sourceType":"datasetVersion","datasetId":7649570},{"sourceId":12145709,"sourceType":"datasetVersion","datasetId":7649581},{"sourceId":12145724,"sourceType":"datasetVersion","datasetId":7649591},{"sourceId":12145731,"sourceType":"datasetVersion","datasetId":7649596},{"sourceId":12145734,"sourceType":"datasetVersion","datasetId":7649598},{"sourceId":12145737,"sourceType":"datasetVersion","datasetId":7649600},{"sourceId":12145743,"sourceType":"datasetVersion","datasetId":7649605},{"sourceId":12145881,"sourceType":"datasetVersion","datasetId":7649694},{"sourceId":12145884,"sourceType":"datasetVersion","datasetId":7649696},{"sourceId":12145888,"sourceType":"datasetVersion","datasetId":7649699},{"sourceId":12145892,"sourceType":"datasetVersion","datasetId":7649702},{"sourceId":12145913,"sourceType":"datasetVersion","datasetId":7649716},{"sourceId":12145917,"sourceType":"datasetVersion","datasetId":7649719},{"sourceId":12145920,"sourceType":"datasetVersion","datasetId":7649722},{"sourceId":12145925,"sourceType":"datasetVersion","datasetId":7649725},{"sourceId":245116990,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Libs","metadata":{"_kg_hide-input":true}},{"cell_type":"code","source":"if 'SKIP' not in globals() or not SKIP:   # torch\n    print('libs torch')\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.0 imports\n    # ----------------------------------------------------------------------------------------------------\n\n    import warnings\n    warnings.filterwarnings('ignore')\n\n    import pyarrow.parquet as pq\n\n    import copy\n    from contextlib import contextmanager, nullcontext\n    \n    import os\n    import gc\n    import time\n    import pickle\n    import psutil\n    import ast\n    import copy\n    import math\n    import random\n    import pdb\n    import json\n    import glob\n    import shutil\n\n    # from tqdm import tqdm\n    from tqdm.autonotebook import tqdm\n    import wandb\n    from collections import deque, defaultdict\n    import numba as nb\n\n    import numpy as np\n    import numpy.matlib\n    import pandas as pd\n    # import datatable as dt\n    # import polars\n\n    pd.set_option('display.max_columns', 300)\n    pd.set_option('display.max_rows', 100)\n    pd.set_option(\"display.width\", 150)\n    np.set_printoptions(linewidth=140)\n\n    from multiprocessing.pool import ThreadPool\n\n    # torch\n    import torch\n    # from torch import cuda\n    from torch import nn\n    import torch.nn.functional as F\n    from torch.utils.data import Dataset, IterableDataset, DataLoader\n    import pytorch_lightning as L\n    from pytorch_lightning.utilities.model_summary import ModelSummary\n    from pytorch_lightning.callbacks import (LearningRateMonitor, ModelCheckpoint, TQDMProgressBar, StochasticWeightAveraging)\n    from pytorch_lightning.loggers import WandbLogger\n    import torchaudio\n    import torchaudio.transforms as T\n    import torchvision\n    from torchvision.transforms import v2\n    from torch.distributions import Beta\n    import cv2\n    from PIL import Image, ImageOps\n\n    import os\n\n    # visual models\n    import timm\n    from timm.models.layers import drop_path\n\n    import albumentations as A\n\n    # metrics\n    import sklearn.metrics\n\n    # multiprocessing\n    # import joblib\n    from joblib import Parallel, delayed\n    from joblib.externals.loky.backend.context import get_context\n    import concurrent.futures\n\n    # onnx\n    try:\n        import onnx\n        import onnxruntime as ort\n    except:\n        pass\n\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.0 constants\n    # ----------------------------------------------------------------------------------------------------\n\n    STRICT = False\n\n    SR = 32_000\n    \n    BIRDS = [\n        '1139490', '1192948', '1194042', '126247', '1346504', '134933', '135045', '1462711', '1462737', '1564122', '21038', '21116',\n        '21211', '22333', '22973', '22976', '24272', '24292', '24322', '41663', '41778', '41970', '42007', '42087', '42113', '46010',\n        '47067', '476537', '476538', '48124', '50186', '517119', '523060', '528041', '52884', '548639', '555086', '555142', '566513',\n        '64862', '65336', '65344', '65349', '65373', '65419', '65448', '65547', '65962', '66016', '66531', '66578', '66893', '67082',\n        '67252', '714022', '715170', '787625', '81930', '868458', '963335', \n        'amakin1', 'amekes', 'ampkin1', 'anhing', 'babwar', 'bafibi1', 'banana', 'baymac', 'bbwduc', 'bicwre1', 'bkcdon', 'bkmtou1', \n        'blbgra1', 'blbwre1', 'blcant4', 'blchaw1', 'blcjay1', 'blctit1', 'blhpar1', 'blkvul', 'bobfly1', 'bobher1', 'brtpar1', 'bubcur1',\n        'bubwre1', 'bucmot3', 'bugtan', 'butsal1', 'cargra1', 'cattyr', 'chbant1', 'chfmac1', 'cinbec1', 'cocher1', 'cocwoo1', 'colara1',\n        'colcha1', 'compau', 'compot1', 'cotfly1', 'crbtan1', 'crcwoo1', 'crebob1', 'cregua1', 'creoro1', 'eardov1', 'fotfly', 'gohman1',\n        'grasal4', 'grbhaw1', 'greani1', 'greegr', 'greibi1', 'grekis', 'grepot1', 'gretin1', 'grnkin', 'grysee1', 'gybmar', 'gycwor1', \n        'labter1', 'laufal1', 'leagre', 'linwoo1', 'littin1', 'mastit1', 'neocor', 'norscr1', 'olipic1', 'orcpar', 'palhor2', 'paltan1',\n        'pavpig2', 'piepuf1', 'pirfly1', 'piwtyr1', 'plbwoo1', 'plctan1', 'plukit1', 'purgal2', 'ragmac1', 'rebbla1', 'recwoo1', 'rinkin1',\n        'roahaw', 'rosspo1', 'royfly1', 'rtlhum', 'rubsee1', 'rufmot1', 'rugdov', 'rumfly1', 'ruther1', 'rutjac1', 'rutpuf1', 'saffin',\n        'sahpar1', 'savhaw1', 'secfly1', 'shghum1', 'shtfly1', 'smbani', 'snoegr', 'sobtyr1', 'socfly1', 'solsan', 'soulap1', 'spbwoo1',\n        'speowl1', 'spepar1', 'srwswa1', 'stbwoo2', 'strcuc1', 'strfly1', 'strher', 'strowl1', 'tbsfin1', 'thbeup1', 'thlsch3', 'trokin',\n        'tropar', 'trsowl', 'turvul', 'verfly', 'watjac1', 'wbwwre1', 'whbant1', 'whbman1', 'whfant1', 'whmtyr1', 'whtdov', 'whttro1',\n        'whwswa1', 'woosto', 'y00678', 'yebela1', 'yebfly1', 'yebsee1', 'yecspi2', 'yectyr1', 'yehbla2', 'yehcar1', 'yelori1', 'yeofly1',\n        'yercac1', 'ywcpar', \n    ]\n\n    N_BIRDS = len(BIRDS)\n    RANGE_BIRDS = range(N_BIRDS)\n    \n    label_to_num = dict(zip(BIRDS, RANGE_BIRDS))\n\n    RARE_BIRDS = [ # 38 birds\n        '1139490', '1192948', '1194042', '126247', '1346504', '134933', '1462711', '1462737', '1564122',\n        '21038', '21116', '24272', '24292', '41778', '42087', '42113', '46010', '47067', '476537',\n        '476538', '523060', '528041', '548639', '555142', '64862', '65336', '65419', '65547', '66016',\n        '66531', '66578', '66893', '67082', '714022', '787625', '81930', '868458', '963335',\n    ]\n    \n    # ----------------------------------------------------------------------------------------------------\n    # 0.1 utils\n    # ----------------------------------------------------------------------------------------------------\n\n    class dotdict(dict):\n        def __getattr__(self, name):\n            return self.get(name, None)\n\n        def __setattr__(self, name, val):\n            self[name] = val\n\n        def __set__(self, name, val):\n            self.__setattr__(name, val)\n            \n        def to_dict(self):\n            return eval(str(self))\n\n\n    class Log:\n        def __init__(self, log_path, time_key=True):\n            self.path = log_path\n            if time_key:\n                self.path = self.path.replace('.','{}.'.format(time.strftime('_%Y%m%d%H%M%S', time.localtime(time.time()))))\n            print(time.strftime('%Y-%m-%d %H:%M:%S',time.localtime(time.time())), file=open(self.path,'a+'))\n            print('log path:', self.path)\n            print('------------ begin ------------', file=open(self.path, 'a+'))\n\n        def __call__(self, *content):\n            t1 = time.strftime('%H:%M:%S', time.localtime(time.time()))\n            print(*content)\n            print(t1, *content, file=open(self.path,'a+'))\n\n        def clean(self):\n            print(time.strftime('%Y-%m-%d %H:%M:%S',time.localtime(time.time())), file=open(self.path,'w'))\n            print('------------ begin -------------', file=open(self.path,'a+'))\n\n\n    class Config:\n        def __getattr__(self, name):\n            \"\"\" retun None if attribute doesn't exist \"\"\"\n            return None\n\n        def __repr__(self):\n            out = dict()\n            names = dir(self)\n            for key in names:\n                if not key.startswith('__') and key != 'to_dict':\n                    out[key] = getattr(self, key)\n            return str(out)\n\n        def to_dict(self):\n            config = dotdict(eval(str(self)))\n            if config.model:\n                config.model = dotdict(config.model)\n            if config.aug:\n                config.aug = dotdict(config.aug)\n            return config\n\n\n    def save_obj(obj, name, protocol=4): # pickle.HIGHEST_PROTOCOL):\n        with open('./'+ name + '.pkl', 'wb') as f:\n            pickle.dump(obj, f, protocol)\n\n\n    def load_obj(name, folder=''):\n        name = name.replace('.pkl', '')\n        with open(folder + name + '.pkl', 'rb') as f:\n            return pickle.load(f)\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.3 Models\n    # ----------------------------------------------------------------------------------------------------\n\n    class SpecNetImg(nn.Module):\n        \"\"\" linear head \"\"\"\n        def __init__(self, cfg, num_classes=N_BIRDS, channels_first=True, resize=None, add_position=False,\n                     pretrained=True, submit=False, in_chans=3, **kwargs):\n            super().__init__()\n            self.submit = submit\n            self.in_chans = in_chans\n            encoder = cfg.encoder or 'efficientnet_b0'\n            head_dropout = cfg.head_dropout or 0.\n            drop_path_rate = cfg.drop_path_rate\n            if isinstance(pretrained, str):\n                self.model = timm.create_model(\n                    encoder, pretrained=False, in_chans=in_chans, drop_path_rate=drop_path_rate)\n                load_encoder_state(self.model, pretrained)\n            else:\n                self.model = timm.create_model(\n                    encoder, pretrained=pretrained, in_chans=in_chans, drop_path_rate=drop_path_rate)\n            n_channels = self.model.num_features\n            self.pool = nn.AdaptiveAvgPool2d(1) if channels_first else lambda x: x.mean([1,2])\n            self.drop = nn.Dropout(p=head_dropout)\n            self.fc = nn.Linear(n_channels, out_features=num_classes, bias=True)\n\n        def forward(self, x):\n            B = x.size()[0]\n                \n            if self.in_chans != x.shape[1]:\n                x = x.expand(-1, 3, -1, -1)     # C3HW\n            x = self.model.forward_features(x)  # BHWC\n            x = self.pool(x)                    # BHWC -> BC or BC11\n            x = x.view(B, -1)  # flatten\n            x = self.drop(x)\n            x = self.fc(x)\n            return x\n           \n           \n    def load_encoder_state(encoder, chkp_path):\n        try:\n            chkp = torch.load(chkp_path)\n        except:\n            chkp = torch.load(chkp_path, map_location=torch.device('cpu'))\n        \n        if 'state_dict' in chkp:\n            encoder_state = {k[6:]: v for k, v in chkp['state_dict'].items()}\n            encoder_state = {k[6:]: v for k, v in encoder_state.items() if k[:6] == 'model.'}\n        else:\n            encoder_state = chkp.copy()\n            for k, v in encoder_state.items():\n                if '.running_mean' in k and torch.isnan(v).any():\n                    encoder_state[k] = torch.ones_like(v, dtype=torch.float32) * 0\n                if '.running_var' in k and torch.isnan(v).any():\n                    encoder_state[k] = torch.ones_like(v, dtype=torch.float32)\n\n        encoder_state = {\n             k.replace('backbone.0', 'stem').replace('backbone.1', 'stages')\n              .replace('backbone.2', 'final_conv').replace('stem_', 'stem.')\n              .replace('stages_', 'stages.').replace('global_pool.', ''): v \n                 for k, v in encoder_state.items()}\n        encoder.load_state_dict(encoder_state, strict=STRICT)  # for Vladimir pretrained encoders\n        print(f\"loaded encoder '{chkp_path}'\")\n        return encoder_state\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.4 augment\n    # ----------------------------------------------------------------------------------------------------\n    \n    def augment_img_train():\n        \"\"\" train augmentation to run on dataset \"\"\"\n        aug = config.aug\n        pipeline = []\n        if aug.Flip > 0:\n            pipeline.append(A.HorizontalFlip(p=aug.Flip))\n        \n        if aug.CoarseDropout is not None:  # (max_height=.375, max_width=.375, max_holes=1, p=.7)\n            max_height = int(config.img_dim[0] * aug.CoarseDropout[0]) # 0.375\n            max_width = int(config.img_dim[1] * aug.CoarseDropout[1]) # 0.375\n            pipeline.append(A.CoarseDropout(max_height=max_height, max_width=max_width,\n                            max_holes=aug.CoarseDropout[2], p=aug.CoarseDropout[3]))\n\n        tranform_fn = A.Compose(pipeline)\n        return tranform_fn\n\n    BETA = Beta(1., 1.)\n    \n    def mixup(batch):\n        \"\"\" mixup augmentation to run during training step \"\"\"\n        x, y = batch['x'], batch['y']\n        perm = torch.randperm(x.size(0))\n\n        # lam = torch.FloatTensor([np.random.beta(alpha, alpha)])\n        lam = BETA.rsample(x.shape[:1]).to(x.device)\n        lam_x = lam.view(-1, 1, 1, 1)\n        lam_y = lam.view(-1, 1)\n        x = x * lam_x + x[perm] * (1 - lam_x)\n        y = y * lam_y + y[perm] * (1 - lam_y)\n        if 'mask' in batch:\n            mask = batch['mask']\n            mask = mask * mask[perm]\n            batch['mask'] = mask\n\n        batch['x'], batch['y'] = x, y\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.5 Datasets\n    # ----------------------------------------------------------------------------------------------------\n\n    hop_length_256_5 = 32_000 * 5 // (256 - 1)\n\n    melspec_128_256_5 = T.MelSpectrogram(\n        n_fft=2048, hop_length=hop_length_256_5, f_min=50, f_max=16_000, sample_rate=32_000,\n        n_mels=128, norm='slaney', mel_scale='slaney', pad_mode='constant')\n    \n    db_transform = T.AmplitudeToDB(stype='power', top_db=80)\n        \n    len_audio_5s = 5 * 32_000\n    len_audio_1m = 1 * 60 * 32_000\n    \n    SPEC_DIMS = {}\n        \n\n    def transform_to_spec(dim=(128, 256), norm=False, duration=5, snipet=True, spec=None, bits=8):\n        \"\"\" converts an audio array to spectrograms. Audio is a 2D array with a single segment to check \n            against train or multiple segments for submit prediction. Transformation is performed per\n            audio snipet, which means that submitting an audio by itself (1, 32_000*5) or as part of a\n            batch (6, 32_000*5), produces the same result for the common audio snipet.\n        \"\"\"\n        if spec == 'p4_1':\n            melspec_fn = melspec_p4_1\n        elif spec == 'p4_1b':\n            melspec_fn = melspec_p4_1b\n        elif spec == 'p4_4':\n            melspec_fn = melspec_p4_4\n        elif spec == 'p4_4b':\n            melspec_fn = melspec_p4_4b\n        elif spec == 'f256':\n            melspec_fn = melspec_256_128_5\n\n        elif spec == 'f128x2':\n            melspec_fn = melspec_128_128_5\n        elif spec == 'f160x2':\n            melspec_fn = melspec_160_160_5\n        elif spec == 'f192x2':\n            melspec_fn = melspec_192_192_5\n        elif spec == 'f192':\n            melspec_fn = melspec_192_256_5\n\n        elif spec == 'f256x2':\n            melspec_fn = melspec_256_256_5\n        elif dim[-1] == 512:\n            if duration == 5:\n                melspec_fn = melspec_128_512_5\n        elif dim[-1] == 256:\n            if duration == 5:\n                melspec_fn = melspec_128_256_5\n        elif dim[-1] == 224:\n            if duration == 10:\n                melspec_fn = melspec_224_224b_10\n                \n        db_transform_ = db_transform_snipet if snipet else db_transform\n        eps = 1e-6\n        \n        def _process(audio):\n            spec = melspec_fn(audio)\n            spec = db_transform(spec)\n            min_ = torch.amin(spec)\n            max_ = torch.amax(spec)\n            spec = (spec - min_) / (max_ - min_)\n            if bits == 8 or bits is None:\n                spec = (spec * 255).to(torch.uint8) / 255\n            elif bits == 16:\n                spec = (spec * 65535).to(torch.uint32) / 65535\n            return spec\n\n        def _process_norm(audio):\n            spec = melspec_fn(audio)\n            spec = db_transform(spec)\n            mean_, std_ = spec.mean(), spec.std()\n            spec = (spec - mean_) / (std_ + eps)\n            min_ = torch.amin(spec)\n            max_ = torch.amax(spec)\n            spec = (spec - min_) / (max_ - min_)\n            if bits == 8 or bits is None:\n                spec = (spec * 255).to(torch.uint8) / 255\n            elif bits == 16:\n                spec = (spec * 65535).to(torch.uint32) / 65535\n            return spec\n\n        def _process_snipet(audio):\n            spec = melspec_fn(audio)\n            spec = db_transform(spec)\n            min_ = torch.amin(spec, dim=(-2,-1), keepdims=True)\n            max_ = torch.amax(spec, dim=(-2,-1), keepdims=True)\n            spec = (spec - min_) / (max_ - min_)\n            if bits == 8 or bits is None:\n                spec = (spec * 255).to(torch.uint8) / 255\n            elif bits == 16:\n                spec = (spec * 65535).to(torch.uint32) / 65535\n            return spec\n\n        def _process_norm_snipet(audio):\n            spec = melspec_fn(audio)\n            spec = db_transform(spec)\n            mean_, std_ = spec.mean((-2,-1), keepdims=True), spec.std((-2,-1), keepdims=True)\n            spec = (spec - mean_) / (std_ + eps)\n            min_ = torch.amin(spec, dim=(-2,-1), keepdims=True)\n            max_ = torch.amax(spec, dim=(-2,-1), keepdims=True)\n            spec = (spec - min_) / (max_ - min_)\n            if bits == 8 or bits is None:\n                spec = (spec * 255).to(torch.uint8) / 255\n            elif bits == 16:\n                spec = (spec * 65535).to(torch.uint32) / 65535\n            return spec\n\n        if snipet:\n            return _process_norm_snipet if norm else _process_snipet\n        else:\n            return _process_norm if norm else _process\n\n\n    def prepare_data(files, num_workers, verbose=1, norm=False, dim=(128,256), duration=5,\n                     snipet=False, spec=None, bits=8):\n        \"\"\" generate specs for submission in parallel \"\"\"\n        transform_to_spec_ = transform_to_spec(\n            dim=dim, norm=norm, duration=duration, snipet=snipet, spec=spec, bits=bits)\n        len_duration = duration * 32_000\n        \n        def _process_chunk(chunk):\n            out = []\n            for file in tqdm(chunk):\n                audio, sr = torchaudio.load(file)\n                audio = audio[:, :len_audio_1m].view(-1, len_duration)\n                out.append(transform_to_spec_(audio))\n            out = torch.cat(out, axis=0)\n            out = out.unsqueeze(1)\n            return out\n\n        t1 = time.time()\n        chunk_size = math.ceil(len(files) / num_workers)\n        chunks = [files[i : i+chunk_size] for i in range(0, len(files), chunk_size)]\n\n        with concurrent.futures.ThreadPoolExecutor(max_workers=num_workers) as executor:\n            out = list(executor.map(_process_chunk, chunks))\n        out = torch.cat(out, axis=0)\n        duration = time.time() - t1\n        if verbose:\n            print(f'specs: {out.shape}')\n            print(f'time: {duration:,.0f} s for {len(files):,.0f} files  -  all: {duration * 700 / len(files) / 60:,.0f} m')\n        return out\n\n\n    class Cache_audio_files:\n        \"\"\" keeps a cache of a random `duration` segment of unlabeled audio files. \n            Refreshes segment every `max_usage` times it is used. \n        \"\"\"\n        \n        def __init__(self, files, duration=5, max_usage=None, size=1000):\n            self.files = files\n            random.shuffle(self.files)\n            self.n_loaded = 0\n            self.duration = duration * 32_000\n            self.max_usage = max_usage or 6\n            self.size = size or 1000\n            self.cache = dict()\n            self.count_used = np.zeros(size)\n            \n        def get_random(self):\n            idx = random.randint(0, self.size - 1)\n            cache = self.cache\n            if not idx in cache or self.count_used[idx] > self.max_usage:\n                file = self.files[self.n_loaded]\n                self.n_loaded += 1\n                if self.n_loaded == len(self.files):\n                    random.shuffle(self.files)\n                    self.n_loaded = 0\n                audio, sr = torchaudio.load(file)\n\n                self.cache[idx] = audio\n\n            self.count_used[idx] = 1\n            audio = self.cache[idx]\n            \n            # adjust size\n            gap = self.duration - audio.shape[-1]\n            if gap > 0:\n                pad_before = random.randint(0, gap - 1)\n                pad_after = gap - pad_before\n                audio = F.pad(audio, pad=(pad_before, pad_after), mode='constant', value=0)\n            elif gap < 0:\n                start = random.randint(0, -gap - 1)\n                audio = audio[:, start : start + self.duration]\n            return audio\n\n\n    class Ba_dataset_val(Dataset):\n        \"\"\" Uses prebuilt specs and labels\n        \"\"\"\n        def __init__(self, x_file, y_file, idx_birds=None, **kwargs):\n            super().__init__()\n            self.x_file = x_file\n            self.y_file = y_file\n            self.x = load_obj(x_file)\n            self.y = load_obj(y_file)\n            if idx_birds is not None:\n                self.y = self.y[:, idx_birds]\n            #self.ws = torch.ones(x.shape[0], dtype=torch.float32) / x.shape[0]\n            self._tensorize_and_pin()\n\n        def __len__(self):\n            return len(self.y)\n\n        def __getitem__(self, index):\n            x = self.x[index]\n            y = self.y[index]\n            return dict(idx=index, x=x, y=y)\n\n        def _tensorize_and_pin(self):\n            if torch.cuda.is_available():\n                self.x = torch.tensor(self.x).pin_memory()\n                self.y = torch.tensor(self.y).pin_memory()\n                # self.ws = torch.tensor(self.ws).pin_memory()\n            else:\n                self.x = torch.tensor(self.x)\n                self.y = torch.tensor(self.y)\n\n\n    def pick_audio_section(desc, train_mode):\n        \"\"\" pick a section at random from those included in `desc`, which contains:\n            - file, bird, curation, sections, sections_prob, soft_labels, hard_labels, mask \n        \"\"\"\n        if len(desc.sections) == 1:\n            return 0\n        if train_mode == '2x7s':\n            return 0 if random.random() < .5 else len(desc.sections) - 1\n        return random.choices(range(len(desc.sections)), weights=desc.sections_prob, k=1)[0]\n        \n        \n    class Ba_preload_audio:\n        \"\"\" preloads audios in `catalog` into a dict of shared tensors. \"\"\"\n        \n        def __init__(self, catalog, species):\n            self.catalog = dict()\n            self.bird_ids = defaultdict(list)\n            self.width = 5\n            \n            # get candidates\n            n_preloaded = defaultdict(int)\n            n_unlabeled = defaultdict(int)\n            candidates = []\n            for desc in catalog:\n                if desc.primary_label in species:\n                    bird = label_to_num[desc.primary_label]\n                    if desc.id_[0] not in ['H', 'O']:\n                        self.catalog[desc.id_] = desc\n                        start, end = desc.sections[0][0], desc.sections[-1][1]\n                        audio_start, audio_end = int(start * SR), int(end * SR)\n                        candidates.append((desc.id_, desc.primary_label, desc.file, audio_start, audio_end))\n                        bird = label_to_num[desc.primary_label]\n                        self.bird_ids[bird].append(desc.id_)\n                        n_preloaded[bird] += 1\n                    else:\n                        n_unlabeled[bird] += 1\n                        \n            # load files\n            self.audios = dict()\n            for id_, primary_label, file, audio_start, audio_end in tqdm(candidates):\n                with h5py.File(file, 'r') as f:\n                    audio = f['raw'][audio_start : audio_end]\n                audio = torch.as_tensor(audio)\n                audio.share_memory_()\n                self.audios[id_] = audio\n                \n            # fill share_preloaded\n            share_preloaded = dict()\n            for bird in n_preloaded:\n                share_preloaded[bird] = n_preloaded[bird] / (n_preloaded[bird] + n_unlabeled[bird])\n            self.share_preloaded = share_preloaded\n\n\n        def exists(self, id_):\n            \"\"\" return full `id` audio as a 1D tensor or None if not loaded \"\"\"\n            return id_ in self.audios\n\n                \n        def get_id_audio(self, id_):\n            \"\"\" return full `id` audio as a 1D tensor or None if not loaded \"\"\"\n            return self.audios.get(id_, None)\n\n\n        def get_id_segment(self, id_):\n            \"\"\" returns random 5s segment for `id_` audio as tensor or None if not loaded \"\"\"\n            audio = self.audios.get(id_, None)\n            if audio is None:\n                return None\n            desc = self.catalog[id_]\n            idx_section = random.choices(range(len(desc.sections_prob)), weights=desc.sections_prob, k=1)[0]\n            section_start, section_end = desc.sections[idx_section]\n            max_offset = max(section_end - section_start - self.width, 0)\n            offset = random.uniform(0, max_offset)\n            offset = int(offset * SR)\n            segment = audio[offset : offset + self.width*SR]\n            \n            # pad if needed\n            gap = self.width*SR - segment.shape[-1]\n            if gap > 0:\n                pad_before = random.randint(0, gap - 1)\n                pad_after = gap - pad_before\n                segment = F.pad(segment, pad=(pad_before, pad_after), mode='constant', value=0)\n\n            label = desc.hard_labels[idx_section]\n            mask = desc.mask[idx_section]\n            return segment, label, mask\n\n\n        def get_bird_segment(self, bird):\n            \"\"\" returns random 5s segment from a random audio for `bird` as tensor or None if not loaded.\n                return also the corresponding label and mask.\n            \"\"\"\n            ids = self.bird_ids.get(bird, None)\n            if ids is None:\n                return None, None, None\n            id_ = random.choice(ids)\n            return self.get_id_segment(id_)\n\n                            \n    class Ba_data_module_montage_catalog_h5(L.LightningDataModule):\n        def __init__(self, catalog, preloaded_audio, train_idx, val_idx, dataset_val, fold, ws=None, train_size=None,\n                     files_background_noise=None, montage_bird_prob=None, train_mode='random'):\n            super().__init__()\n            self.catalog = catalog\n            self.preloaded_audio = preloaded_audio\n            self.birds = np.array([label_to_num[x.primary_label] for x in catalog])\n            self.ws = ws\n            self.train_idx = train_idx\n            self.val_idx = val_idx\n            self.dataset_val = dataset_val\n            self.fold = fold\n            self.files_background_noise = files_background_noise\n            self.prefetch_factor = config.prefetch_factor\n            self.train_size = train_size or len(files)\n            self.montage_bird_prob = montage_bird_prob\n            self.train_mode = train_mode\n      \n        def prepare_data(self):\n            \"\"\" processing to be done only in one GPU \"\"\"\n            train_idx = self.train_idx\n            val_idx = self.val_idx\n            train_catalog = [x for x, is_train in zip(self.catalog, train_idx) if is_train]\n            self.train_dataset = Ba_dataset_montage_catalog_unlabeled_h5(\n                train_catalog, self.preloaded_audio, self.birds[train_idx], shuffle=True, ws=self.ws,\n                augment=True, cache=False, montage_cache_size=config.montage_cache_size,\n                files_background_noise=self.files_background_noise, train_size=self.train_size,\n                montage_bird_prob=self.montage_bird_prob, train_mode=self.train_mode)\n            folder = f'{ROOT}/b5-data-{self.dataset_val}'\n            if self.val_idx is not None:\n                self.val_dataset = Ba_dataset_val(f'{folder}/x_f{self.fold}', f'{folder}/y_f{self.fold}')\n                size = self.val_idx.sum()\n                if size != self.val_dataset.x.shape[0]:\n                    self.val_dataset.x = self.val_dataset.x[:size]\n                    self.val_dataset.y = self.val_dataset.y[:size]\n      \n        def setup(self, stage=None):\n            \"\"\" processing to be done in all GPUs \"\"\"\n            pass\n      \n        def train_dataloader(self):\n            return DataLoader(\n                self.train_dataset, batch_size=config.train_batch_size, shuffle=False, \n                num_workers=config.num_workers, prefetch_factor=self.prefetch_factor,\n                collate_fn=None, drop_last=True)\n\n        def val_dataloader(self):\n            if self.val_idx is None:\n                return None\n            return DataLoader(\n                self.val_dataset, batch_size=config.val_batch_size, shuffle=False, \n                num_workers=config.num_workers, prefetch_factor=self.prefetch_factor,\n                collate_fn=None, drop_last=False)\n      \n        def test_dataloader(self):\n            return None\n\n\n    class Ba_dataset_montage_catalog_unlabeled_h5(Dataset):\n        \n        def __init__(self, catalog, preloaded_audio, birds, shuffle=True, ws=None, augment=None,\n                     cache=False, montage_cache_size=10, files_background_noise=None,\n                     train_size=None, montage_bird_prob=None, train_mode='random', **kwargs):\n            super().__init__()\n            \n            self.preloaded_audio = preloaded_audio\n            self.birds = birds\n            self.shuffle = shuffle\n            self.ws = ws / ws.sum()\n            self.mode = mode\n            self.train_size = train_size\n            self.montage_bird_prob = montage_bird_prob if montage_bird_prob is not None else ([1]*N_BIRDS)\n            self.train_mode = train_mode\n\n            self.montage_unlabeled_prob = config.montage_unlabeled_prob\n            self.catalog = catalog\n\n            self.augment = augment_img_train() if augment else None\n            if config.background_noise_prob and files_background_noise is not None:\n                self.background_noise_cache = Cache_audio_files(\n                    files_background_noise, config.img_duration, config.background_noise_max_usage, config.background_noise_cache_size)\n                self.background_noise_prob = config.background_noise_prob\n            else:\n                self.background_noise_prob = 0\n            self.aug_volume = (math.log(config.aug.volume[0]), math.log(config.aug.volume[1])) \\\n                if config.aug.volume is not None else (math.log(.5), math.log(2.))\n            self.background_noise_reference = config.background_noise_reference\n            self.norm_audio = config.norm_audio or False\n\n            self.label_smoothing = config.label_smoothing\n            self.montage = config.montage if config.montage is not None else (0, 0)\n            self.montage_cache_size = montage_cache_size\n            self.montage_cache = {bird: [] for bird in RANGE_BIRDS}\n            self.snipet = config.snipet if config.snipet is not None else False\n            self.width = config.img_duration\n            self.idxs = np.arange(len(self.catalog))\n            bits = config.bits or 8\n            self.transform_to_spec_ = transform_to_spec(\n                dim=config.img_dim, norm=config.spec_norm, duration=config.img_duration,\n                snipet=self.snipet, spec=config.spec, bits=bits)\n            self.on_epoch_end()\n            self.all_true = torch.ones(N_BIRDS, dtype=torch.float32)\n            \n            \n        def __len__(self):\n            return self.train_size\n\n\n        def __getitem__(self, index):\n            # pick primary audio\n            idx = self.idxs[index]  # idxs is already randomized before each epoch\n            x, y, mask = self._get_audio(idx)\n            x, y, mask = self._get_montage(x, y, mask, idx)\n            if random.random() < self.background_noise_prob:\n                background_noise_audio = self.background_noise_cache.get_random()\n                if self.background_noise_reference:\n                    max_x = x.abs().max()\n                    max_background_noise = background_noise_audio.abs().max()\n                    x = x / max_x * max_background_noise + background_noise_audio\n                else:\n                    x = x + background_noise_audio * math.exp(random.uniform(*self.aug_volume))\n                mask = (y > 0).to(torch.float32)\n            else:\n                mask = self.all_true\n                \n            y = torch.clamp(y, 0., 1.)\n            if self.label_smoothing:\n                tp = y.sum(-1)\n                y = y * (1 - self.label_smoothing) + self.label_smoothing * tp / y.shape[-1]\n\n            if self.norm_audio:\n                x /= torch.abs(x).max()\n\n            x = self.transform_to_spec_(x)\n\n            if self.augment:\n                dtype = x.dtype\n                x = self.augment(image=x.permute(1,2,0).numpy())['image'].astype(np.float32)\n                x = torch.tensor(x, dtype=dtype).permute(2,0,1)\n\n            return dict(idx=idx, x=x, y=y, mask=mask)\n\n\n        def on_epoch_end(self):\n            if self.shuffle and self.ws is not None:\n                n = len(self.catalog)\n                self.idxs = np.random.choice(np.arange(n), size=self.train_size, replace=True, p=self.ws)\n\n\n        def _get_audio(self, idx):\n            desc = self.catalog[idx]\n            width = self.width\n            width_sr = width * SR\n            bird = label_to_num[desc.primary_label]\n            \n            if desc.id_[0] not in ['H', 'O']:\n                x, y, mask = self.preloaded_audio.get_bird_segment(bird)\n            else:\n                x = None\n            \n            if x is None:\n                # pick segment start and read it from file\n                section_idx = pick_audio_section(desc, self.train_mode)\n                section_start, section_end = desc.sections[section_idx]\n                if self.train_mode == '2x7s':\n                    if section_idx == 0:\n                        if len(desc.sections) > 1 or desc.sections_prob[0] <= 7:  \n                            # multiple sections and we already picked first or \n                            # single section <= 7s\n                            section_end = min(section_end, section_start + 7)\n                        else:  # single section; select firt or last 7s\n                            if random.random() < .5: # pick start of section\n                                section_end = min(section_end, section_start + 7)\n                            else:  # pick end of section\n                                section_start = max(section_start, section_end - 7)\n                    else:\n                        section_start = max(section_start, section_end - 7)\n                \n                # pick a random offset within segment\n                max_offset = max(section_end - section_start - width, 0)\n                offset = random.uniform(0, max_offset)\n                # pick montage offset to be within 2s of montage to minimize size of audio file to load\n                if self.montage[1] > 0:\n                    montage_min_offset = max(0, offset - 2)\n                    montage_max_offset = min(offset + 2, max_offset)\n                    offset_montage = random.uniform(montage_min_offset, montage_max_offset)\n                else:\n                    offset_montage = offset\n                    \n                min_offset = min(offset, offset_montage)\n                max_offset = max(offset, offset_montage)\n                audio_start = int((section_start + min_offset) * SR)\n                audio_end = int((section_start + width + max_offset) * SR)\n                start_montage = int((offset_montage - min_offset) * SR)\n                    \n                with h5py.File(desc.file, 'r') as f:\n                    audio = f['raw'][audio_start : audio_end]\n                audio = torch.as_tensor(audio)\n                \n                # add montage\n                if audio.shape[-1] < 5 * SR and start_montage > 0:\n                    print('######', idx, desc.file)\n                    print(section_start, section_end, audio_start/SR, audio_end/SR, audio.shape, audio.shape[-1]/SR)\n                    print(audio_start, audio_end, offset, offset_montage, start_montage/SR)\n                    assert 1 == 0\n\n                self._add_montage_cache(audio.unsqueeze(0), start_montage, idx, section_idx)\n                start = int((offset - min_offset) * SR)\n                x = audio[start : start + width_sr]\n                y = desc.hard_labels[section_idx]\n                mask = desc.mask[section_idx]\n                \n                gap = width_sr - x.shape[-1]\n                if gap > 0:\n                    pad_before = random.randint(0, gap - 1)\n                    pad_after = gap - pad_before\n                    x = F.pad(x, pad=(pad_before, pad_after), mode='constant', value=0)\n            \n            return x.unsqueeze(0), y, mask\n\n\n        def _get_montage(self, x, y, mask, idx):\n            \"\"\" build montage. Applies a random volume change to all signals, including the original one \"\"\"\n            aug_volume = self.aug_volume\n            x = x * math.exp(random.uniform(*aug_volume))\n            bird = self.birds[idx]\n            n = random.randint(*self.montage)\n            sample = random.choices(RANGE_BIRDS, weights=self.montage_bird_prob, k=n)\n            # sample = random.sample(RANGE_BIRDS, n)\n            # sample = [x for x in sample if x != bird]\n            if len(sample) > 0:\n                preloaded_audio = self.preloaded_audio\n                share_preloaded = preloaded_audio.share_preloaded\n                x = x.clone()\n                y = y.clone()\n                mask = mask.clone()\n                for i in sample:\n                    if i in preloaded_audio.bird_ids and random.random() <= share_preloaded[i]:\n                        mx, my, mmask = preloaded_audio.get_bird_segment(i)\n                    else:\n                        mx, my, mmask = self._get_montage_cache(i)\n                    if mx is not None:\n                        x += mx * math.exp(random.uniform(*aug_volume))\n                        y += my\n                        mask &= mmask\n            return x, y, mask\n\n\n        def _add_montage_cache(self, audio, offset, idx, section_idx):\n            \"\"\" add bird instance to cache \"\"\"\n            if self.montage_unlabeled_prob:  # montage uses only unlabeled audios loaded in real time\n                return\n            desc = self.catalog[idx]\n            if self.preloaded_audio.exists(desc.id_):\n                return  # no need to cache preloaded audios\n            width_sr = self.width * SR\n            x = audio[..., offset : offset + width_sr]\n            gap = width_sr - x.shape[-1]\n            if gap > 0:\n                pad_before = random.randint(0, gap - 1)\n                pad_after = gap - pad_before\n                x = F.pad(x, pad=(pad_before, pad_after), mode='constant', value=0)\n            \n            bird = label_to_num[desc.primary_label]\n            y = desc.hard_labels[section_idx]\n            mask = desc.mask[section_idx]\n            data = (x, y, mask)\n            if len(self.montage_cache[bird]) <= self.montage_cache_size:\n                self.montage_cache[bird].append(data)\n            else:\n                idx_cache = random.randint(0, self.montage_cache_size - 1)\n                self.montage_cache[bird][idx_cache] = data\n\n\n        def _get_montage_cache(self, bird):\n            \"\"\" get random instance of bird from those cached \"\"\"\n            options = self.montage_cache[bird]\n            if len(options) == 0:\n                return None, None, None\n            else:\n                idx_cache = random.randint(0, len(options) - 1)\n                return self.montage_cache[bird][idx_cache]\n\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.6 Loss & metrics\n    # ----------------------------------------------------------------------------------------------------\n\n    def sigmoid(x):\n        return np.where(x >= 0, 1 / (1 + np.exp(-x)), np.exp(x) / (1 + np.exp(x)))\n\n\n    def get_loss_fn(name):\n        reduction = 'none' if config.use_mask or config.mask_outliers else 'mean'\n        if name == 'focal_volodymyr':\n            return FocalLossBCEVolodymyr(reduction=reduction)\n\n\n    def auc_multi(preds, labels):\n        \"\"\" Version of macro-averaged ROC-AUC score that ignores all classes that have no true positive labels. \"\"\"\n        scored_cols = labels.sum(axis=0) > 0\n        return sklearn.metrics.roc_auc_score(\n            labels[:, scored_cols], preds[:, scored_cols], average='macro')\n\n\n    class FocalLossBCEVolodymyr(torch.nn.Module):\n        def __init__(\n            self,\n            alpha: float = 0.25,\n            gamma: float = 2,\n            reduction: str = 'mean',\n            bce_weight: float = 1.0,\n            focal_weight: float = 1.0,\n        ):\n            super().__init__()\n            self.alpha = alpha\n            self.gamma = gamma\n            self.reduction = reduction\n            self.bce = torch.nn.BCEWithLogitsLoss(reduction=reduction)\n            self.bce_weight = bce_weight\n            self.focal_weight = focal_weight\n\n        def forward(self, inputs, targets):\n            focall_loss = torchvision.ops.focal_loss.sigmoid_focal_loss(\n                inputs=inputs,\n                targets=targets,\n                alpha=self.alpha,\n                gamma=self.gamma,\n                reduction=self.reduction,\n            )\n            bce_loss = self.bce(inputs, targets)\n            return self.bce_weight * bce_loss + self.focal_weight * focall_loss\n            \n    # ----------------------------------------------------------------------------------------------------\n    # 0.7 Optimizers & schedulers\n    # ----------------------------------------------------------------------------------------------------\n\n\n    def fetch_optimizer(model, cfg):\n        no_decay = [\"bias\", \"LayerNorm.bias\", \"LayerNorm.weight\"]\n        param_optimizer = list(model.named_parameters())\n\n        param_optimizer_encoder = [x for x in param_optimizer if x[0].startswith('model.')]\n        param_optimizer_head = [x for x in param_optimizer if not x[0].startswith('model.')]\n        optimizer_parameters = [\n            {'params': [p for n, p in param_optimizer_encoder if not any(nd in n for nd in no_decay)],\n             'weight_decay': cfg.weight_decay, 'lr': cfg.lr[0][0]},\n            {'params': [p for n, p in param_optimizer_encoder if any(nd in n for nd in no_decay)],\n             'weight_decay': 0.0, 'lr': cfg.lr[0]},\n            {'params': [p for n, p in param_optimizer_head if not any(nd in n for nd in no_decay)],\n             'weight_decay': cfg.weight_decay, 'lr': cfg.lr[0][1]},\n            {'params': [p for n, p in param_optimizer_head if any(nd in n for nd in no_decay)],\n             'weight_decay': 0.0, 'lr': cfg.lr[1]},\n        ]\n\n        if cfg.optimizer == 'AdamW':\n            optimizer = torch.optim.AdamW(optimizer_parameters, eps=cfg.eps, betas=cfg.betas)\n\n        if cfg.previous_step is not None and cfg.previous_step != -1 and not optimizer_loaded:\n            for group in optimizer.param_groups:\n                group.setdefault('initial_lr', group['lr'])\n        return optimizer\n\n\n    def fetch_scheduler(cfg, optimizer):\n        previous_step = cfg.previous_step or -1\n        \n        if cfg.schedule_type == 'multi_lr':  # lr = [(1e-5, 1e-3), (3, 2.5e-4)]\n            pct_start = cfg.num_warmup_steps / cfg.num_train_steps\n            # scheduler 1: split encoder / head\n            lr1 = cfg.lr[0]\n            div_factor = 25 if cfg.warmup_share else 1\n            final_div_factor = 25\n            max_lr = (lr1[0], lr1[0], lr1[1], lr1[1])\n            scheduler_1 = torch.optim.lr_scheduler.OneCycleLR(\n                optimizer, anneal_strategy='cos', last_epoch=previous_step,\n                total_steps=cfg.num_train_steps, pct_start=pct_start,\n                max_lr=max_lr, div_factor=div_factor, final_div_factor=final_div_factor)\n            \n            # scheduler 2: unique lr\n            lr2_epoch, lr2 = cfg.lr[1]\n            lr2_step = int(lr2_epoch * config.train_size / (config.train_batch_size))\n            start_lr = cfg.start_lr or (lr2 / 25 if cfg.warmup_share else lr2)\n            end_lr = cfg.end_lr or lr2 / 25\n            div_factor = 25 if cfg.warmup_share else 1\n            final_div_factor = lr2 / end_lr\n            max_lr = lr2\n            scheduler_2 = OneCycleLRCustom(\n                optimizer, anneal_strategy='cos', last_epoch=lr2_step - 1,\n                total_steps=cfg.num_train_steps - lr2_step,\n                max_lr=max_lr, div_factor=div_factor, final_div_factor=final_div_factor)\n            \n            # combined scheduler\n            scheduler = torch.optim.lr_scheduler.SequentialLR(\n                optimizer, schedulers=[scheduler_1, scheduler_2],\n                milestones=[lr2_step], last_epoch=-1)\n        \n        # restore state if applicable\n        if cfg.saved_scheduler:\n            filename = cfg.saved_scheduler\n            if os.path.exists(filename):\n                print(f\"Loading scheduler '{filename}'\")\n                try:\n                    checkpoint = torch.load(filename)\n                except RuntimeError:\n                    checkpoint = torch.load(filename, map_location=torch.device('cpu'))\n                scheduler.load_state_dict(checkpoint['lr_schedulers'][-1])\n            else:\n                print('scheduler not found.')\n        return scheduler\n\n\n    class OneCycleLRCustom(torch.optim.lr_scheduler._LRScheduler):\n        \"\"\" similar to default OneCycleLR but uses it's own lr rather than the optimizers.\n            To use as second scheduler in a SequentialLR.\n        \"\"\"\n\n        def __init__(self, optimizer, max_lr, total_steps=None, epochs=None, steps_per_epoch=None,\n                     pct_start = 0.3, anneal_strategy=\"cos\", cycle_momentum=True, base_momentum=0.85,\n                     max_momentum=0.95, div_factor=25.0, final_div_factor=1e4, three_phase=False,\n                     last_epoch=-1, verbose=\"deprecated\"):\n            # Validate optimizer\n            # if not isinstance(optimizer, Optimizer):\n            #     raise TypeError(f\"{type(optimizer).__name__} is not an Optimizer\")\n            self.optimizer = optimizer\n\n            # Validate total_steps\n            if total_steps is not None:\n                if total_steps <= 0 or not isinstance(total_steps, int):\n                    raise ValueError(\n                        f\"Expected positive integer total_steps, but got {total_steps}\"\n                    )\n                self.total_steps = total_steps\n            elif epochs is not None and steps_per_epoch is not None:\n                if not isinstance(epochs, int) or epochs <= 0:\n                    raise ValueError(f\"Expected positive integer epochs, but got {epochs}\")\n                if not isinstance(steps_per_epoch, int) or steps_per_epoch <= 0:\n                    raise ValueError(\n                        f\"Expected positive integer steps_per_epoch, but got {steps_per_epoch}\"\n                    )\n                self.total_steps = epochs * steps_per_epoch\n            else:\n                raise ValueError(\n                    \"You must define either total_steps OR (epochs AND steps_per_epoch)\"\n                )\n\n            self._schedule_phases: List[_SchedulePhase]\n            self._schedule_phases = [{\n                \"end_step\": self.total_steps - 1,\n                \"start_lr\": \"max_lr\",\n                \"end_lr\": \"min_lr\",\n                \"start_momentum\": \"base_momentum\",\n                \"end_momentum\": \"max_momentum\"}]\n\n            # Validate anneal_strategy\n            if anneal_strategy not in [\"cos\", \"linear\"]:\n                raise ValueError(f\"anneal_strategy must be one of 'cos' or 'linear', instead got {anneal_strategy}\")\n            else:\n                self._anneal_func_type = anneal_strategy\n\n            # Initialize learning rate variables\n            max_lrs = max_lr\n            if last_epoch == -1:\n                for idx, group in enumerate(self.optimizer.param_groups):\n                    group[\"initial_lr\"] = max_lrs[idx] / div_factor\n                    group[\"max_lr\"] = max_lrs[idx]\n                    group[\"min_lr\"] = group[\"initial_lr\"] / final_div_factor\n\n            # Initialize momentum variables\n            self.cycle_momentum = cycle_momentum\n            if self.cycle_momentum:\n                if (\"momentum\" not in self.optimizer.defaults and \"betas\" not in self.optimizer.defaults):\n                    raise ValueError(\"optimizer must support momentum or beta1 with `cycle_momentum` option enabled\")\n                self.use_beta1 = \"betas\" in self.optimizer.defaults\n                max_momentums = torch.optim.lr_scheduler._format_param(\"max_momentum\", optimizer, max_momentum)\n                base_momentums = torch.optim.lr_scheduler._format_param(\"base_momentum\", optimizer, base_momentum)\n                if last_epoch == -1:\n                    for m_momentum, b_momentum, group in zip(max_momentums, base_momentums, optimizer.param_groups):\n                        if self.use_beta1:\n                            group[\"betas\"] = (m_momentum, *group[\"betas\"][1:])\n                        else:\n                            group[\"momentum\"] = m_momentum\n                        group[\"max_momentum\"] = m_momentum\n                        group[\"base_momentum\"] = b_momentum\n\n            self.start_lr = max_lr\n            self.end_lr = max_lr / final_div_factor\n            super().__init__(optimizer, last_epoch, verbose)\n\n        def _anneal_func(self, *args, **kwargs):\n            if hasattr(self, \"_anneal_func_type\"):\n                if self._anneal_func_type == \"cos\":\n                    return self._annealing_cos(*args, **kwargs)\n                elif self._anneal_func_type == \"linear\":\n                    return self._annealing_linear(*args, **kwargs)\n                else:\n                    raise ValueError(f\"Unknown _anneal_func_type: {self._anneal_func_type}\")\n            else:\n                # For BC\n                return self.anneal_func(*args, **kwargs)  # type: ignore[attr-defined]\n\n        @staticmethod\n        def _annealing_cos(start, end, pct):\n            \"\"\"Cosine anneal from `start` to `end` as pct goes from 0.0 to 1.0.\"\"\"\n            cos_out = math.cos(math.pi * pct) + 1\n            return end + (start - end) / 2.0 * cos_out\n\n        @staticmethod\n        def _annealing_linear(start, end, pct):\n            \"\"\"Linearly anneal from `start` to `end` as pct goes from 0.0 to 1.0.\"\"\"\n            return (end - start) * pct + start\n\n        def get_lr(self):\n            \"\"\"Compute the learning rate of each parameter group.\"\"\"\n            torch.optim.lr_scheduler._warn_get_lr_called_within_step(self)\n\n            lrs = []\n            step_num = self.last_epoch\n\n            if step_num > self.total_steps:\n                raise ValueError(\n                    f\"Tried to step {step_num} times. The specified number of total steps is {self.total_steps}\"  # noqa: UP032\n                )\n\n            for group in self.optimizer.param_groups:\n                start_step = 0.0\n                for i, phase in enumerate(self._schedule_phases):\n                    end_step = phase[\"end_step\"]\n                    if step_num <= end_step or i == len(self._schedule_phases) - 1:\n                        pct = (step_num - start_step) / (end_step - start_step)\n                        computed_lr = self._anneal_func(self.start_lr, self.end_lr, pct)\n                        if self.cycle_momentum:\n                            computed_momentum = self._anneal_func(\n                                group[phase[\"start_momentum\"]],\n                                group[phase[\"end_momentum\"]],\n                                pct,\n                            )\n                        break\n                    start_step = phase[\"end_step\"]\n\n                lrs.append(computed_lr)  # type: ignore[possibly-undefined]\n                if self.cycle_momentum:\n                    if self.use_beta1:\n                        group[\"betas\"] = (computed_momentum, *group[\"betas\"][1:])  # type: ignore[possibly-undefined]\n                    else:\n                        group[\n                            \"momentum\"\n                        ] = computed_momentum  # type: ignore[possibly-undefined]\n\n            return lrs\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.8 Lightning training\n    # ----------------------------------------------------------------------------------------------------\n\n    def get_amp_context(device_type, datatype):\n        \"\"\" get context for \"\"\"\n        ptdtype = {'float32': torch.float32, 'bfloat16': torch.bfloat16, 'float16': torch.float16}[datatype]\n        device_type = 'cuda' if 'cuda' in str(device_type) else str(device_type)\n        #context = nullcontext() if device_type == 'cpu' else torch.amp.autocast(device_type=device_type, dtype=ptdtype)\n        context = nullcontext() if device_type == 'cpu' else torch.cuda.amp.autocast(dtype=ptdtype)  # for torch v1.11\n        return context\n\n\n    def pb_format_num(n):\n        \"\"\" format numbers in progress bar \"\"\"\n        f = '{0:.4g}'.format(n).replace('+0', '+').replace('-0', '-')\n        n = str(n)\n        return f if len(f) < len(n) else n\n\n\n    class My_progress_bar(TQDMProgressBar):\n        \"\"\" customize TQDMProgressBar to add lr \"\"\"\n        def get_metrics(self, trainer, pl_module):\n            items = super().get_metrics(trainer, pl_module)\n            items.pop('v_num', None)  # added in base class; contains version number\n            items = {k: pb_format_num(v) for k, v in items.items()}\n            return items\n\n        def init_validation_tqdm(self):\n            \"\"\" disable validation progress bar \"\"\"\n            bar = tqdm(disable=True)\n            return bar\n            \n\n    class AverageMeter:\n        def __init__(self, labels=None):\n            self.reset()\n            self.auc = None\n            self.labels = labels if labels is not None else None\n            \n        def reset(self):\n            self.loss_sum = 0\n            self.acc_sum = 0\n            self.count = 0\n            self.preds = []\n        \n        def update(self, loss, acc):\n            self.loss_sum += loss\n            self.acc_sum += acc\n            self.count += 1\n            return loss\n        \n        def update_preds(self, preds):\n            self.preds.append(preds.detach().cpu())\n        \n        @property\n        def loss(self):\n            return self.loss_sum / self.count if self.count > 0 else 0\n\n        @property\n        def acc(self):\n            return self.acc_sum / self.count if self.count > 0 else 0\n        \n        def calculate_auc(self):\n            self.preds = np.concatenate(self.preds, axis=0)\n            len_preds = self.preds.shape[0]  # to support sanity check\n            self.auc = auc_multi(self.preds, self.labels[:len_preds])\n            \n\n    class Pl_model(L.LightningModule):\n        def __init__(self, model, cfg, val_labels=None, silent=None):\n            super().__init__()\n            self.model = model\n            self.cfg = cfg\n            self.use_mask = cfg.use_mask\n            self.automatic_optimization = False\n            self.current_amp_dtype = None\n            self.current_loss_name = None\n            self._setup_amp()\n            self._setup_loss()\n            self.train_meter = AverageMeter()\n            self.val_meter = AverageMeter(val_labels) if val_labels is not None else None\n            self.t0 = time.time()\n\n        def _setup_amp(self):\n            \"\"\" adjust automatic mode precision dtype based on epoch \"\"\"\n            dtype = self.cfg.precision_schedule.get(self.current_epoch, None)\n            if dtype is not None and self.current_amp_dtype != dtype:\n                self.amp_context = get_amp_context(self.device, dtype)\n                self.scaler = torch.cuda.amp.GradScaler(enabled=(dtype == 'float16'),\n                                                        growth_interval=600)\n            assert hasattr(self, 'amp_context'), \"Amp context not set, check precision schedule\"\n            self.current_amp_dtype = dtype\n            os.environ['EPOCH'] = str(self.current_epoch)\n\n        def _setup_loss(self):\n            \"\"\" adjust loss based on epoch \"\"\"\n            loss_name = self.cfg.loss_schedule.get(self.current_epoch, None)\n            if loss_name is not None and self.current_loss_name != loss_name:\n                self.loss_fn = get_loss_fn(loss_name)\n                self.current_loss_name = loss_name\n\n        def _compute_loss_metrics(self, pred, data, meter):\n            y = data['y']\n            loss = self.loss_fn(pred, y)\n            mask = data.get('mask', None)\n            if self.use_mask and mask is not None:\n                loss = (loss * mask).sum() / mask.sum()\n            else:\n                loss = loss.mean()\n            meter.update(loss.item(), 0)\n            return loss\n\n        def on_train_epoch_start(self):\n            self._setup_loss()\n            self.train_meter.reset()\n            if self.val_meter:\n                self.val_meter.reset()\n            self.t0 = time.time()                     \n        \n        def training_step(self, batch, batch_idx):\n            cfg = self.cfg\n            meter = self.train_meter\n            if cfg.aug.mixup:\n                mixup(batch)\n            with self.amp_context:\n                preds = self(batch['x'])\n                loss = self._compute_loss_metrics(preds, batch, meter)\n            out = dict(loss=meter.loss, acc=0)\n            if torch.isnan(loss).any():\n                my_log('ERROR: loss is nan. Interrupting training.')\n                self.trainer.should_stop = True\n            self.log_dict(out, prog_bar=True)\n            self.manual_backward(self.scaler.scale(loss))\n\n            score = self.val_meter.auc if self.val_meter else self.train_meter.loss\n            optimizer = self.optimizers()\n            self.scaler.unscale_(optimizer)\n            if cfg.max_grad_value is not None:\n                torch.nn.utils.clip_grad_value_(self.parameters(), cfg.max_grad_value)\n            elif cfg.max_grad_norm is not None:\n                norm_grad = torch.nn.utils.clip_grad_norm_(parameters=self.parameters(),\n                                                           max_norm=cfg.max_grad_norm)\n            self.scaler.step(optimizer)\n            self.scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n            self.lr_schedulers().step()\n\n        def on_train_epoch_end(self):\n            \"\"\" invoked after validation to process outputs of training_step. \"\"\"\n            ds = self.trainer.train_dataloader.dataset\n            if hasattr(ds, 'on_epoch_end'):\n                ds.on_epoch_end()\n            if self.val_meter is None:\n                loss_metrics_str = f'loss: {self.train_meter.loss:.4f}'\n                scheduler = self.lr_schedulers()\n                my_log(f'epoch:{self.current_epoch:>3} - {loss_metrics_str}')\n            return super().on_train_epoch_end()\n\n        def validation_step(self, batch, batch_idx):\n            \"\"\" no logging per step \"\"\"\n            meter = self.val_meter\n            with self.amp_context:\n                preds = self(batch['x'])\n                self._compute_loss_metrics(preds, batch, meter)\n                meter.update_preds(preds)\n\n        def on_validation_epoch_end(self):\n            \"\"\" process outputs of validation_step \"\"\"\n            self.val_meter.calculate_auc()\n            out = dict(val_loss=self.val_meter.loss, val_auc=self.val_meter.auc)\n            self.log_dict(out, prog_bar=False)\n            loss_metrics_str = f'loss: {self.train_meter.loss:.4f} - ' \\\n                               f'val_loss: {self.val_meter.loss:.4f} - ' \\\n                               f'val_auc: {self.val_meter.auc:.4f}'\n            scheduler = self.lr_schedulers()\n            duration = time.time() - self.t0\n            my_log(f'epoch:{self.current_epoch:>3} - {loss_metrics_str} - {duration:,.0f}s', ' '*40)\n            return super().on_validation_epoch_end()\n\n        def forward(self, batch):\n            return self.model(batch)\n\n        def configure_optimizers(self):\n            cfg = self.cfg\n            optimizer = fetch_optimizer(self.model, cfg)\n            scheduler = fetch_scheduler(self.cfg, optimizer)\n            lr_scheduler_config = {\"scheduler\": scheduler, \"interval\": self.cfg.step_scheduler_after}\n            if cfg.monitor:\n                lr_scheduler_config['monitor'] = cfg.monitor\n            return {\"optimizer\": optimizer, \"lr_scheduler\": lr_scheduler_config}\n\n    \n    def train(meta_raw, catalog_d, no_try=False, debug=None):\n        global config, RANK, my_log\n        \n        RANK = torch.range(1, N_BIRDS).reshape(1, -1).expand(config.train_batch_size, -1)\n        config.img_dim = SPEC_DIMS.get(config.spec, (128, 256))\n        \n        my_log = Log(f'{config.name}.log', time_key=False)\n        logger = get_logger(config)\n        my_log(f'config: {config}\\n')\n        folds = config.folds\n        idx_birds = None\n        montage_bird_prob = None\n            \n        # split per fold\n        meta, catalog = get_meta_folds_catalog(meta_raw, catalog_d)\n\n        for fold in folds:\n            my_log(f'{\"-\"*40} {config.name}_f{fold} {\"-\"*40}')\n            if config.monitor == 'val_auc':\n                train_idx = (meta.fold.values != fold)\n                val_idx = (meta.fold.values == fold)\n                if config.filter[-1] > 0:\n                    train_idx &= meta.rating >= config.filter[-1]\n                val_labels = load_obj(f'{ROOT}/b5-data-{config.dataset_val}/y_f{fold}')\n            else:\n                train_idx, val_idx = np.full(meta.shape[0], True), None\n                val_labels = None\n            if debug:\n                my_log(f'\\n*** DEBUG: {debug} ***\\n')\n                n = meta.shape[0]\n                train_idx = np.array([True] * debug + [False] * (n - debug))\n                np.random.shuffle(train_idx)\n                val_idx = np.array([True] * debug + [False] * (n - debug))\n                np.random.shuffle(val_idx)\n            meta_train = meta[train_idx]\n\n            train_size = config.train_size or int(meta_train.ws.sum())\n            if debug:\n                train_size = debug\n            if val_idx is not None:\n                my_log(f'train: {train_size:,} | {len(meta_train):,} -  val: {val_idx.sum():,}')\n            else:\n                my_log(f'train: {train_size:,} | {len(meta_train):,} - val: n/a')\n            \n            # set ws for train sample\n            ws_power = config.ws_power or 1/2\n            meta_train['counter'] = meta_train.groupby('primary_label')['ws'].transform('sum')\n            ws = meta_train.ws.values / np.power(meta_train.counter.values, ws_power)\n            times_picked = ws * meta_train.counter.values\n            my_log(f'raw weights: {times_picked.min():.2f} - {times_picked.max():.2f}')\n            ws /= ws.sum()\n            \n            epochs = config.epochs\n            config.num_train_steps = int(train_size / (config.train_batch_size) * epochs)\n            config.num_warmup_steps = int(config.num_train_steps * config.warmup_share)\n            my_log(f'train steps: {config.num_train_steps:,} - warmup steps: {config.num_warmup_steps:,}')\n\n            count_audios = meta_train[meta_train.sample_ == 't'].groupby('primary_label')['ws'].count()\n            n_preload_species = config.n_preload_species\n            preload_species = set(count_audios[count_audios <= n_preload_species].index)\n            my_log(f'Preload {len(preload_species)} species')\n            preloaded_audio = Ba_preload_audio(catalog, preload_species)\n            data_module = Ba_data_module_montage_catalog_h5(\n                catalog, preloaded_audio, train_idx, val_idx, config.dataset_val, fold, ws=ws, train_size=train_size,\n                files_background_noise=files_background_noise if config.background_noise_prob else None,\n                montage_bird_prob=montage_bird_prob, train_mode=config.train_mode)\n\n            model_class = globals()[config.model_name]\n            pretrained = config.pretrained_encoder if config.pretrained_encoder is not None else True\n            in_chans = config.in_chans or 3\n            model = model_class(\n                config, resize=config.resize, add_position=config.add_position, num_classes=N_BIRDS,\n                out_indices=config.out_indices, pretrained=pretrained, in_chans=in_chans)\n            inp_dim = (1, config.img_dim[0], config.img_dim[1] * config.img_duration // 5)\n            example_input_array = [torch.zeros(8, *inp_dim)]\n            out_ = model(example_input_array[0])  # sanity check\n            my_log(f'model out shape: {out_.shape}')\n\n            if config.checkpoint:\n                my_log('load existing model')\n                try:\n                    chk = torch.load(config.checkpoint)\n                except:\n                    chk = torch.load(config.checkpoint, map_location=torch.device('cpu'))\n                \n                # model_state = {k.replace('model.', ''): v for k, v in chk['state_dict'].items()}\n                model_state = {k[6:]: v for k, v in chk['state_dict'].items()}\n                model.load_state_dict(model_state)\n\n            pl_model = Pl_model(model, config, val_labels=val_labels)\n            pl_model.example_input_array = example_input_array\n            # display(summary(model, col_names=[\"num_params\",\"trainable\"]))  # show layer sizes\n            # model  # show dimensions\n\n            loss_file = '{epoch}-{val_loss:.4f}'\n            loss_monitor = 'val_loss'\n            metric_file = '{epoch}-{val_auc:.4f}'\n            metric_monitor = 'val_auc'\n            chk_callback_metric = ModelCheckpoint(\n                dirpath='./', filename=f'{config.name}_f{fold}-{metric_file}',\n                save_top_k=3, monitor=metric_monitor, mode=config.mode, \n                save_last=True, save_on_train_epoch_end=True, save_weights_only=True)\n            chk_callback_loss = ModelCheckpoint(\n                dirpath='./', filename=f'{config.name}_f{fold}-{loss_file}',\n                save_top_k=3, monitor=loss_monitor, mode='min', \n                save_last=False, save_on_train_epoch_end=True, save_weights_only=True)\n            callbacks = [chk_callback_metric, chk_callback_loss, \n                         LearningRateMonitor(logging_interval='step'),\n                         My_progress_bar(refresh_rate=config.log_every_n_steps),\n                        ]\n                \n            max_epochs = config.max_epochs if config.max_epochs is not None else epochs\n            num_sanity_val_steps = 1\n            limit_val_batches = None\n            trainer = L.Trainer(max_epochs=max_epochs, max_time=config.duration, logger=logger,\n                                log_every_n_steps=config.log_every_n_steps,\n                                check_val_every_n_epoch=config.check_val_every_n_epoch,\n                                accelerator='auto', devices='auto',\n                                num_sanity_val_steps=num_sanity_val_steps,\n                                limit_train_batches=None, limit_val_batches=limit_val_batches,\n                                enable_checkpointing=True, callbacks=callbacks,\n                                # gradient_clip_val=1.0,\n                               )\n            my_log('start training...')\n            trainer.fit(pl_model, data_module)\n\n        my_log('Finished training...')\n        del data_module\n        gc.collect()\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.11 dummy lightning logger\n    # ----------------------------------------------------------------------------------------------------\n\n    from pytorch_lightning.loggers.logger import Logger, rank_zero_experiment\n    from pytorch_lightning.utilities import rank_zero_only\n\n    class Dummy_logger(Logger):\n        \"\"\" Needed to bypass bug in default logger class.\n            for pytorch_lightning==1.9.4. \n        \"\"\"\n        \n        @property\n        def name(self):\n            return \"MyLogger\"\n\n        @property\n        def version(self):\n            return \"0.1\"\n\n        @rank_zero_only\n        def log_hyperparams(self, params):\n            # params is an argparse.Namespace\n            pass\n\n        @rank_zero_only\n        def log_metrics(self, metrics, step):\n            # metrics is a dictionary of metric names and values\n            pass\n\n        @rank_zero_only\n        def save(self):\n            # Optional. Any code necessary to save logger data\n            pass\n\n        @rank_zero_only\n        def finalize(self, status):\n            # Optional. Any code that needs to be run after training finishes.\n            pass\n\n    def get_logger(config, group='att', job_type='train'):\n        logger = Dummy_logger()\n        return logger\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.13 end\n    # ----------------------------------------------------------------------------------------------------\n\nprint('done')","metadata":{"trusted":true,"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-06-12T17:30:56.893236Z","iopub.execute_input":"2025-06-12T17:30:56.893587Z","iopub.status.idle":"2025-06-12T17:31:13.209860Z","shell.execute_reply.started":"2025-06-12T17:30:56.893537Z","shell.execute_reply":"2025-06-12T17:31:13.209081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 'SKIP' not in globals() or not SKIP:     # features\n    print('features')\n    \n    from typing import Any, Dict, List, Optional, Sequence, Tuple, Union, NamedTuple\n    import copy\n    from contextlib import contextmanager\n    import warnings\n    warnings.filterwarnings('ignore')\n    \n    import os\n    import sys\n    import gc\n    import time\n    import torch\n    import joblib\n    import pickle\n    import psutil\n    import ast\n    import copy\n    import glob\n    import random\n    import pdb\n    import math\n    import numpy as np\n    import pandas as pd\n    from tqdm.autonotebook import tqdm\n    import multiprocessing as mp\n    import concurrent.futures\n    from itertools import groupby as grp_by\n    from datetime import datetime, timedelta\n    import h5py\n\n    from itertools import permutations\n    import math\n    import heapq\n    import re\n\n    from numba import njit\n\n    # scipy\n    from scipy.stats import linregress\n    from scipy.interpolate import interp1d\n    from scipy import signal\n    from scipy.signal import argrelmax\n    from scipy.signal import spectrogram\n    from scipy.signal import find_peaks\n    from scipy.signal import butter, lfilter\n    from scipy import optimize\n    \n    # sklearn\n    from sklearn import linear_model\n    from sklearn.preprocessing import normalize, QuantileTransformer, MinMaxScaler, RobustScaler, normalize, minmax_scale\n    from sklearn.model_selection import train_test_split\n    from sklearn.model_selection import KFold, GroupKFold, GroupShuffleSplit, StratifiedKFold, StratifiedGroupKFold\n    from sklearn.model_selection._split import _BaseKFold\n    from sklearn.metrics import make_scorer, \\\n                                mean_squared_error, mean_absolute_error, r2_score, \\\n                                accuracy_score, classification_report, precision_score, recall_score, f1_score, \\\n                                roc_auc_score, precision_recall_curve, roc_curve, auc, average_precision_score, confusion_matrix\n    from sklearn.decomposition import PCA, TruncatedSVD\n    from sklearn.utils.validation import _num_samples, check_array\n    from sklearn.utils.multiclass import type_of_target\n    from sklearn.utils import check_random_state\n\n    # competition\n    try:\n        import pywt\n    except:\n        pass\n    import librosa\n    \n    COLAB = \"google.colab\" in sys.modules\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.1 utils\n    # ----------------------------------------------------------------------------------------------------\n\n    def wait():\n        for i in range(1000):\n            print('.', end='')\n            time.sleep(60)\n\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.3 meta\n    # ----------------------------------------------------------------------------------------------------\n\n    def read_meta(split_author=True, pivot_folds=False, author_to_num=True):\n        file = f'{ROOT}/birdclef-2025/train.csv'\n        meta = pd.read_csv(file).sort_values(['filename'])\n        meta['i'] = np.arange(meta.shape[0])\n        authors = meta.author.unique()\n        if author_to_num:\n            author_to_num = dict(zip(authors, range(len(authors))))\n            meta['author'] = meta.author.map(author_to_num)\n        \n        # split author in case there is only one per primary_label\n        if split_author:\n            meta['count_unique_author'] = meta.groupby('primary_label')['author'].transform('nunique')\n            meta['order'] = meta.groupby('primary_label').cumcount()\n            meta['author'] = np.where(meta.count_unique_author == 1, meta.author + meta.order / 1000, meta.author)\n\n        meta = meta[['i', 'primary_label', 'secondary_labels', 'author', 'latitude', 'longitude', 'rating', 'filename']]\n        \n        # merge labels\n        labels = []\n        for row in meta.to_dict(orient='records'):\n            row = dotdict(row)\n            secondary = row.secondary_labels[1:-1].replace(\"'\", '')\n            secondary = secondary.split(',') if len(secondary) > 0 else []\n            secondary = [x.replace(' ', '') for x in secondary]\n            labels.append([row.primary_label] + secondary)\n        meta['labels'] = labels\n        meta['n_labels'] = [len(x) for x in labels]\n        meta = meta.drop(['secondary_labels'], axis='columns')\n\n        return meta\n        \n\n    def get_meta(geo_label=0, split_author=False, author_to_num=True):\n        meta = read_meta(split_author=split_author, author_to_num=author_to_num)\n        audios = pd.read_csv(f'{ROOT}/b5-cache/len_train_audios.csv')\n        audios.columns = ['filename', 'dim']\n        meta = meta.merge(audios)\n        return meta\n\n\n    def get_train_add_meta(additional_data_path=None):\n        \"\"\" get merged train meta + meta for additional data \"\"\"\n        meta_train = get_meta(geo_label=0, split_author=False, author_to_num=False)\n        if additional_data_path is not None:\n            meta_add = load_obj(additional_data_path)\n            if config.add_exclude_new is not None:\n                if config.add_exclude_new == 'common':\n                    new_ids = load_obj('/kaggle/input/b5-cache/add new_ids.pkl')\n                    meta_ids = pd.Series([x.split('.')[0] for x in meta_add.filename])\n                    f_selected = (~meta_ids.isin(new_ids)).values | (meta_add.primary_label.isin(RARE_BIRDS)).values\n                    print(f'add files use {f_selected.sum():,} of {len(meta_add):,}')\n                    meta_add = meta_add[f_selected]\n            \n            meta = pd.concat([meta_train, meta_add])\n        meta = meta.drop_duplicates(subset=['filename'], keep='first')\n        meta['source'] = np.where(meta.fold.isnull(), 'b25', 'previous')\n        meta.loc[meta.i.isnull(), 'source'] = 'other'\n\n        # set group (author) for stratification. Split group in case there is only one per primary_label\n        meta['group'] = meta.author.astype('str')\n        meta['count_unique_group'] = meta.groupby('primary_label')['group'].transform('nunique')\n        meta['order'] = meta.groupby('primary_label').cumcount().astype('str')\n        meta['group'] = np.where(meta.count_unique_group == 1, meta.group + meta.order, meta.group)\n        meta = meta.drop(['order'], axis='columns')\n\n        meta['i'] = np.arange(len(meta))\n        meta = meta.drop(['fold'], axis='columns')\n        return meta\n\n\n    def filename_to_id(filename):\n        parts = filename.split('/')[-2:]\n        if parts[-1][0] in ['H', 'O']:\n            return f'{parts[-1][:-4]}'\n        return f'{parts[0]}/{parts[1][:-4]}'\n        \n\n    def get_meta_unlabeled():\n        \"\"\" build meta and pseudo labels for pseudo labeled data (soundscapes).\n            Uses predition files in sequence: all audios with hits from first file, then \n            all audios with hits from second file (excluding the audios included in the first),\n            and so on.\n        \"\"\"\n        \n        # build meta\n        merged_meta, merged_unlabeled_probs = None, None\n        for file in config.unlabeled_pseudo_preds:\n            oof = load_obj(f'{ROOT}/{file}')\n            unlabeled_probs = oof[BIRDS].values.astype(np.float32)\n            ids = oof.row_id.values\n            meta = []\n            for id_ in ids:\n                filename = '_'.join(id_.split('_')[:-1]) + '.ogg'\n                order = int(id_.split('_')[-1]) // 5 - 1\n                meta.append(dict(filename=filename, order=order))\n            meta = pd.DataFrame(meta)        \n        \n            # filter and adjust pseudo labels\n            f_primary = np.max(unlabeled_probs, axis=1) > config.unlabeled_primary_th\n            unlabeled_probs = unlabeled_probs[f_primary]\n            unlabeled_probs[unlabeled_probs < config.pseudo_label_zero_th] = 0\n            meta = meta[f_primary]\n            meta['id_'] = [x[:-4] for x in meta.filename]\n            if merged_meta is None:\n                merged_meta = meta\n                merged_unlabeled_probs = unlabeled_probs\n            else:\n                existing_ids = set(merged_meta.id_.unique())\n                f_valid = ~meta.id_.isin(existing_ids)\n                merged_meta = pd.concat([merged_meta, meta[f_valid]], ignore_index=True)\n                merged_unlabeled_probs = np.concatenate([merged_unlabeled_probs, unlabeled_probs[f_valid]])\n            print(f'unlabeled file: {file}: {len(merged_meta):,} - {merged_meta.id_.nunique():,}')\n                \n        merged_meta['primary_label'] = [BIRDS[x] for x in np.argmax(merged_unlabeled_probs, axis=1)]\n        merged_meta['labels'] = [[x] for x in merged_meta.primary_label]\n        merged_meta['n_labels'] = 1\n        merged_meta['i'] = np.arange(len(merged_meta))\n        merged_meta['group'] = np.arange(len(merged_meta))\n        merged_meta['ws'] = config.unlabeled_weight / merged_meta.groupby('filename')['i'].transform('count')\n\n        return merged_meta, merged_unlabeled_probs\n\n\n    def read_train(cfg):\n        \"\"\" read train data with default labels \"\"\"\n        \n        meta = get_train_add_meta(additional_data_path=f'{ROOT}/b5-data-{cfg.previous_dataset}-l-0/meta')\n        meta['filename'] = [x.split('.')[0]+'.ogg' for x in meta.filename]\n        meta['counter'] = meta.groupby('primary_label')['filename'].transform('count')\n        print(meta.source.value_counts(), '\\n')\n\n        datasets = [cfg.dataset, cfg.previous_dataset, cfg.train_unlabeled_dataset]\n        filename_to_path = get_filename_to_path(datasets)\n        files_edges = load_obj(cfg.files_edges)\n\n        # fix duplication (2 audios with same filename and different extensions and lengths)\n        meta = meta[meta.filename != 'compot1/104826.ogg']\n\n        # build train+data catalog\n        meta['id_'] = [filename_to_id(x) for x in meta.filename.values]\n\n        # remove repeat audios with different ids\n        meta = meta[~meta.id_.isin(set(DUPLICATED_AUDIOS))]\n\n        id_to_source = dict(zip(meta['id_'], meta['source']))\n        id_to_labels = dict(zip(meta['id_'], meta['labels']))\n        id_to_labels_num = dict()\n        for id_, labels in id_to_labels.items():\n            id_to_labels_num[id_] = tuple(BIRDS.index(x) for x in labels if x in BIRDS)\n\n        catalog = get_catalog_train(meta, id_to_labels, id_to_labels_num, files_edges)\n        \n        # Exclude from meta files not in catalog (duplicated)\n        meta = meta[meta.id_.isin(catalog.keys())]\n        assert len(set(meta.id_.values) - set(files_edges.keys())) == 0, 'missing filenames in files_edges'\n\n        # add unlabeled with pseudo soft labels\n        if cfg.train_unlabeled_dataset is not None:\n            meta_unlabeled, unlabeled_probs = get_meta_unlabeled()\n            print(f'Equivalent pseudo audios: {meta_unlabeled.ws.sum():,.0f}')\n            meta = pd.concat([meta, meta_unlabeled], axis=0)\n            catalog_unlabeled = get_catalog_pseudo_unlabeled(\n                meta_unlabeled, unlabeled_probs, th_min=config.pseudo_label_zero_th, th_pos=config.unlabeled_primary_th)\n            catalog = catalog | catalog_unlabeled\n\n        # add path to meta and catalog\n        meta['path'] = meta.filename.map(filename_to_path)\n        for id_, rec in catalog.items():\n            filename = f'{id_}.ogg' if id_[0] not in ['H', 'O'] else f'{id_.split(\"#\")[0]}.ogg'\n            rec.file = filename_to_path[filename]\n\n        if 'noise_path' in globals():\n            files_background_noise = glob.glob(f'{noise_path}/ff1010bird_nocall/nocall/*.ogg')\n        else:\n            files_background_noise = glob.glob(f'{ROOT}/birdclef2021-background-noise/ff1010bird_nocall/nocall/*.ogg')\n        assert meta.primary_label.isnull().sum() == 0, 'No support for null primary labels'\n\n        print(f'\\nmeta shape: {meta.shape}  catalog: {len(catalog):,.0f}  '\n              f'background noise: {len(files_background_noise):,.0f}')\n        \n        return meta, catalog, files_background_noise\n    \n\n    def add_folds_strat(meta, strat_col='primary_label', group_col='author', seed=42, n_folds=5):\n        \"\"\" add folds to meta using stratification (split equally between folds) and groups\n            (each group in a single fold). Use add_mlstratified_folds because it's better.\n        \"\"\"\n\n        def _get_folds_strat(index, groups, strats, random_state=42, n_folds=5):\n            if n_folds is None:\n                n_folds = config.n_folds\n            sgkf = StratifiedGroupKFold(n_splits=n_folds, shuffle=True, random_state=random_state)\n            folds = pd.Series(0, index=index)\n            for fold, (_, val_idx) in enumerate(sgkf.split(index, y=strats, groups=groups)):\n                folds.iloc[val_idx] = fold\n            return folds\n\n        groups = meta[group_col] if group_col is not None else None\n        strats = meta[strat_col]\n        index = np.arange(meta.shape[0])\n        folds = _get_folds_strat(index, groups, strats, random_state=seed, n_folds=n_folds)\n        folds_df = pd.DataFrame(dict(fold=folds, i=meta.i), index=meta.index)\n        return folds_df\n        \n    \n    def get_meta_folds_catalog(meta, catalog):\n        \"\"\" take a meta df and a catalog as dict and return catalog as a list and a meta\n            that contains an entry for each catalog position with the author and primary_label\n            for fold assigment.\n        \"\"\"\n        # assign fold\n        meta = meta.copy()\n        meta.loc[meta.ws.isnull(), 'ws'] = 1.\n        meta['sample_'] = ['t' if x[0] not in ['H', 'O'] else 'u' for x in meta.id_.values]\n        meta['i'] = np.arange(len(meta))\n        meta['split_candidate'] = meta.counter >= 10\n        folds = meta.loc[meta.split_candidate, ['i', 'primary_label', 'group']]\n        folds = add_folds_strat(\n            folds, strat_col='primary_label', group_col='group', n_folds=config.n_folds, seed=config.seed)\n        meta = pd.merge(meta, folds[['i', 'fold']], on='i', how='left')\n        meta.loc[meta.fold.isnull() & (meta.sample_ == 't'), 'fold'] = -2\n        meta.loc[meta.fold.isnull() & (meta.sample_ == 'u'), 'fold'] = -3\n\n        print(meta.groupby('fold')['id_'].size())\n        meta.to_csv('meta.csv')\n        save_obj(meta, 'meta')\n\n        catalog_l = list(catalog.values())\n        return meta, catalog_l\n\n    # ----------------------------------------------------------------------------------------------------\n    # 0.6 audio sections\n    # ----------------------------------------------------------------------------------------------------\n\n    def standardize_audio_filters(raw_audio_filters):\n        \"\"\" finalize a dict of audio dilters by converting audio hits into audio sections (start, end),\n            merging them when applicable (if the distance to the previous hit <= 5s).\n            section for a hit is defined a (hit - BAND, hit + BAND)\n            The list sections is prefixed with the type of curation:\n                - 'm': the bird is vocalizing in every 5s segment of every section\n                - 'a': the bird is not guaranted to vocalize in every 5s segment of every section\n                - 'i': the bird is not vocalizing; ignore the audio\n        \"\"\"\n        BAND = 4\n        \n        audio_filters = dict()\n        for id_, raw_sections in raw_audio_filters.items():\n            prior_hit = -100\n            if raw_sections[0] in ['m', 'i']:\n                curation = raw_sections[0]\n                start = 1\n            else:\n                curation = 'a'\n                start = 0\n            sections = [curation]\n            for sec in raw_sections[start:]:\n                if isinstance(sec, tuple):\n                    sections.append(sec)\n                    if sec[1] is not None:\n                        prior_hit = sec[1] - 4\n                else:\n                    if prior_hit + 5 >= sec:  # merge with previous section\n                        sections[-1] = (sections[-1][0], sec + BAND)\n                    else:\n                        start = max(sec - BAND, 0)\n                        end = 5 if start == 0 else sec + BAND\n                        sections.append((start, end))\n                    prior_hit = sec\n                        \n            audio_filters[id_] = sections\n        return audio_filters   \n    \n\n    def get_audio_sections(meta, curation_mode, files_edges):\n        \"\"\" build dict of audio sections dict[id_, [curation, (start, end)*]].\n            `curation_mode` is one of:\n                - none: no curation\n                \n            The start and end values are offsets in seconds within the original audio.\n            curation: (m)anual, (a)utomatic, (i)gnore.\n        \"\"\"\n       \n        # complete train filters and get file edges\n        if curation_mode == 'none':\n            raw_audio_filters = dict()\n        else:\n            raw_audio_filters = standardize_audio_filters(TRAIN_AUDIO_FILTERS)\n\n        # prepare custom filter format\n        audio_sections = dict()\n        for row in meta.itertuples():\n            id_ = row.filename[:-4]\n            \n            # if not in files_edges, then it should be unlabeled audio with 1m length\n            audio_start, audio_end, audio_size = files_edges.get(row.id_, (0, 60 * SR, 60 * SR))\n            sections = raw_audio_filters.get(id_, None)\n            if sections is None:\n                if row.primary_label in RARE_BIRDS:\n                    sections = ['a', (0, audio_size / SR)]\n                else:\n                    if curation_mode == 'speech':\n                        sections = ['a', (audio_start / SR, audio_end / SR)]\n                    else:\n                        sections = ['a', (0, audio_size / SR)]\n            elif sections[-1][-1] is None:  # end of audio\n                sections[-1] = (sections[-1][0], audio_size / SR)\n\n            # make sure that end of last section is not larger than the length of the audio file\n            sections[-1] = (sections[-1][0], min(sections[-1][1], audio_size / SR))\n            audio_sections[id_] = sections\n        return audio_sections\n\n\n    def get_filename_to_path(datasets):\n        filename_to_path = dict()\n        for dataset in datasets:\n            if dataset is None:\n                continue\n            if 'unlabeled' in dataset:\n                files = glob.glob(f'{ROOT}/b5-data-{dataset}*/*.h5')\n                id_offset = 1\n            else:\n                files = glob.glob(f'{ROOT}/b5-data-{dataset}*/*/*.h5')\n                id_offset = 2\n            if len(files) == 0:\n                break\n            for file in files:\n                filename = '/'.join(file.split('/')[-id_offset:])[:-3] + '.ogg'\n                filename_to_path[filename] = file\n        return filename_to_path\n        \n    # ------------------------------------------------------------------------------------------------------\n    # 0.7 pseudo labels\n    # ----------------------------------------------------------------------------------------------------\n\n    def get_catalog_train(meta, id_to_labels, id_to_labels_num, files_edges):\n        \"\"\" get catalog for training data, which is:\n            dict[id_, dict(file, bird, curation, sections, sections_prob, hard_labels, mask)]\n        \"\"\"\n        audio_sections = get_audio_sections(meta, config.curation_mode, files_edges)\n        catalog = dict()\n        mask_ones = np.ones(N_BIRDS)\n        mask_ones_t = torch.tensor(mask_ones, dtype=torch.int32)\n        for id_ in tqdm(meta.id_.values):\n            audio_sec = audio_sections.get(id_, None)\n            curation = audio_sec[0] if audio_sec is not None else 'a'\n            if curation == 'i':\n                continue\n\n            primary = id_to_labels[id_][0]\n            labels_num = list(id_to_labels_num[id_])\n            expected_labels = np.zeros(N_BIRDS)\n            expected_labels[labels_num] = 1\n            filename = f'{id_}.ogg'\n            rec = dotdict(id_=id_, primary_label=primary, curation=curation, ws=1)\n                \n            sections = audio_sec[1:]\n            rec.sections = sections\n            rec.sections_prob = [e - s for s, e in rec.sections]\n            rec.hard_labels = torch.tensor([expected_labels] * len(rec.sections), dtype=torch.float32)\n            rec.mask = [mask_ones_t] * len(sections)\n            catalog[id_] = rec\n        \n        return catalog\n\n\n    def get_catalog_pseudo_unlabeled(\n            meta_unlabeled, unlabeled_probs, th_min=.1, th_pos=.5):\n        \"\"\" get catalog of pseudo labels for unlabeled data, which is:\n            dict[id_, dict(file, bird, curation, sections, sections_prob, soft_labels, hard_labels, mask)]\n            each 5s segment gets it's own entry.\n        \"\"\"\n        catalog = dict()\n        mask_ones = np.ones(N_BIRDS)\n        mask_ones_t = torch.tensor(mask_ones, dtype=torch.int32)\n        for row, p in tqdm(zip(meta_unlabeled.itertuples(index=False), unlabeled_probs), total=len(unlabeled_probs)):\n            id_ = f'{row.filename[:-4]}#{row.order}'\n            rec = dotdict(id_=row.filename[:-4], curation='a')\n            segment_start = row.order * 5\n            rec.sections = [(segment_start, segment_start + 5)]\n            p = p.copy()\n            p[p < th_min] = 0\n            pt = torch.tensor(p, dtype=torch.float32)\n            rec.soft_labels = [pt]\n            rec.hard_labels = [pt]\n            rec.mask = [mask_ones_t]\n            rec.primary_label = row.primary_label\n            rec.sections_prob = [1]\n            rec.ws = row.ws\n            catalog[id_] = rec\n        return catalog        \n\n","metadata":{"trusted":true,"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-06-12T17:31:13.211064Z","iopub.execute_input":"2025-06-12T17:31:13.211662Z","iopub.status.idle":"2025-06-12T17:31:13.610975Z","shell.execute_reply.started":"2025-06-12T17:31:13.211639Z","shell.execute_reply":"2025-06-12T17:31:13.610331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 'SKIP' not in globals() or not SKIP:     # audio filters\n    print('audios filters exclude speech v3.0')\n\n    # 3.0: rare audios for all spcies in `UNCOMMON_BIRDS`\n\n    UNCOMMON_BIRDS = [\n    # <= 30 audios\n        '1139490', '1192948', '1194042', '126247', '1346504', '134933', '135045', '1462711', '1462737', '1564122',\n        '21038', '21116', '24272', '24292', '24322', '41778', '41970', '42007', '42087', '42113', '46010', '47067',\n        '476537', '476538', '48124', '50186', '523060', '528041', '548639', '64862', '65336', '65344', '65349',\n        '65419', '65962', '66016', '66531', '66578', '66893', '67082', '67252', '714022', '715170', '787625',\n        '81930', '868458', '963335', 'plctan1', 'sahpar1', 'shghum1', 'turvul', 'woosto',\n    ]\n\n    TRAIN_AUDIO_FILTERS = {  # list of tuple (start, end) and/or float (time of vocalization)\n        '1139490/CSA36385': ['m', (0, 8)],\n        '1139490/CSA36389': ['m', (0, 8)],\n\n        '1192948/CSA36358': ['m', (0, 8)],\n        '1192948/CSA36366': ['m', (0, 8)],\n        '1192948/CSA36373': ['m', (0, 8)],\n        '1192948/CSA36388': ['m', (0, 8)],\n\n        '1194042/CSA18783': ['m', (0, 27)],\n        '1194042/CSA18794': ['m', (0, 8.5)],\n        '1194042/CSA18802': ['m', (0, 15), (23, 35)],\n\n        '126247/111388': ['m', (0, 83)],\n        '126247/111389': ['m', (0, 35), (45, 63)],\n        '126247/111390': ['m', (0, 97)],\n        '126247/113953': ['m', (0, 54)],\n        '126247/114600': ['m', (0, 156), (160, 191)],\n        '126247/115310': ['m', (0, 56), (60, 85)],\n        '126247/115314': ['m', (0, 67)],\n        '126247/115377': ['m', (0, 68)],\n\n        '1346504/CSA18784': ['m', (0, 38)],\n        '1346504/CSA18791': ['m', (0, 71)],\n        '1346504/CSA18792': ['m', (0, 23)], \n        '1346504/CSA18793': ['m', (0, 62)],\n        '1346504/CSA18803': ['m', (0, 94)],\n\n        '134933/114020': ['m', (0, 57)],\n        \n        '135045/114263': ['m', (0, 80)],\n        '135045/114265': ['m', (0, 182)],\n        '135045/114266': ['m', (0, 95)],\n        '135045/114268': ['m', (0, 90)],\n\n        '1462711/CSA36371': ['m', (0, 8)],\n        '1462711/CSA36379': ['m', (0, 8)],\n        '1462711/CSA36390': ['m', (0, 8)],\n\n        '1462737/CSA36341': ['m', (0, 8)],\n        '1462737/CSA36369': ['m', (0, 8)],\n        '1462737/CSA36380': ['m', (0, 8)],\n        '1462737/CSA36381': ['m', (0, 8)],\n        '1462737/CSA36386': ['m', (0, 9)],\n        '1462737/CSA36391': ['m', (0, 8)],\n        '1462737/CSA36395': ['m', (0, 8)],\n\n        '21116/114410': ['m', (0, 148)],\n        '21116/114490': ['m', (0, 278)],\n\n        '24272/114033': ['m', (0, 292)],          \n        '24272/115328': ['m', (0, 190)],          \n        '24272/XC882885': ['m', (0, 62)],\n\n        '24292/115323': ['m', (0, 50)],\n        '24292/115324': ['m', (0, 46)],\n        '24292/115325': ['m', (0, 33)],\n        '24292/115326': ['m', (0, 28)],\n        '24292/115327': ['m', (0, 47)],\n        '24292/115338': ['m', (0, 126)],\n        '24292/115909': ['m', (0, 76)],\n        '24292/115921': ['m', (0, 115)],\n        '24292/116126': ['m', (0, 65)],\n        '24292/CSA34649': ['m', (0, 130), (140, 170)],\n        '24292/CSA34651': ['m', (0, 95)],\n        '24292/CSA35021': ['m', (0, 76)],\n\n        '476537/115311': ['m', (0, 71)],\n        '476537/115313': ['m', (0, 27)],\n        '476537/CSA35459': ['m', (0, 135)],\n        '476537/CSA35461': ['m', (0, 260)],\n\n        '476538/114086': ['m', (0, 164)],\n\n        '528041/CSA36359': ['m', (0, 8)],\n        '528041/CSA36365': ['m', (0, 8)],\n\n        '64862/CSA18218': ['m', (3, None)],\n        '64862/CSA18222': ['m', (4, None)],\n\n        '555142/110064': ['m', (6, None)],\n        '555142/111435': ['m', (0, 78)],\n        '555142/111438': ['m', (0, 62)],\n        '555142/114329': ['m', (0, 125)],\n        '555142/114339': ['m', (0, 214)],\n        '555142/114437': ['m', (0, 160)],\n        '555142/114443': ['m', (0, 92)],\n        '555142/114489': ['m', (0, 60), (75, 115)],\n        '555142/114541': ['m', (0, 112)],\n        '555142/114542': ['m', (0, 35)],\n        '555142/114543': ['i', (0, 0)],  # repeat of 114542\n        '555142/114544': ['m', (0, 35)],\n        '555142/114627': ['m', (0, 41)],\n        '555142/114631': ['m', (0, 110)],\n        '555142/114632': ['i', (0, 0)],  # repeat of 114631\n        '555142/115330': ['m', (0, 77)],\n        '555142/115331': ['m', (0, 102)],\n        '555142/115343': ['m', (0, 33)],\n        '555142/115350': ['m', (0, 60)],\n        '555142/115353': ['m', (0, 44)],\n        '555142/115356': ['m', (0, 18)],\n        '555142/115357': ['m', (0, 40)],\n        '555142/115358': ['m', (0, 18)],\n        '555142/115360': ['m', (0, 65)],\n        '555142/115361': ['m', (0, 11)],\n        '555142/115369': ['m', (0, 16)],\n        '555142/115903': ['m', (0, 185)],\n        '555142/115904': ['m', (0, 142)],\n        '555142/115906': ['m', (0, 47)],\n        '555142/115922': ['m', (0, 115)],\n        '555142/115924': ['m', (0, 265)],\n        '555142/115936': ['m', (0, 57), (135, 165)],  # excluded section with speech and possible vocalization\n\n        '65419/114028': ['m', (0, 134)],\n        '65419/114682': ['m', (0, 66)],\n\n        '65547/111437': ['m', (0, 75)],\n        '65547/113941': ['m', (5, 20)],\n        '65547/113942': ['m', (5, 36)],\n        '65547/113943': ['m', (5, 36)],\n        '65547/113944': ['m', (5, 38)],\n        '65547/113945': ['m', (5, 32)],\n        '65547/113946': ['m', (5, 32)],\n        '65547/113947': ['m', (5, 32)],\n        '65547/113948': ['m', (5, 21)],\n        '65547/113949': ['m', (5, 28)],\n        '65547/113950': ['m', (5, 45)],\n        '65547/114018': ['m', (0, 122)],\n        '65547/114022': ['m', (0, 182)],\n        '65547/114023': ['m', (0, 125)],\n        '65547/114030': ['m', (0, 270)],\n        '65547/114031': ['m', (0, 185)],\n        '65547/114032': ['m', (0, 157)],\n        '65547/114035': ['m', (0, 165)],\n        '65547/111437': ['m', (0, 127)],\n        '65547/114331': ['m', (0, 85)],\n        '65547/115940': ['m', (0, 130)],\n        '65547/115941': ['m', (0, 111)],\n        '65547/115942': ['m', (0, 52)],\n        '65547/115943': ['m', (0, 171)],\n        '65547/115944': ['m', (0, 350)],\n        '65547/115945': ['m', (0, 87)],\n        '65547/115946': ['m', (0, 47)],\n        '65547/115947': ['m', (0, 305)],\n        '65547/115948': ['m', (0, 76)],\n        '65547/115949': ['m', (0, 76)],\n        '65547/115950': ['m', (0, 68)],\n        '65547/115951': ['m', (0, 86)],\n        '65547/115952': ['m', (0, 140)],\n        '65547/115953': ['m', (0, 77)],\n        '65547/115954': ['m', (0, 142)],\n        '65547/115955': ['m', (0, 46)],\n        '65547/115956': ['m', (0, 155)],\n        '65547/115957': ['m', (0, 53)],\n        '65547/115958': ['m', (0, 215)],\n        '65547/115959': ['m', (0, 92)],\n        '65547/115960': ['m', (0, 145)],\n        '65547/115961': ['m', (0, 82)],\n        '65547/115962': ['m', (0, 16)],\n        '65547/115963': ['m', (0, 71)],\n        '65547/115964': ['m', (0, 132)],\n        '65547/115965': ['m', (0, 502)],\n        '65547/115966': ['m', (0, 106)],\n        '65547/115967': ['m', (0, 12)],\n        '65547/115968': ['m', (0, 37)],\n        '65547/115969': ['m', (0, 5)],\n        \n        '66531/114088': ['m', (2, 153)],\n        '66531/114089': ['m', (0, 114)],\n\n        '66578/111440': ['m', (0, 16), (31, 77)],\n        '66578/115454': ['m', (0, 65)],\n        \n        '67082/111436': ['m', (0, 24)],\n        '67082/114091': ['m', (0, 146)],\n        '67082/114092': ['m', (0, 340)],\n        '67082/115410': ['m', (0, 80)],\n        '67082/115411': ['m', (0, 83)],\n        \n        '787625/110487': ['m', (6, None)],\n        '787625/113311': ['m', (5, None)],\n        '787625/113312': ['m', (4, None)],\n        '787625/114267': ['m', (0, 35)],\n        '787625/114889': ['m', (0, 60)],\n        \n        '81930/114081': ['m', (0, 26)],\n        '81930/114623': ['m', (0, 175)],\n        '81930/114624': ['i', (0, 0)],      # repeat of 114623\n        '81930/114636': ['m', (0, 37)],\n        '81930/115322': ['m', (0, 30)],\n        '81930/115351': ['m', (0, 37)],\n        '81930/115370': ['m', (0, 56)],\n        '81930/115378': ['m', (0, 28)],\n\n\n        '963335/CSA36372': ['m', (0, 9)],\n        '963335/CSA36374': ['m', (0, 9)],\n        '963335/CSA36375': ['m', (0, 9)],\n        '963335/CSA36377': ['m', (0, 8)],\n        '963335/CSA36393': ['m', (0, 8)],\n        \n        'plctan1/114860': ['m', (0, 72)],\n        'plctan1/114875': ['m', (0, 40)],\n\n        'colcha1/XC337020': ['m', (45, 230)],\n        'colcha1/XC532406': ['m', (0, 10)],\n        \n        # v3 <= 30 audios\n        '42007/iNat1241987': ['i', (0, 0)],  # Ignore, puma spotted but not heard\n        '42007/iNat15105': ['i', (0, 0)],    # Ignore, puma spotted but not heard\n        '42007/iNat500217': ['m', (0, 114)],\n        \n        '48124/CSA03598': ['m', (0, 47)],\n        '48124/CSA18785': ['m', (0, 61)],\n        '48124/CSA18795': ['m', (0, 126)],\n        '48124/CSA18798': ['m', (0, 192)],\n        '48124/CSA34485': ['m', (0, 78)],\n        '48124/CSA35111': ['m', (0, 120)],\n        '48124/CSA35116': ['m', (0, 155)],\n        '48124/CSA35118': ['m', (0, 105)],\n        '48124/CSA35157': ['m', (0, 222)],\n        '48124/CSA35160': ['m', (0, 250)],\n        '48124/CSA35173': ['m', (0, 290)],\n        '48124/CSA35179': ['m', (0, 236)],\n        '48124/CSA35181': ['m', (0, 86)],\n        '48124/CSA35183': ['m', (0, 7)],\n        '48124/CSA35184': ['m', (0, 134)],\n        '48124/CSA35185': ['m', (0, 115)],\n        '48124/CSA35187': ['m', (0, 5)],\n        '48124/CSA35188': ['m', (0, 18)],\n        '48124/CSA35190': ['m', (0, 7)],\n        '48124/CSA36346': ['m', (0, 29)],\n        \n        '67252/114602': ['m', (0, 195)],\n        '67252/115376': ['m', (0, 42)],\n        '67252/XC882988': ['m', (0, 7), (20, 105)],\n        '67252/XC882992': ['m', (0, 42)],\n        '67252/XC882993': ['m', (0, 18)],\n        '67252/XC882994': ['m', (0, 28)],\n        '67252/XC882995': ['m', (0, 78)],\n        '67252/XC882996': ['m', (0, 59)],\n        '67252/XC882997': ['m', (0, 20)],\n        '67252/XC882999': ['m', (0, 300)],\n        \n        '715170/CSA34510': ['m', (0, 75)],\n        '715170/CSA34511': ['m', (0, 58)],\n        '715170/CSA34512': ['m', (0, 95)],\n        '715170/CSA34513': ['m', (0, 35)],\n        '715170/CSA34514': ['m', (0, 58)],\n        '715170/CSA34515': ['m', (0, 66)],\n        '715170/CSA35799': ['m', (0, 60)],\n        '715170/CSA35800': ['m', (0, 55)],\n        '715170/CSA35801': ['m', (0, 64)],\n        '715170/CSA35802': ['m', (0, 74)],\n        '715170/CSA35803': ['m', (0, 70)],\n        '715170/CSA35804': ['m', (0, 52)],\n        '715170/CSA35805': ['m', (0, 89)],\n        '715170/CSA35806': ['m', (0, 55)],\n        '715170/CSA35807': ['m', (0, 49)],\n        '715170/CSA35808': ['m', (0, 59)],\n        '715170/CSA35809': ['m', (0, 61)],\n\n        '65344/111439': ['m', (0, 63)],\n        '65344/111481': ['m', (0, 300)],\n        '65344/114888': ['m', (0, 63)],\n\n        '24322/110493': ['m', (5, 49)],\n        '24322/114082': ['m', (2, 14)],\n        '24322/114083': ['m', (0, 38)],\n        '24322/114619': ['m', (0, 225)],\n        '24322/114620': ['m', (0, 240)],\n        '24322/115379': ['m', (0, 37)],\n        '24322/115417': ['m', (0, 107)],\n        '24322/115419': ['m', (0, 118)],\n        '24322/115420': ['m', (0, 56)],\n        '24322/115905': ['m', (0, 47)],\n        '24322/XC882877': ['m', (0, 67)],\n        '24322/XC882880': ['m', (0, 97)],\n        \n        '50186/CSA04279': ['m', (0, 12)],\n        '50186/CSA04299': ['m', (0, 61)],\n        '50186/CSA14874': ['m', (3, None)],\n        '50186/CSA18166': ['m', (5, 18)],\n        '50186/CSA18185': ['m', (5, 110)],\n        '50186/CSA18191': ['m', (5, 46), (67, 75), (85, 95)],\n        '50186/CSA18201': ['m', (5, 61)],\n        '50186/CSA18202': ['m', (5, 46)],\n        '50186/CSA18282': ['m', (6, None)],\n        '50186/CSA18916': ['m', (5, None)],\n        '50186/CSA20037': ['m', (4, None)],\n        '50186/CSA28885': ['m', (3, None)],\n        '50186/CSA34402': ['m', (0, 68)],\n        '50186/CSA34428': ['m', (0, 25)],\n        '50186/CSA34456': ['m', (0, 38)],\n        '50186/CSA34622': ['m', (0, 26)],\n        '50186/CSA34678': ['m', (0, 48)],\n        '50186/CSA35126': ['m', (0, 90)],\n        '50186/CSA35128': ['m', (0, 215)],\n        '50186/CSA35133': ['m', (0, 188)],\n        '50186/CSA35140': ['m', (0, 185)],\n        '50186/CSA35146': ['m', (0, 267)],\n        '50186/CSA35148': ['m', (0, 232)],\n        '50186/CSA35149': ['m', (0, 202)],\n        '50186/CSA35159': ['m', (0, 293)],\n        '50186/CSA35174': ['m', (0, 265)],\n        '50186/CSA35175': ['m', (0, 135)],\n        \n        '65962/iNat137091': ['m', (16, None)],\n        \n        'shghum1/114578': ['m', (0, 56)],\n        'shghum1/114583': ['m', (0, 16), (36, 52)],\n        'shghum1/114586': ['m', (0, 82)],\n        'shghum1/114641': ['m', (0, 85)],\n\n        'turvul/989811': ['m', (0, 82)],\n        \n        'sahpar1/110750': ['m', (5, 33)],\n        'sahpar1/110751': ['m', (5, 36)],\n        'sahpar1/112982': ['m', (3, 63)],\n        'sahpar1/113010': ['m', (3, 42)],\n        'sahpar1/113037': ['m', (5, 142)],\n    }\n\n    DUPLICATED_AUDIOS = [\n        '555142/114543', '555142/114632', '81930/114624'\n        #'65547/iNat1103224', '66893/iNat1109827',\n    ]\n\n","metadata":{"trusted":true,"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-06-12T17:31:13.612500Z","iopub.execute_input":"2025-06-12T17:31:13.612978Z","iopub.status.idle":"2025-06-12T17:31:13.647669Z","shell.execute_reply.started":"2025-06-12T17:31:13.612955Z","shell.execute_reply":"2025-06-12T17:31:13.646818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Config","metadata":{}},{"cell_type":"code","source":"class Aug1(Config):\n    Flip = .5\n    CoarseDropout = (.375, .375, 1, .7)\n    mixup = 1\n    volume = (1/3, 3)\n    audio = False\n\n\nclass Config_train_tr1(Config):\n    seed = 10\n\n    # dataset\n    dataset = 'h5-6-s'\n    previous_dataset = 'add-h5-6-s'\n    train_unlabeled_dataset = 'h5-unlabeled-2-s'\n    files_edges = '/kaggle/input/b5-cache/files_edges speech 5'\n    curation_mode = 'speech'\n    \n    pseudo_label_zero_th = .1\n    unlabeled_primary_th = .5\n    train_primary_th = .5\n    mask_non_labels = True\n    sample_unlabeled_prob = .0   # sample unlabeld \n    unlabeled_weight = 1\n    dataset_val = 'val-128256-m1-4-data'\n    extension = '.ogg'  # .ogg _3.png\n    background_noise_prob = .5\n    background_noise_reference = True\n    background_noise_cache_size = 1_000\n    background_noise_max_usage = 6\n    montage_cache_size = 10\n    n_preload_species = 15\n    use_mask = False\n    ws_power = 1/2\n    \n    filter = (False, 0)    # filter_secondary, min_quality\n    spec = None\n    img_dim = (128, 256)\n    img_duration = 5\n    spec_norm = False\n    montage = (0, 1)\n    train_mode = '2x7s'\n    norm_audio = False\n    fold_col = 'author'  # geo_label author\n    train_batch_size = 64\n    val_batch_size = 64\n    num_workers = 6\n    prefetch_factor = 2\n\n    aug = Aug1()\n\n    # model\n    model_name = 'SpecNetImg'\n    in_chans = 1\n    out_indices = 2\n    encoder = 'eca_nfnet_l0'  # tf_efficientnetv2_s  tf_efficientnetv2_b2  eca_nfnet_l0\n    add_position = False\n    drop_path_rate = None\n    head_dropout = 0.\n    gem_p = 1.8\n\n    # train\n    # wandb_project = 'test_2',\n    train_size = 28_000\n    accumulation_steps = 1\n    n_folds = 5\n    duration = {'hours': 11, 'minutes': 55}\n    log_every_n_steps = 50\n    check_val_every_n_epoch = 1\n    # adam8 = False,\n    precision_schedule = {0: 'float32'}\n    # max_grad_value = 0.5,\n    max_grad_norm = 10\n\n    # loss & metrics\n    loss_schedule = {0: 'focal_volodymyr'}\n    label_smoothing = .0\n    monitor = 'val_auc'  # acc\n    mode = 'max'\n\n    # optimizer\n    optimizer = 'AdamW'\n    saved_optimizer = None # ''\n    weight_decay = 1e-6\n    eps = 1e-8\n    betas = (0.9, 0.999)\n\n    # scheduler\n    epochs = 50\n    # config.max_epochs = 20\n    step_scheduler_after = 'step'  #'step', 'epoch'\n    schedule_type = 'cos'\n    saved_scheduler = None\n    warmup_share = 1 / epochs\n    lr = 2.5e-4    # 1e-3, 2.5e-4 for eca\n    start_lr = None\n    end_lr = 1e-6\n\n    # swa\n    swa = None # (2e-4, 2, 4),\n    \n    # awp\n    awp = False\n    adv_th = 0.3  # only apply awp if tr_score < threshold\n    adv_lr = .005  # 1\n    adv_eps = 1e-2 # 0.001\n    \nconfig = Config_train_tr1().to_dict()\nmode = ['train']\n\nTYPE = 'train'\nROOT = '/kaggle/input'","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-06-12T17:35:01.714714Z","iopub.execute_input":"2025-06-12T17:35:01.715090Z","iopub.status.idle":"2025-06-12T17:35:01.726532Z","shell.execute_reply.started":"2025-06-12T17:35:01.715058Z","shell.execute_reply":"2025-06-12T17:35:01.725614Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Train","metadata":{"_kg_hide-input":true}},{"cell_type":"code","source":"if 'train' in mode:\n    config = Config_train_tr1().to_dict()\n    config.name = 'ebs.1'\n    config.curation_mode = 'speech'\n    config.folds = [0]\n    config.schedule_type = 'multi_lr'\n    config.lr = [(1e-5, 1e-3), (3, 2.5e-4)]\n    config.montage = (0, 3)\n    config.train_mode = 'random'\n    config.unlabeled_pseudo_preds = ['b5-data-pseudo-pred-v3-data/unlabeled pseudo pred v2']\n    config.montage_bird_prob = False\n    config.train_batch_size = 64\n    config.encoder = 'tf_efficientnetv2_s'\n    config.pretrained_encoder = '/kaggle/input/b5-pretrained-weights-s/tf_efficientnetv2_s_in21k_Pretrainversion1.pth'\n    config.log_every_n_steps = 1\n    meta, catalog, files_background_noise = read_train(config)\n    train(meta, catalog, True)\n    print('done')","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-06-12T17:35:04.132595Z","iopub.execute_input":"2025-06-12T17:35:04.132909Z","execution_failed":"2025-06-12T17:45:13.812Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"_____","metadata":{"_kg_hide-input":true}}]}