{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":11754970,"sourceType":"datasetVersion","datasetId":7287224}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nfrom pathlib import Path\nfrom types import SimpleNamespace\n\ncfg = SimpleNamespace(**{})\ncfg.num_folds = 5\ncfg.fold = -1\ncfg.gpu = \"0\"\nos.environ['CUDA_VISIBLE_DEVICES'] = cfg.gpu\n\ncfg.fname = 'efficientvit_b0'\ncfg.seed = 2025\n\ncfg.input_path = Path('../input')\ncfg.comp_data_path = cfg.input_path / 'birdclef-2025'\ncfg.save_path = Path('../working')\ncfg.soundscape_path = cfg.comp_data_path / 'train_soundscapes'\ncfg.audio_path = Path(\"/kaggle/input/birdclef-2025/train_audio\")\n\ncfg.logger_file = True\n\n# image size\ncfg.image_height = 224\ncfg.image_width = 224\n\n# audio\ncfg.duration = 5\ncfg.sr = 32000\ncfg.fmin = 40\ncfg.fmax = 15000\ncfg.n_fft = 1536\ncfg.n_mels = cfg.image_height\ncfg.win_length = 1024\ncfg.hop_length = int((cfg.duration * cfg.sr - cfg.win_length + cfg.n_fft) / (cfg.image_width)) + 1 \n\n# training HP\ncfg.num_epochs = 10\ncfg.train_batch_size = 128\ncfg.valid_batch_size = 128\ncfg.workers = 2\ncfg.grad_value = 2.0\ncfg.grad_norm = 0.0\ncfg.grad_norm_type = 2\ncfg.device = \"cuda\"\ncfg.accumulate = 1\n\n# optimizer\ncfg.lr = 7e-5\ncfg.decay = 0.01\ncfg.opt_beta1 = 0.9\ncfg.opt_beta2 = 0.999\ncfg.opt_eps = 1e-8\ncfg.optimizer = 'AdamW'\ncfg.no_decay = True\n\n# scheduler\ncfg.pct_start = 0.1\ncfg.max_lr = 3e-3\ncfg.final_div_factor = 100\n\n# augmentations\ncfg.resample_train = 10\ncfg.max_shift = 1\ncfg.mixup_p = 0.2\ncfg.other_samples = 1\ncfg.db_range = 10.0\ncfg.loudness_range = 10.0\n\n# logging\ncfg.local_rank = 0\ncfg.verbose=True\n\n# model\ncfg.backbone = 'efficientvit_b0.r224_in1k'\ncfg.gem_pooling = False\ncfg.bce = True\ncfg.drop_rate = 0.1\n\n# tasks hp\ncfg.train_model = True\ncfg.num_folds = 5\ncfg.pretrained_path = None \n\ncfg.pl_batch_size = 4\ncfg.pl_dup = 10\ncfg.pl = [Path(\"/kaggle/input/birdclef-data/\")]\ncfg.max_pl = True","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2025-05-01T17:11:13.803103Z","iopub.execute_input":"2025-05-01T17:11:13.803873Z","iopub.status.idle":"2025-05-01T17:11:13.815532Z","shell.execute_reply.started":"2025-05-01T17:11:13.803835Z","shell.execute_reply":"2025-05-01T17:11:13.814847Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport random\nfrom tqdm import tqdm\nfrom logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\nimport gc\nimport pickle as pkl\n\nimport librosa\n\nfrom torch.utils.data import DataLoader, Dataset\nimport torchaudio\nimport torchaudio.transforms as T\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.amp import autocast, GradScaler\nfrom torch.optim import lr_scheduler, Adam, AdamW\n\nimport timm\n\nfrom glob import glob\nfrom sklearn.model_selection import KFold, GroupKFold\nfrom sklearn.metrics import roc_auc_score\nfrom scipy.special import logit, expit","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:13.816945Z","iopub.execute_input":"2025-05-01T17:11:13.817266Z","iopub.status.idle":"2025-05-01T17:11:13.835952Z","shell.execute_reply.started":"2025-05-01T17:11:13.817233Z","shell.execute_reply":"2025-05-01T17:11:13.835289Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pprint import pprint\nmodel_names = timm.list_models(pretrained=True)\npprint(model_names)","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:13.836986Z","iopub.execute_input":"2025-05-01T17:11:13.837309Z","iopub.status.idle":"2025-05-01T17:11:13.874115Z","shell.execute_reply.started":"2025-05-01T17:11:13.837292Z","shell.execute_reply":"2025-05-01T17:11:13.873403Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cfg.train_model:\n    checkpoint_path = cfg.save_path / cfg.fname\n    if not cfg.save_path.exists():\n        cfg.save_path.mkdir()\n    if not checkpoint_path.exists():\n        checkpoint_path.mkdir()\n    checkpoint_path = checkpoint_path / \"exp_0\"\n    exp = 0\n    while(checkpoint_path.exists()):\n        exp += 1\n        checkpoint_path = cfg.save_path / f\"exp_{exp}\"\n    checkpoint_path.mkdir()\n    cfg.checkpoint_path = checkpoint_path\n    print(f\"Saving checkpoint path to {checkpoint_path}\")","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:13.875030Z","iopub.execute_input":"2025-05-01T17:11:13.875253Z","iopub.status.idle":"2025-05-01T17:11:13.888637Z","shell.execute_reply.started":"2025-05-01T17:11:13.875237Z","shell.execute_reply":"2025-05-01T17:11:13.888062Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_logger(cfg):\n    logger = getLogger(cfg.fname)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    if cfg.logger_file and cfg.train_model:\n        filename= cfg.checkpoint_path / \"run.log\"\n        handler2 = FileHandler(filename=filename)\n        handler1.setFormatter(Formatter(\"%(message)s\"))\n        logger.addHandler(handler2)\n    return logger\n\ndef seed_torch(seed_value):\n    random.seed(seed_value) # Python\n    np.random.seed(seed_value) # cpu vars\n    torch.manual_seed(seed_value) # cpu  vars    \n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(seed_value)\n        torch.cuda.manual_seed_all(seed_value) # gpu vars\n    if torch.backends.cudnn.is_available:\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:13.889833Z","iopub.execute_input":"2025-05-01T17:11:13.890017Z","iopub.status.idle":"2025-05-01T17:11:13.905304Z","shell.execute_reply.started":"2025-05-01T17:11:13.890004Z","shell.execute_reply":"2025-05-01T17:11:13.904599Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv(cfg.comp_data_path / 'train.csv')\ntrain","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:13.905977Z","iopub.execute_input":"2025-05-01T17:11:13.906499Z","iopub.status.idle":"2025-05-01T17:11:14.046751Z","shell.execute_reply.started":"2025-05-01T17:11:13.906474Z","shell.execute_reply":"2025-05-01T17:11:14.046178Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train[\"species\"] = [filename.split(\"/\")[0] for filename in train[\"filename\"]]\ntrain[\"record\"] = [filename.split(\"/\")[1] for filename in train[\"filename\"]]\ntrain[\"secondary_labels\"] = [eval(sls) for sls in train[\"secondary_labels\"]]","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:14.048281Z","iopub.execute_input":"2025-05-01T17:11:14.048497Z","iopub.status.idle":"2025-05-01T17:11:14.243708Z","shell.execute_reply.started":"2025-05-01T17:11:14.048481Z","shell.execute_reply":"2025-05-01T17:11:14.243199Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = train.groupby(\"record\").size()\ndf = df[df > 1]\ndf","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:14.244372Z","iopub.execute_input":"2025-05-01T17:11:14.244569Z","iopub.status.idle":"2025-05-01T17:11:14.276576Z","shell.execute_reply.started":"2025-05-01T17:11:14.244553Z","shell.execute_reply":"2025-05-01T17:11:14.275851Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = train.groupby(\"record\").agg({\n    'species': ['first', 'last'],\n    'secondary_labels': ['first', 'last'],\n})\n\ndf.columns = [\"first_species\", \"last_species\", \"first_secondary\", \"last_secondary\"]\ndf.reset_index(inplace=True)\ndf","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:14.277348Z","iopub.execute_input":"2025-05-01T17:11:14.277596Z","iopub.status.idle":"2025-05-01T17:11:14.329591Z","shell.execute_reply.started":"2025-05-01T17:11:14.277580Z","shell.execute_reply":"2025-05-01T17:11:14.328970Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = train.merge(df[['record', 'first_species', 'last_species',]],\n                   how='left',\n                   on='record')","metadata":{"execution":{"iopub.status.busy":"2025-05-01T17:11:14.330297Z","iopub.execute_input":"2025-05-01T17:11:14.330554Z","iopub.status.idle":"2025-05-01T17:11:14.359728Z","shell.execute_reply.started":"2025-05-01T17:11:14.330531Z","shell.execute_reply":"2025-05-01T17:11:14.359190Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_audio(file_name, isfirst, istrain, cfg):\n    filepath = file_name.split(\"/\")[0]\n    fname = file_name.split(\"/\")[1].split(\".\")[0]\n    filepath = cfg.input_path / \"birdclef-data\" / \"birdclef_data\" / \"birdclef_data\" / filepath\n\n    if istrain:\n        max_duration = int((cfg.duration + cfg.max_shift) * cfg.sr)\n    else:\n        max_duration = cfg.duration * cfg.sr\n\n    if isfirst:\n        filepath = filepath / f\"first10_{fname}.npy\"\n        audio = np.load(filepath)\n        audio = audio[:max_duration]\n    else:\n        filepath = filepath / f\"last10_{fname}.npy\"\n        audio = np.load(filepath)\n        audio = audio[-max_duration:]\n\n    return audio","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:14.362361Z","iopub.execute_input":"2025-05-01T17:11:14.362585Z","iopub.status.idle":"2025-05-01T17:11:14.367699Z","shell.execute_reply.started":"2025-05-01T17:11:14.362569Z","shell.execute_reply":"2025-05-01T17:11:14.367186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"new_train = []\n\nkf = KFold(n_splits=cfg.num_folds, shuffle=True, random_state=0)\nfor species, df in train.groupby('species'):\n    df = df.reset_index(drop=True)\n    df['fold'] = -1\n    if len(df) < 5:\n        df['fold'] = 0\n    else:\n        for fold, (train_idx, valid_idx) in enumerate(kf.split(df, df.primary_label)):\n            df.loc[valid_idx, \"fold\"] = fold\n    new_train.append(df)\nnew_train = pd.concat(new_train).reset_index(drop=True)    \nnew_train.fold.value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:14.368413Z","iopub.execute_input":"2025-05-01T17:11:14.368692Z","iopub.status.idle":"2025-05-01T17:11:14.920674Z","shell.execute_reply.started":"2025-05-01T17:11:14.368676Z","shell.execute_reply":"2025-05-01T17:11:14.920043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"new_train.groupby(['species', 'fold']).size().unstack()\ntrain = new_train","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:14.921345Z","iopub.execute_input":"2025-05-01T17:11:14.921570Z","iopub.status.idle":"2025-05-01T17:11:14.931385Z","shell.execute_reply.started":"2025-05-01T17:11:14.921554Z","shell.execute_reply":"2025-05-01T17:11:14.930732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def metric(preds, targets, cfg):\n    score = {}\n    for j, label in enumerate(cfg.labels):\n        y_true = targets[:, j]\n        y_pred = preds[:, j]\n        if len(np.unique(y_true)) < 2:\n            score[label] = np.nan\n            continue\n        score[label] = roc_auc_score(y_true, y_pred)\n    score_avg = np.nanmean([v for v in score.values()])\n    return score_avg, score\n\ndef metric_db(train, oofs, cfg):\n    score = {}\n    for j,label in enumerate(cfg.labels):\n        score[label] = roc_auc_score(train.primary_label == label, oofs[label])\n    score_avg = np.mean([v for k,v in score.items()])\n    return score_avg, score\n    \ndef my_softmax(preds):\n    preds = preds - preds.max(1, keepdims=True)\n    preds = np.exp(preds.clip(-20, 0))\n    preds = preds / preds.sum(1, keepdims=True)\n    return preds\n\ndef bce_with_mask(preds, targets, mask):\n    loss = nn.BCEWithLogitsLoss(reduction='none')(preds, targets)\n    loss = loss * mask\n    loss = loss.mean()\n    return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:14.932350Z","iopub.execute_input":"2025-05-01T17:11:14.932637Z","iopub.status.idle":"2025-05-01T17:11:14.942837Z","shell.execute_reply.started":"2025-05-01T17:11:14.932614Z","shell.execute_reply":"2025-05-01T17:11:14.941918Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cfg.device = torch.device(\"cuda\")\nif cfg.bce:\n    cfg.loss = bce_with_mask\nelse:\n    cfg.loss = nn.CrossEntropyLoss()\n    \ncfg.labels = np.array(sorted(train.primary_label.unique()))\ncfg.num_labels = len(cfg.labels)\ncfg.targets = {v: i for i, v in enumerate(cfg.labels)}\n\ncfg.logger = get_logger(cfg)\nseed_torch(cfg.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:14.943546Z","iopub.execute_input":"2025-05-01T17:11:14.943785Z","iopub.status.idle":"2025-05-01T17:11:14.963034Z","shell.execute_reply.started":"2025-05-01T17:11:14.943767Z","shell.execute_reply":"2025-05-01T17:11:14.962291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdPLDataset():\n    def __init__(self, pl_preds, fold, cfg):\n        self.cfg = cfg\n        pl_filenames = [k for k,pred in pl_preds.items() if len(pred[0]) == 12]\n        self.filename = np.array(pl_filenames)\n        self.preds = pl_preds\n        self.fold = fold\n        self.savepath = cfg.input_path / 'birdclef-data' / 'unlabeled_soundscapes' / 'unlabeled_soundscapes'\n        \n    def __len__(self):\n        return len(self.filename) * self.cfg.pl_dup\n    \n    def __getitem__(self, idx):\n        idx = idx % len(self.filename)\n        filepath = self.filename[idx]\n        cfg = self.cfg\n        filename = filepath.split('/')[-1].split('.')[0]\n        audio = np.load(self.savepath / (filename + '.npy'))\n        duration = int(len(audio) / cfg.sr)\n        periods = duration // cfg.duration\n        audio = audio[:periods * cfg.duration * cfg.sr]\n        audio = torch.from_numpy(audio).view(periods, -1)\n        targets = self.preds[filepath][self.fold].astype(np.float32)\n        targets = torch.from_numpy(targets)\n        secondary_mask = torch.ones(*targets.shape)\n        out = {\n            'audio' : audio,\n            'targets' : targets,\n            'secondary_mask' : secondary_mask,\n        }\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:14.963902Z","iopub.execute_input":"2025-05-01T17:11:14.964996Z","iopub.status.idle":"2025-05-01T17:11:14.974032Z","shell.execute_reply.started":"2025-05-01T17:11:14.964973Z","shell.execute_reply":"2025-05-01T17:11:14.973138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(cfg.pl[0] / f\"pl_all_b1.pkl\", \"rb\") as file:\n    pl_preds = pkl.load(file)\nlen(pl_preds)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:14.974712Z","iopub.execute_input":"2025-05-01T17:11:14.974969Z","iopub.status.idle":"2025-05-01T17:11:15.784217Z","shell.execute_reply.started":"2025-05-01T17:11:14.974948Z","shell.execute_reply":"2025-05-01T17:11:15.783485Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for k,pred in pl_preds.items():\n    break\nk, len(pred)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:15.785102Z","iopub.execute_input":"2025-05-01T17:11:15.785342Z","iopub.status.idle":"2025-05-01T17:11:15.790091Z","shell.execute_reply.started":"2025-05-01T17:11:15.785326Z","shell.execute_reply":"2025-05-01T17:11:15.789577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.plot(np.mean(pred, 0)[0], alpha=0.5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:15.790754Z","iopub.execute_input":"2025-05-01T17:11:15.790953Z","iopub.status.idle":"2025-05-01T17:11:15.937667Z","shell.execute_reply.started":"2025-05-01T17:11:15.790938Z","shell.execute_reply":"2025-05-01T17:11:15.936934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pl_filenames = [k for k,pred in pl_preds.items() if len(pred[0]) == 12]\nlen(pl_filenames)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:15.938513Z","iopub.execute_input":"2025-05-01T17:11:15.938694Z","iopub.status.idle":"2025-05-01T17:11:15.947485Z","shell.execute_reply.started":"2025-05-01T17:11:15.938680Z","shell.execute_reply":"2025-05-01T17:11:15.946946Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cfg.max_pl:\n    for k in pl_filenames:\n        preds = pl_preds[k]\n        preds = [logit(pred.clip(1e-7, 1 - 1e-7)) for pred in preds]\n        preds_max = np.max(preds, 0)\n        pl_preds[k] = [expit((pred + preds_max + pred.mean() - preds_max.mean()) / 2.0) for pred in preds]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:15.948238Z","iopub.execute_input":"2025-05-01T17:11:15.948614Z","iopub.status.idle":"2025-05-01T17:11:19.762543Z","shell.execute_reply.started":"2025-05-01T17:11:15.948598Z","shell.execute_reply":"2025-05-01T17:11:19.761965Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pl_dataset = BirdPLDataset(pl_preds, 2, cfg)\nfor k,v in pl_dataset[1].items():\n    print(k, v.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:19.763238Z","iopub.execute_input":"2025-05-01T17:11:19.763438Z","iopub.status.idle":"2025-05-01T17:11:19.788980Z","shell.execute_reply.started":"2025-05-01T17:11:19.763423Z","shell.execute_reply":"2025-05-01T17:11:19.788277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_pl_iter(pl_preds, fold, istrain, cfg):\n    pl_dataset = BirdPLDataset(pl_preds, fold, cfg)\n    pl_data_loader = DataLoader(\n        pl_dataset,\n        batch_size=cfg.pl_batch_size,\n        num_workers=0,\n        shuffle=istrain,\n        pin_memory=False,\n        #collate_fn=collate_pad,\n        drop_last = False,\n    )\n    \n    return iter(pl_data_loader), pl_data_loader\n\ndef get_pl_batch(pl_dataloader):\n    try:\n        batch = next(pl_dataloader[0])\n    except StopIteration:\n        pl_data_loader = pl_dataloader[1]\n        pl_data_loader = iter(pl_data_loader), pl_data_loader\n        batch = next(pl_dataloader[0])\n    new_batch = {k : v.view(v.shape[0] * v.shape[1], v.shape[2]) \n                 for k,v in batch.items()}\n    return new_batch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:19.791714Z","iopub.execute_input":"2025-05-01T17:11:19.792157Z","iopub.status.idle":"2025-05-01T17:11:19.797298Z","shell.execute_reply.started":"2025-05-01T17:11:19.792127Z","shell.execute_reply":"2025-05-01T17:11:19.796584Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pl_dataloader = get_pl_iter(pl_preds, 1, False, cfg)\nbatch = get_pl_batch(pl_dataloader)\nfor k,v in batch.items():\n    print(k, v.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:19.797871Z","iopub.execute_input":"2025-05-01T17:11:19.798028Z","iopub.status.idle":"2025-05-01T17:11:19.885498Z","shell.execute_reply.started":"2025-05-01T17:11:19.798015Z","shell.execute_reply":"2025-05-01T17:11:19.884856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdDataset():\n    def __init__(self, train, istrain, cfg):\n        self.cfg = cfg\n        self.istrain = istrain\n        self.filename = train.filename.values\n        #self.primary_label = train.primary_label.values\n        self.secondary_labels = train.secondary_labels.values\n        self.first_species = train.first_species.values\n        self.last_species = train.last_species.values\n        \n    def __len__(self):\n        return len(self.filename)\n    \n    def get_audio(self, idx):\n        filename = self.filename[idx]\n        duration = self.cfg.sr * self.cfg.duration\n        if self.istrain:\n            first = np.random.rand() < 0.5\n            audio = load_audio(filename, first, True, self.cfg)\n            if len(audio) < duration:\n                pad_length = np.random.randint(0, duration - len(audio) + 1) \n                audio = np.pad(audio, \n                               ((pad_length, duration - len(audio) - pad_length),), \n                               mode='constant')\n            else:\n                start = np.random.randint(0, len(audio) - duration + 1)\n                audio = audio[start : start + duration]\n        else:\n            audio = load_audio(filename, True, False, self.cfg)\n            audio = audio[:duration]\n            if len(audio) < duration:\n                pad_length = (duration - len(audio)) // 2\n                audio = np.pad(audio, \n                               ((pad_length, duration - len(audio) - pad_length),), \n                               mode='constant')\n        return audio\n      \n    def __getitem__(self, idx):\n        audio = self.get_audio(idx)\n        targets = np.zeros(len(cfg.labels), dtype=np.float32)        \n        targets[self.cfg.targets[self.first_species[idx]]] = 1.0\n        targets[self.cfg.targets[self.last_species[idx]]] = 1.0\n        secondary_mask = np.ones(len(cfg.labels), dtype=np.float32)\n        secondary_labels = self.secondary_labels[idx]\n        if len(secondary_labels) > 0:\n            for label in secondary_labels:\n                if label in cfg.targets:\n                    secondary_mask[cfg.targets[label]] = 0\n        secondary_mask = np.maximum(secondary_mask, targets)        \n        out = {\n            'audio' : torch.from_numpy(audio),\n            'targets' : torch.from_numpy(targets),\n            'secondary_mask' : secondary_mask,\n        }\n        return out\n\ndef batch_to_device(batch, device):\n    return {k:batch[k].to(device, non_blocking=True) for k in batch.keys() if k not in []}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:19.886201Z","iopub.execute_input":"2025-05-01T17:11:19.886422Z","iopub.status.idle":"2025-05-01T17:11:19.899084Z","shell.execute_reply.started":"2025-05-01T17:11:19.886399Z","shell.execute_reply":"2025-05-01T17:11:19.898362Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = BirdDataset(train, True, cfg)\nelt = dataset[0]\nfor k,v in elt.items():\n    print(k, v.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:19.899903Z","iopub.execute_input":"2025-05-01T17:11:19.900172Z","iopub.status.idle":"2025-05-01T17:11:19.921709Z","shell.execute_reply.started":"2025-05-01T17:11:19.900132Z","shell.execute_reply":"2025-05-01T17:11:19.921185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_data_loader(dataset, istrain, cfg):\n    if istrain:\n        batch_size = cfg.train_batch_size\n    else:\n        batch_size = cfg.valid_batch_size\n    data_loader = DataLoader(\n        dataset,\n        batch_size=batch_size,\n        num_workers=cfg.workers,\n        shuffle=istrain,\n        pin_memory=False,\n        #collate_fn=collate_pad,\n        drop_last = istrain,\n    )\n    return data_loader ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:19.922301Z","iopub.execute_input":"2025-05-01T17:11:19.922548Z","iopub.status.idle":"2025-05-01T17:11:19.933340Z","shell.execute_reply.started":"2025-05-01T17:11:19.922533Z","shell.execute_reply":"2025-05-01T17:11:19.932643Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def gem(x, p=3, eps=1e-6):\n    return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6, p_trainable=True):\n        super(GeM,self).__init__()\n        if p_trainable:\n            self.p = nn.Parameter(torch.ones(1)*p)\n        else:\n            self.p = p\n        self.eps = eps\n\n    def forward(self, x):\n        return gem(x, p=self.p, eps=self.eps)       \n    def __repr__(self):\n        return self.__class__.__name__ + '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + ', ' + 'eps=' + str(self.eps) + ')'\n\nclass BirdModel(nn.Module):\n    def __init__(self, cfg, pretrained: bool = True):\n        super(BirdModel, self).__init__()\n        self.cfg = cfg\n        self.mel = T.MelSpectrogram(\n            sample_rate=cfg.sr, n_fft=cfg.n_fft, win_length=cfg.win_length, \n            hop_length= cfg.hop_length, f_min=cfg.fmin, f_max=cfg.fmax, \n            n_mels=cfg.n_mels, mel_scale='htk', power=2.0)\n        self.A2DB = T.AmplitudeToDB(stype=\"power\")\n        self.backbone = timm.create_model(\n            cfg.backbone,\n            pretrained=pretrained,\n            drop_rate = cfg.drop_rate,\n            num_classes=cfg.num_labels\n        )\n        #if cfg.gem_pooling == \"gem\":\n        #    self.backbone.head.global_pool = GeM(p_trainable=args.p_trainable)\n         \n    def forward(self, input_dict):\n        x = input_dict['audio']\n        with autocast(enabled=False, device_type=\"cuda\"), torch.no_grad():\n            x = x / torch.std(x, 1, keepdim=True)\n            x = x.float()\n            x = self.mel(x)\n            x = self.A2DB(x)\n            x = (x - 40) / 80\n        with torch.no_grad():\n            x = x.unsqueeze(1)\n            pos = torch.linspace(0., 1., x.size(2)).to(x.device)\n            pos = pos.unsqueeze(0).unsqueeze(0).unsqueeze(-1)\n            pos = pos.expand(x.size(0), 1, x.size(2), x.size(3))\n            x = x.expand(-1, 2, -1, -1)\n            x = torch.cat([x, pos], 1)\n\n        x = self.backbone(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:19.933945Z","iopub.execute_input":"2025-05-01T17:11:19.934167Z","iopub.status.idle":"2025-05-01T17:11:19.949744Z","shell.execute_reply.started":"2025-05-01T17:11:19.934125Z","shell.execute_reply":"2025-05-01T17:11:19.949068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_loader = get_data_loader(dataset, False, cfg)\nmodel = BirdModel(cfg)\nfor batch in data_loader:\n    break\nout = model(batch)\nout.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:19.950356Z","iopub.execute_input":"2025-05-01T17:11:19.950537Z","iopub.status.idle":"2025-05-01T17:11:26.856177Z","shell.execute_reply.started":"2025-05-01T17:11:19.950523Z","shell.execute_reply":"2025-05-01T17:11:26.855203Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(loader, pl_data_loader, model, optimizer, scheduler, scaler, device, cfg):\n    model.train()\n    model.zero_grad()\n    if cfg.verbose:\n        bar = tqdm(range(len(loader)))\n    else:\n        bar = range(len(loader))\n    load_iter = iter(loader)\n    loss_l = []\n    grad_norm_l = []\n    \n    accumulate = cfg.accumulate\n    \n    for i, batch in zip(bar, load_iter):\n        input_dict0 = batch_to_device(batch, device)\n        pl_batch = get_pl_batch(pl_data_loader)\n        pl_dict = batch_to_device(pl_batch, device)\n        input_dict = {\n            k : torch.cat([input_dict0[k], pl_dict[k]], 0) for k in input_dict0.keys()\n        }\n        with autocast(enabled=True, device_type=\"cuda\"):\n            targets = input_dict['targets']\n            if cfg.loudness_range:\n                loudness = - np.log(cfg.loudness_range)\n                bs = targets.shape[0]\n                #weight = 0.1 ** (cfg.loundness_range *  np.random.rand() / 10)\n                weight = torch.rand(bs, 1).to(targets.device)\n                weight = torch.exp(weight * loudness)\n                audio = input_dict['audio']\n                audio = audio * weight\n                input_dict['audio'] = audio\n            secondary_mask = input_dict['secondary_mask']\n            mixup = (np.random.rand() < cfg.mixup_p)\n            if mixup:\n                bs = targets.shape[0]\n                perm = torch.randperm(bs).to(targets.device)\n                weight = 0.1 ** (cfg.db_range *  np.random.rand() / 10)\n                audio = input_dict['audio']\n                audio = audio + weight * audio[perm]\n                input_dict['audio'] = audio\n                secondary_mask = torch.minimum(secondary_mask, secondary_mask[perm])\n                targets = torch.maximum(targets, targets[perm])\n                secondary_mask = torch.maximum(secondary_mask, targets)        \n            preds = model(input_dict)\n            loss = cfg.loss(preds, targets, secondary_mask).mean()\n        loss_l.append(loss.detach().cpu().item())\n        scaler.scale(loss / cfg.accumulate).backward() \n        accumulate -= 1\n        if accumulate == 0:\n            if cfg.grad_value:\n                scaler.unscale_(optimizer)\n                nn.utils.clip_grad_value_(model.parameters(), cfg.grad_value)\n            if cfg.grad_norm:\n                scaler.unscale_(optimizer)\n                total_norm = nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_norm).item()\n                if np.isnan(total_norm):\n                    total_norm = cfg.grad_norm\n                else:\n                    total_norm = np.clip(total_norm, 0, cfg.grad_norm)\n                grad_norm_l.append(total_norm)\n            scaler.step(optimizer)     \n            scaler.update()\n            optimizer.zero_grad()\n            accumulate = cfg.accumulate\n            scheduler.step()  \n        del preds, targets, loss, input_dict\n        if cfg.verbose:\n            if cfg.grad_norm:\n                bar.set_description('loss: %.4f grad norm %.1f' % (np.mean(loss_l), np.mean(grad_norm_l),))\n            else:\n                bar.set_description('loss: %.4f ' % np.mean(loss_l))\n    optimizer.zero_grad()\n    del loss_l, bar\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:26.857415Z","iopub.execute_input":"2025-05-01T17:11:26.857646Z","iopub.status.idle":"2025-05-01T17:11:26.869413Z","shell.execute_reply.started":"2025-05-01T17:11:26.857623Z","shell.execute_reply":"2025-05-01T17:11:26.868508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def valid_epoch(loader, model, device, cfg):\n    model.eval()\n    model.zero_grad()\n    if cfg.verbose:\n        bar = tqdm(range(len(loader)))\n    else:\n        bar = range(len(loader))\n    load_iter = iter(loader)\n    preds_l = []\n    targets_l = []\n    with torch.no_grad():\n        for i, batch in zip(bar, load_iter):      \n            input_dict = batch_to_device(batch, device)\n            with autocast(enabled=False, device_type=\"cuda\"):\n                preds = model(input_dict)                \n            preds_l.append(preds.detach().cpu())\n            del preds, input_dict\n            targets = batch['targets']\n            targets_l.append(targets)\n        preds = torch.cat(preds_l)\n        targets = torch.cat(targets_l)\n        return preds.numpy(), targets.numpy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:26.870119Z","iopub.execute_input":"2025-05-01T17:11:26.870862Z","iopub.status.idle":"2025-05-01T17:11:26.891523Z","shell.execute_reply.started":"2025-05-01T17:11:26.870837Z","shell.execute_reply":"2025-05-01T17:11:26.890832Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_optimizer(model, cfg):\n    no_decay = [\"bias\", \"norm\"]\n    if cfg.no_decay:\n        optimizer_parameters = [\n            {'params': [p for n, p in model.named_parameters() \n                        if not any(nd in n for nd in no_decay)],\n             'weight_decay': cfg.decay},\n            {'params': [p for n, p in model.named_parameters() \n                        if any(nd in n for nd in no_decay)],\n             'weight_decay': 0.0},\n        ] \n    else:\n        optimizer_parameters = model.parameters()\n    optimizer = torch.optim.AdamW(optimizer_parameters, lr=cfg.lr)\n    return optimizer\n\ndef get_scheduler(optimizer, train_data_loader, cfg):\n    scheduler = lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=cfg.max_lr,\n        epochs=cfg.num_epochs,\n        steps_per_epoch=len(train_data_loader),\n        pct_start=cfg.pct_start,\n        anneal_strategy=\"cos\",\n        final_div_factor=cfg.final_div_factor,\n    )\n    return scheduler","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:26.892259Z","iopub.execute_input":"2025-05-01T17:11:26.892462Z","iopub.status.idle":"2025-05-01T17:11:26.909053Z","shell.execute_reply.started":"2025-05-01T17:11:26.892447Z","shell.execute_reply":"2025-05-01T17:11:26.908436Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_checkpoint(model, fold, cfg):\n    checkpoint = {\n        'model' : model.state_dict(),\n        'fold' : fold,\n        'seed' : seed,\n        }\n    checkpoint_path = cfg.checkpoint_path\n    save_path = checkpoint_path / ('%s_fold%d.pth' % (cfg.fname, fold))\n    if cfg.local_rank == 0:\n        cfg.logger.info('saving %s ...' % save_path)\n    torch.save(checkpoint, save_path)\n    if cfg.local_rank == 0:\n        cfg.logger.info('done')\n\ndef load_checkpoint(fold, cfg):\n    if cfg.pretrained_path:\n        checkpoint_path = cfg.pretrained_path\n    else:\n        checkpoint_path = cfg.checkpoint_path\n    save_path = checkpoint_path / ('%s_fold%d.pth' % (cfg.fname, fold))\n    cfg.logger.info('loading %s ...' % save_path)\n    checkpoint = torch.load(save_path, map_location='cpu')\n    model = BirdModel(cfg, pretrained=False).to(cfg.device)\n    model.load_state_dict(checkpoint['model'], strict=True)\n    model.eval()\n    cfg.logger.info('done')\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:26.909747Z","iopub.execute_input":"2025-05-01T17:11:26.909944Z","iopub.status.idle":"2025-05-01T17:11:26.919031Z","shell.execute_reply.started":"2025-05-01T17:11:26.909928Z","shell.execute_reply":"2025-05-01T17:11:26.918505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def resample(train_fold, cfg):\n    new_train = []\n    for species, df in train.groupby('species'):\n        new_train.append(df)\n        if len(df) < cfg.resample_train:\n            df = df.sample(n=(cfg.resample_train - len(df)), replace=True, random_state=cfg.seed)\n            new_train.append(df)\n    new_train = pd.concat(new_train).reset_index(drop=True)  \n    return new_train  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:26.920012Z","iopub.execute_input":"2025-05-01T17:11:26.920267Z","iopub.status.idle":"2025-05-01T17:11:26.936140Z","shell.execute_reply.started":"2025-05-01T17:11:26.920251Z","shell.execute_reply":"2025-05-01T17:11:26.935516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scores = []\noofs = pd.DataFrame(columns=cfg.labels, index=train.index, data = 0.0)\noofs['filename'] = train.filename","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:11:26.936839Z","iopub.execute_input":"2025-05-01T17:11:26.937096Z","iopub.status.idle":"2025-05-01T17:11:27.009023Z","shell.execute_reply.started":"2025-05-01T17:11:26.937079Z","shell.execute_reply":"2025-05-01T17:11:27.008437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for fold in range(cfg.num_folds):\n    if cfg.fold >= 0 and fold != cfg.fold:\n        continue\n    seed = cfg.seed + fold\n    seed_torch(seed)\n    train_fold = train[train.fold != fold]\n    valid_dataset = BirdDataset(train[train.fold == fold], False, cfg)\n    valid_dataloader = get_data_loader(valid_dataset, istrain=False, cfg=cfg)\n    device = cfg.device\n    \n    if cfg.pretrained_path:\n        model = load_checkpoint(fold, seed, cfg)\n    else:\n        model = BirdModel(cfg, pretrained=True).to(device)\n    optimizer = get_optimizer(model, cfg)\n    scheduler = None\n    scaler = GradScaler()\n    result = None\n    pl_data_loader = get_pl_iter(pl_preds, fold, False, cfg)\n    for epoch in range(cfg.num_epochs):\n        if cfg.resample_train:\n            train_fold = resample(train_fold, cfg)\n        train_dataset = BirdDataset(train_fold, True, cfg)\n        train_dataloader = get_data_loader(train_dataset, istrain=True, cfg=cfg)\n\n        if scheduler is None:\n            scheduler = get_scheduler(optimizer, train_dataloader, cfg)\n        train_epoch(train_dataloader, pl_data_loader, model, \n                    optimizer, scheduler, scaler, device, cfg)\n        if valid_dataset is not None:\n            preds, targets = valid_epoch(valid_dataloader, model, device, cfg)\n            if cfg.bce:\n                preds = expit(preds) # model uses logits\n            else:\n                preds = my_softmax(preds) # model uses logits\n            result, _ = metric(preds, targets, cfg)\n            msg = f\"seed {cfg.seed} fold {fold} epoch {epoch} metric {result:.4f}\"\n            cfg.logger.info(msg)\n        else:\n            msg = f\"seed {cfg.seed} fold {fold} epoch {epoch}\"\n            cfg.logger.info(msg)\n    if cfg.local_rank == 0:\n        save_checkpoint(model, fold, cfg)\n    del model, optimizer, scheduler, scaler, train_dataloader, \n    if valid_dataset is not None:\n        del valid_dataloader\n    gc.collect()\n    torch.cuda.empty_cache()\n    scores.append(result)\n    \n    for j,c in enumerate(cfg.labels):\n        oofs.loc[train.fold == fold, c] = preds[:, j]\noofs.to_csv(cfg.checkpoint_path / 'oofs.csv', index=False)\n\nnp.mean(scores), metric_db(train, oofs, cfg)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:15:04.613611Z","iopub.execute_input":"2025-05-01T17:15:04.614341Z","iopub.status.idle":"2025-05-01T17:27:42.520439Z","shell.execute_reply.started":"2025-05-01T17:15:04.614310Z","shell.execute_reply":"2025-05-01T17:27:42.519745Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"res = metric_db(train, oofs, cfg)\nres","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:14:39.383345Z","iopub.status.idle":"2025-05-01T17:14:39.383715Z","shell.execute_reply.started":"2025-05-01T17:14:39.383529Z","shell.execute_reply":"2025-05-01T17:14:39.383544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dfe = train.groupby('primary_label').size()\ndfe.name = 'size'\ndfe = dfe.reset_index().sort_values('primary_label').reset_index(drop=True)\ndfe","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:14:39.384861Z","iopub.status.idle":"2025-05-01T17:14:39.385088Z","shell.execute_reply.started":"2025-05-01T17:14:39.384984Z","shell.execute_reply":"2025-05-01T17:14:39.384994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.DataFrame({'label' : cfg.labels, \n                   'roc_auc' : [res[1][label] for label in cfg.labels],\n                   'size' : dfe['size'].values,\n                  })\ndf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:14:39.386228Z","iopub.status.idle":"2025-05-01T17:14:39.386485Z","shell.execute_reply.started":"2025-05-01T17:14:39.386363Z","shell.execute_reply":"2025-05-01T17:14:39.386376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.scatter(np.log(df['size']), df['roc_auc'], marker='+')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T17:14:39.387645Z","iopub.status.idle":"2025-05-01T17:14:39.387896Z","shell.execute_reply.started":"2025-05-01T17:14:39.387776Z","shell.execute_reply":"2025-05-01T17:14:39.387788Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}