{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":25954,"databundleVersionId":2091745,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13744602,"sourceType":"datasetVersion","datasetId":8745663},{"sourceId":13744609,"sourceType":"datasetVersion","datasetId":8745670},{"sourceId":13744751,"sourceType":"datasetVersion","datasetId":8745781},{"sourceId":13744766,"sourceType":"datasetVersion","datasetId":8745794}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Подготовка","metadata":{}},{"cell_type":"markdown","source":"## Задача\n\n- Создать классификатор звуков, издаваемых птицами.\n- В каждом аудиофайле может быть несколько интервалов, в каждом интервале - сигнал различных птиц .\t\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"## Анализ описания конкурса\n\n- Оценка по F1\n- Результ записать в файл submission.csv. Формат выходных данных:\n```\nrow_id,birds\n3575_COL_5,wewpew batpig1\n3575_COL_10,wewpew batpig1\n3575_COL_15,wewpew batpig1\n```\n\n## Анализ набора данных\n- train_short_audio. Файлы с записями отдельных птиц. Частота дискретизации 32 кГц. Формат файлов OGG.\n- train_soundscapes. Файлы с записями групп птиц. Файлы по содержанию сходны с тестовыми. Длительность каждого файла ~ 10 минут. Формат файлов OGG.\n- test_soundscapes. При отправке на проверку будет заполнено тестовыми аудиофайлами (80 шт). Длительность каждого файла ~ 10 минут. Формат файлов OGG. В этой папке также содержатся текстовые файлы с названием и примерными координатами места записи, а также CSV-файл с набором дат записи тестовых звуковых ландшафтов (test_set_recording_dates.csv).\n- test.csv.\n   - row_id: уникальный идентификатор\n   - site: идентификатор местности\n   - seconds: окончание временного окна\n   - audio_id: идентификатор аудиофайла\n- train_metadata.csv.\n   - primary_label: код для названия птицы\n   - filename: имя файла\n   - secondary_labels: коды для других птиц в фрагменте\n","metadata":{}},{"cell_type":"markdown","source":"## Формирование спектрограмм\n\nДля того, чтобы не считать спектрограммы каждый раз заново и не забивать диск и оперативную память, их нужно посчитать заранее и сохранить в виде kaggle-датасетов\nКод для формирования датасетов спектрограмм в https://www.kaggle.com/code/alekseysavinov/sound-2-mels-prepare.","metadata":{}},{"cell_type":"markdown","source":"## Загрузка сформированных датасетов со спектрограммами и метаданными","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport librosa as lb\nimport librosa.display as lbd\nimport soundfile as sf\nfrom  soundfile import SoundFile\nimport pandas as pd\nfrom  IPython.display import Audio\nfrom pathlib import Path\n\nimport torch\nfrom torch import nn, optim\nfrom  torch.utils.data import Dataset, DataLoader\nimport torchvision.models as models\n\nfrom matplotlib import pyplot as plt\n\nimport os, random, gc\nimport re, time, json\nfrom  ast import literal_eval\n\nfrom IPython.display import Audio\nfrom sklearn.metrics import label_ranking_average_precision_score\n\nfrom tqdm.notebook import tqdm\nimport joblib","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:22:48.381149Z","iopub.execute_input":"2025-11-16T08:22:48.381330Z","iopub.status.idle":"2025-11-16T08:22:58.525635Z","shell.execute_reply.started":"2025-11-16T08:22:48.381313Z","shell.execute_reply":"2025-11-16T08:22:58.524992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\nseed_everything()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:22:58.527101Z","iopub.execute_input":"2025-11-16T08:22:58.527511Z","iopub.status.idle":"2025-11-16T08:22:58.538915Z","shell.execute_reply.started":"2025-11-16T08:22:58.527460Z","shell.execute_reply":"2025-11-16T08:22:58.538055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_CLASSES = 397\nSR = 32000\nDURATION = 7\n\nMAX_READ_SAMPLES = 5\n\nDATA_ROOT = Path('/kaggle/input/birdclef-2021')\n\nMEL_PATHS = sorted(Path('/kaggle/input').glob('birdclef-2021-mels-?/rich_train_metadata.csv'))\nTRAIN_LABEL_PATHS = sorted(Path('/kaggle/input').glob('birdclef-2021-mels-?/LABEL_IDS.json'))\n\nMODEL_ROOT = Path('.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:22:58.540058Z","iopub.execute_input":"2025-11-16T08:22:58.540328Z","iopub.status.idle":"2025-11-16T08:22:58.555830Z","shell.execute_reply.started":"2025-11-16T08:22:58.540309Z","shell.execute_reply":"2025-11-16T08:22:58.555266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_BATCH_SIZE = 100\nTRAIN_NUM_WORKERS = 2\n\nVAL_BATCH_SIZE = 128\nVAL_NUM_WORKERS = 2\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nprint('Device:', DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:22:58.556509Z","iopub.execute_input":"2025-11-16T08:22:58.556722Z","iopub.status.idle":"2025-11-16T08:22:58.648051Z","shell.execute_reply.started":"2025-11-16T08:22:58.556686Z","shell.execute_reply":"2025-11-16T08:22:58.647243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_df(mel_paths=MEL_PATHS, train_label_paths=TRAIN_LABEL_PATHS):\n  df_list = []\n  LABEL_IDS = {}\n    \n  for file_path in mel_paths:\n    temp = pd.read_csv(str(file_path), index_col=0)\n    temp['impath'] = temp.apply(\n        lambda row: file_path.parent/'audio_images/{}/{}.npy'.format(row.primary_label, row.filename), \n        axis=1\n    ) \n    df_list.append(temp)\n\n  df = pd.concat(df_list, ignore_index=True)\n  df['secondary_labels'] = df['secondary_labels'].apply(literal_eval)\n\n  for file_path in train_label_paths:\n    with open(str(file_path)) as f:\n      LABEL_IDS.update(json.load(f))\n\n  return LABEL_IDS, df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:22:58.649075Z","iopub.execute_input":"2025-11-16T08:22:58.649393Z","iopub.status.idle":"2025-11-16T08:22:58.659317Z","shell.execute_reply.started":"2025-11-16T08:22:58.649352Z","shell.execute_reply":"2025-11-16T08:22:58.658437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LABEL_IDS, df = get_df()\n\nprint(df.shape)\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:22:58.660218Z","iopub.execute_input":"2025-11-16T08:22:58.660531Z","iopub.status.idle":"2025-11-16T08:23:01.421864Z","shell.execute_reply.started":"2025-11-16T08:22:58.660504Z","shell.execute_reply":"2025-11-16T08:23:01.421031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df['primary_label'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:23:01.424419Z","iopub.execute_input":"2025-11-16T08:23:01.424674Z","iopub.status.idle":"2025-11-16T08:23:01.435999Z","shell.execute_reply.started":"2025-11-16T08:23:01.424654Z","shell.execute_reply":"2025-11-16T08:23:01.435271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df['label_id'].min(), df['label_id'].max()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:23:01.437028Z","iopub.execute_input":"2025-11-16T08:23:01.437319Z","iopub.status.idle":"2025-11-16T08:23:01.449421Z","shell.execute_reply.started":"2025-11-16T08:23:01.437298Z","shell.execute_reply":"2025-11-16T08:23:01.448643Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_model(name, num_classes=NUM_CLASSES):\n    if 'resnet' in name:\n        model = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n    else:\n        raise RuntimeError('Незнакомая модель')\n    \n\n    if hasattr(model, 'fc'):\n        nb_ft = model.fc.in_features\n        model.fc = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, '_fc'):\n        nb_ft = model._fc.in_features\n        model._fc = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, 'classifier'):\n        nb_ft = model.classifier.in_features\n        model.classifier = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, 'last_linear'):\n        nb_ft = model.last_linear.in_features\n        model.last_linear = nn.Linear(nb_ft, num_classes)\n\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:23:01.450287Z","iopub.execute_input":"2025-11-16T08:23:01.450627Z","iopub.status.idle":"2025-11-16T08:23:01.457896Z","shell.execute_reply.started":"2025-11-16T08:23:01.450608Z","shell.execute_reply":"2025-11-16T08:23:01.457262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_data(df):\n    def load_row(row):\n        return row.filename, np.load(str(row.impath))[:MAX_READ_SAMPLES]\n\n    pool = joblib.Parallel(4)\n    mapper = joblib.delayed(load_row)\n    tasks = [mapper(row) for row in df.itertuples(False)]\n    res = pool(tqdm(tasks))\n    res = dict(res)\n    return res","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:23:01.458790Z","iopub.execute_input":"2025-11-16T08:23:01.459052Z","iopub.status.idle":"2025-11-16T08:23:01.467772Z","shell.execute_reply.started":"2025-11-16T08:23:01.459028Z","shell.execute_reply":"2025-11-16T08:23:01.467028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"audio_image_store = load_data(df)\nlen(audio_image_store)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:23:01.468501Z","iopub.execute_input":"2025-11-16T08:23:01.468813Z","execution_failed":"2025-11-16T08:23:11.900Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('shape:', next(iter(audio_image_store.values())).shape)\nlbd.specshow(next(iter(audio_image_store.values()))[0])","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pd.Series([len(x) for x in audio_image_store.values()]).value_counts()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdClefDataset(Dataset):\n\n    def __init__(self, audio_image_store, meta, sr=SR, is_train=True, num_classes=NUM_CLASSES, duration=DURATION):\n        \n        self.audio_image_store = audio_image_store\n        self.meta = meta.copy().reset_index(drop=True)\n        self.sr = sr\n        self.is_train = is_train\n        self.num_classes = num_classes\n        self.duration = duration\n        self.audio_length = self.duration*self.sr\n    \n    @staticmethod\n    def normalize(image):\n        image = image.astype(\"float32\", copy=False) / 255.0\n        image = np.stack([image, image, image])\n        return image\n\n    def __len__(self):\n        return len(self.meta)\n    \n    def __getitem__(self, idx):\n        row = self.meta.iloc[idx]\n        image = self.audio_image_store[row.filename]\n\n        image = image[np.random.choice(len(image))]\n        image = self.normalize(image)\n        \n        t = np.zeros(self.num_classes, dtype=np.float32) + 0.0025\n        t[row.label_id] = 0.995\n        \n        return image, t","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = BirdClefDataset(audio_image_store, meta=df, sr=SR, duration=DURATION, is_train=True)\nlen(df)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = ds[np.random.choice(len(ds))]\nx.shape, y.shape, np.where(y >= 0.5)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lbd.specshow(x[0])","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y[:5]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def one_step(xb, yb, net, criterion, optimizer, scheduler=None):\n  xb, yb = xb.to(DEVICE), yb.to(DEVICE)\n        \n  optimizer.zero_grad()\n  o = net(xb)\n  loss = criterion(o, yb)\n  loss.backward()\n  optimizer.step()\n  \n  with torch.no_grad():\n      l = loss.item()\n\n      o = o.sigmoid()\n      yb = (yb > 0.5 )*1.0\n      lrap = label_ranking_average_precision_score(yb.cpu().numpy(), o.cpu().numpy())\n\n      o = (o > 0.5)*1.0\n\n      prec = (o*yb).sum()/(1e-6 + o.sum())\n      rec = (o*yb).sum()/(1e-6 + yb.sum())\n      f1 = 2*prec*rec/(1e-6+prec+rec)\n\n  if  scheduler is not None:\n    scheduler.step()\n\n  return l, lrap, f1.item(), rec.item(), prec.item()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef evaluate(net, criterion, val_laoder):\n    net.eval()\n    os, y = [], []\n    val_laoder = tqdm(val_laoder, leave = False, total=len(val_laoder))\n\n    for icount, (xb, yb) in  enumerate(val_laoder):\n        y.append(yb.to(DEVICE))\n        xb = xb.to(DEVICE)\n        o = net(xb)\n        os.append(o)\n\n    y = torch.cat(y)\n    o = torch.cat(os)\n\n    l = criterion(o, y).item()\n    \n    o = o.sigmoid()\n    y = (y > 0.5)*1.0\n\n    lrap = label_ranking_average_precision_score(y.cpu().numpy(), o.cpu().numpy())\n\n    o = (o > 0.5)*1.0\n\n    prec = ((o*y).sum()/(1e-6 + o.sum())).item()\n    rec = ((o*y).sum()/(1e-6 + y.sum())).item()\n    f1 = 2*prec*rec/(1e-6+prec+rec)\n\n    return l, lrap, f1, rec, prec, ","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.908Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AutoSave:\n  def __init__(self, top_k=2, metric=\"f1\", mode=\"min\", root=None, name=\"ckpt\"):\n    self.top_k = top_k\n    self.logs = []\n    self.metric = metric\n    self.mode = mode\n    self.root = Path(root or MODEL_ROOT)\n    assert self.root.exists()\n    self.name = name\n\n    self.top_models = []\n    self.top_metrics = []\n\n  def log(self, model, metrics):\n    metric = metrics[self.metric]\n    rank = self.rank(metric)\n\n    self.top_metrics.insert(rank+1, metric)\n    if len(self.top_metrics) > self.top_k:\n      self.top_metrics.pop(0)\n\n    self.logs.append(metrics)\n    self.save(model, metric, rank, metrics[\"epoch\"])\n\n\n  def save(self, model, metric, rank, epoch):\n    t = time.strftime(\"%Y%m%d%H%M%S\")\n    name = \"{}_epoch_{:02d}_{}_{:.04f}_{}\".format(self.name, epoch, self.metric, metric, t)\n    name = re.sub(r\"[^\\w_-]\", \"\", name) + \".pth\"\n    path = self.root.joinpath(name)\n\n    old_model = None\n    self.top_models.insert(rank+1, name)\n    if len(self.top_models) > self.top_k:\n      old_model = self.root.joinpath(self.top_models[0])\n      self.top_models.pop(0)      \n\n    torch.save(model.state_dict(), path.as_posix())\n\n    if old_model is not None:\n      old_model.unlink()\n\n    self.to_json()\n\n\n  def rank(self, val):\n    r = -1\n    for top_val in self.top_metrics:\n      if val <= top_val:\n        return r\n      r += 1\n\n    return r\n  \n  def to_json(self):\n    # t = time.strftime(\"%Y%m%d%H%M%S\")\n    name = \"{}_logs\".format(self.name)\n    name = re.sub(r\"[^\\w_-]\", \"\", name) + \".json\"\n    path = self.root.joinpath(name)\n\n    with path.open(\"w\") as f:\n      json.dump(self.logs, f, indent=2)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def one_fold(model_name, fold, train_set, val_set, epochs=20, save=True, save_root=None):\n  save_root = Path(save_root) or MODEL_ROOT\n  saver = AutoSave(root=save_root, name=f\"birdclef_{model_name}_fold{fold}\", metric=\"f1_val\")\n  net = get_model(model_name).to(DEVICE)\n  criterion = nn.BCEWithLogitsLoss()\n  optimizer = optim.Adam(net.parameters(), lr=8e-4)\n  scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, eta_min=1e-5, T_max=epochs)\n  train_data = BirdClefDataset(audio_image_store, meta=df.iloc[train_set].reset_index(drop=True),\n                           sr=SR, duration=DURATION, is_train=True)\n  train_laoder = DataLoader(train_data, batch_size=TRAIN_BATCH_SIZE, num_workers=TRAIN_NUM_WORKERS, shuffle=True, pin_memory=True)\n  val_data = BirdClefDataset(audio_image_store, meta=df.iloc[val_set].reset_index(drop=True),  sr=SR, duration=DURATION, is_train=False)\n  val_laoder = DataLoader(val_data, batch_size=VAL_BATCH_SIZE, num_workers=VAL_NUM_WORKERS, shuffle=False)\n  epochs_bar = tqdm(list(range(epochs)), leave=False)\n  for epoch  in epochs_bar:\n    epochs_bar.set_description(f\"--> [EPOCH {epoch:02d}]\")\n    net.train()\n\n    (l, l_val), (lrap, lrap_val), (f1, f1_val), (rec, rec_val), (prec, prec_val) = one_epoch(\n        net=net,\n        criterion=criterion,\n        optimizer=optimizer,\n        scheduler=scheduler,\n        train_laoder=train_laoder,\n        val_laoder=val_laoder,\n      )\n\n    epochs_bar.set_postfix(\n    loss=\"({:.6f}, {:.6f})\".format(l, l_val),\n    prec=\"({:.3f}, {:.3f})\".format(prec, prec_val),\n    rec=\"({:.3f}, {:.3f})\".format(rec, rec_val),\n    f1=\"({:.3f}, {:.3f})\".format(f1, f1_val),\n    lrap=\"({:.3f}, {:.3f})\".format(lrap, lrap_val),\n    )\n\n    print(\n        \"[{epoch:02d}] loss: {loss} lrap: {lrap} f1: {f1} rec: {rec} prec: {prec}\".format(\n            epoch=epoch,\n            loss=\"({:.6f}, {:.6f})\".format(l, l_val),\n            prec=\"({:.3f}, {:.3f})\".format(prec, prec_val),\n            rec=\"({:.3f}, {:.3f})\".format(rec, rec_val),\n            f1=\"({:.3f}, {:.3f})\".format(f1, f1_val),\n            lrap=\"({:.3f}, {:.3f})\".format(lrap, lrap_val),\n        )\n    )\n\n    if save:\n      metrics = {\n          \"loss\": l, \"lrap\": lrap, \"f1\": f1, \"rec\": rec, \"prec\": prec,\n          \"loss_val\": l_val, \"lrap_val\": lrap_val, \"f1_val\": f1_val, \"rec_val\": rec_val, \"prec_val\": prec_val,\n          \"epoch\": epoch,\n      }\n\n      saver.log(net, metrics)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def one_epoch(net, criterion, optimizer, scheduler, train_laoder, val_laoder):\n  net.train()\n  l, lrap, prec, rec, f1, icount = 0.,0.,0.,0., 0., 0\n  train_laoder = tqdm(train_laoder, leave = False)\n  epoch_bar = train_laoder\n  \n  for (xb, yb) in  epoch_bar:\n      _l, _lrap, _f1, _rec, _prec = one_step(xb, yb, net, criterion, optimizer)\n      l += _l\n      lrap += _lrap\n      f1 += _f1\n      rec += _rec\n      prec += _prec\n      icount += 1\n        \n      if hasattr(epoch_bar, \"set_postfix\") and not icount%10:\n          epoch_bar.set_postfix(\n            loss=\"{:.6f}\".format(l/icount),\n            lrap=\"{:.3f}\".format(lrap/icount),\n            prec=\"{:.3f}\".format(prec/icount),\n            rec=\"{:.3f}\".format(rec/icount),\n            f1=\"{:.3f}\".format(f1/icount),\n          )\n  \n  scheduler.step()\n\n  l /= icount\n  lrap /= icount\n  f1 /= icount\n  rec /= icount\n  prec /= icount\n  \n  l_val, lrap_val, f1_val, rec_val, prec_val = evaluate(net, criterion, val_laoder)\n  \n  return (l, l_val), (lrap, lrap_val), (f1, f1_val), (rec, rec_val), (prec, prec_val)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.909Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(model_name, epochs=20, save=True, n_splits=5, seed=177, save_root=None, suffix=\"\", folds=None):\n  gc.collect()\n  torch.cuda.empty_cache()\n\n  save_root = save_root or MODEL_ROOT/f\"{model_name}{suffix}\"\n  save_root.mkdir(exist_ok=True, parents=True)\n  \n  fold_bar = tqdm(df.reset_index().groupby(\"fold\").index.apply(list).items(), total=df.fold.max()+1)\n  \n  for fold, val_set in fold_bar:\n      if folds and not fold in folds:\n        continue\n      \n      print(f\"\\n############################### [FOLD {fold}]\")\n      fold_bar.set_description(f\"[FOLD {fold}]\")\n      train_set = np.setdiff1d(df.index, val_set)\n        \n      one_fold(model_name, fold=fold, train_set=train_set , val_set=val_set , epochs=epochs, save=save, save_root=save_root)\n    \n      gc.collect()\n      torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_NAMES = [\n    \"resnet\",\n] ","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for model_name in MODEL_NAMES:\n  print(\"\\n\\n###########################################\", model_name.upper())\n  try:\n    train(model_name, epochs=30, suffix=f\"_sr{SR}_d{DURATION}_v1_v1\", folds=[0])\n  except Exception as e:\n    # print(f\"Error {model_name} : \\n{e}\")\n    raise ValueError() from  e","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-16T08:23:11.912Z"}},"outputs":[],"execution_count":null}]}