{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":21669,"databundleVersionId":1692278,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Классификация животных в тропическом лесу — RFCx Species Audio Detection\n\n**Соревнование:** Rainforest Connection Species Audio Detection (Kaggle)  \n**Цель:** по каждому аудиофайлу предсказать вероятности присутствия каждого вида (птицы и лягушки).  \n\n**Особенности:**\n- задача **multi-label**: в одном файле может быть *несколько* видов или один\n- много посторонних звуков (насекомые, дождь, ветер) → модель должна научиться **игнорировать шум**\n- разметка `train_tp.csv` содержит **временную локализацию** сигналов (t_min/t_max), `train_fp.csv` — **ложные срабатывания** (hard negatives)\n\n**Метрика:** label-weighted label-ranking average precision (**lwlrap**).  \n\nВ ноутбуке ниже мы:\n1) читаем данные и строим датасет “окна вокруг событий”  \n2) считаем log-mel спектрограммы  \n3) обучаем компактную CNN (без интернета)  \n4) валидируем по lwlrap  \n5) делаем sliding-window инференс и собираем `submission.csv`\n\n\n","metadata":{}},{"cell_type":"code","source":"# === (опционально) timm ===\ntry:\n    import timm\nexcept ImportError:\n    !pip -q install timm\n    import timm\n\nimport os, random, math, warnings\nfrom dataclasses import dataclass\nfrom pathlib import Path\nfrom functools import lru_cache\n\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchaudio\nimport torchaudio.functional as AF\n\nwarnings.filterwarnings(\"ignore\")\n\nprint(\"torch:\", torch.__version__, \"| cuda:\", torch.cuda.is_available())\nprint(\"torchaudio:\", torchaudio.__version__)\n\n@dataclass\nclass CFG:\n    seed: int = 42\n    device: str = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    # audio\n    sr: int = 32000\n    win_sec: float = 8.0          # окно обучения (сек)\n    hop_sec_test: float = 2.0     # шаг на тесте (сек) - меньше => лучше, но медленнее\n\n    # mel\n    n_mels: int = 128\n    n_fft: int = 1024\n    hop_length: int = 320\n    fmin: int = 20\n    fmax: int = 16000\n\n    # sampling\n    fp_ratio: float = 2.0         # сколько FP берем относительно TP\n    bg_ratio: float = 1.0         # сколько background окон относительно TP\n\n    # training\n    train_bs: int = 32\n    valid_bs: int = 64\n    epochs: int = 12\n    lr: float = 3e-4\n    wd: float = 1e-2\n    num_workers: int = 2\n\n    # regularization\n    specaug_p: float = 0.5\n    fp_focus: float = 3.0         # усиление лосса по FP-классу\n\n    # model\n    backbone: str = \"tf_efficientnet_b0.ns_jft_in1k\"  # сильный b0, можно b1/b2\n    pretrained: bool = True\n\ndef set_seed(seed: int):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\nset_seed(CFG.seed)\nprint(\"device:\", CFG.device)\n\nWIN_SAMPLES = int(CFG.sr * CFG.win_sec)\nHOP_SAMPLES_TEST = int(CFG.sr * CFG.hop_sec_test)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:05.831624Z","iopub.execute_input":"2025-12-27T04:10:05.831941Z","iopub.status.idle":"2025-12-27T04:10:05.846568Z","shell.execute_reply.started":"2025-12-27T04:10:05.831916Z","shell.execute_reply":"2025-12-27T04:10:05.845775Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Пути к данным и чтение CSV\n\nОжидаем стандартную структуру Kaggle:\n\n","metadata":{}},{"cell_type":"code","source":"DATA = Path(\"/kaggle/input/rfcx-species-audio-detection\")\nTRAIN_DIR = DATA / \"train\"\nTEST_DIR  = DATA / \"test\"\nTP_CSV = DATA / \"train_tp.csv\"\nFP_CSV = DATA / \"train_fp.csv\"\nSUB_CSV = DATA / \"sample_submission.csv\"\n\nfor p in [DATA, TRAIN_DIR, TEST_DIR, TP_CSV, FP_CSV, SUB_CSV]:\n    assert p.exists(), f\"Не найдено: {p}\"\n\ntp = pd.read_csv(TP_CSV)\nfp = pd.read_csv(FP_CSV)\nsub = pd.read_csv(SUB_CSV)\n\nspecies_cols = [c for c in sub.columns if c != \"recording_id\"]\nnum_classes = len(species_cols)\n\nprint(\"train_tp:\", tp.shape, \"| train_fp:\", fp.shape)\nprint(\"sample_submission:\", sub.shape, \"| num_classes:\", num_classes)\nprint(\"species cols:\", species_cols[:5], \"...\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:05.848020Z","iopub.execute_input":"2025-12-27T04:10:05.848248Z","iopub.status.idle":"2025-12-27T04:10:05.894317Z","shell.execute_reply.started":"2025-12-27T04:10:05.848228Z","shell.execute_reply":"2025-12-27T04:10:05.893814Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1.1 Быстрая проверка аудиофайлов\n\nВ соревновании имена файлов совпадают с `recording_id`. Расширение обычно `.flac`, но на всякий случай ищем среди нескольких вариантов.\n","metadata":{}},{"cell_type":"code","source":"def audio_path(root: Path, recording_id: str) -> Path:\n    for ext in (\".flac\", \".wav\", \".ogg\", \".mp3\"):\n        p = root / f\"{recording_id}{ext}\"\n        if p.exists():\n            return p\n    p = root / recording_id\n    if p.exists():\n        return p\n    raise FileNotFoundError(f\"Не нашёл аудио для {recording_id} в {root}\")\n\nrid0 = tp[\"recording_id\"].iloc[0]\np0 = audio_path(TRAIN_DIR, rid0)\nprint(\"example:\", rid0, \"->\", p0.name)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:05.895058Z","iopub.execute_input":"2025-12-27T04:10:05.895300Z","iopub.status.idle":"2025-12-27T04:10:05.901177Z","shell.execute_reply.started":"2025-12-27T04:10:05.895280Z","shell.execute_reply":"2025-12-27T04:10:05.900466Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Вспомогательные функции: пути, чтение, кроп окна\n","metadata":{}},{"cell_type":"code","source":"def audio_path(root: Path, recording_id: str) -> Path:\n    # В RFCx обычно .flac\n    for ext in (\".flac\", \".wav\", \".ogg\", \".mp3\"):\n        p = root / f\"{recording_id}{ext}\"\n        if p.exists():\n            return p\n    raise FileNotFoundError(f\"Не нашёл аудио для {recording_id} в {root}\")\n\n@lru_cache(maxsize=512)\ndef _load_audio_cached(path_str: str):\n    wav, sr = torchaudio.load(path_str)  # [ch, T]\n    wav = wav.mean(dim=0)                # mono [T]\n    return wav, sr\n    \ndef load_segment(path: Path, start_s: float, win_sec: float, target_sr=CFG.sr) -> torch.Tensor:\n    info = torchaudio.info(str(path))\n    sr0 = info.sample_rate\n    frame_offset = max(0, int(start_s * sr0))\n    num_frames = int(win_sec * sr0)\n\n    wav, sr = torchaudio.load(str(path), frame_offset=frame_offset, num_frames=num_frames)\n    wav = wav.mean(0)\n    if sr != target_sr:\n        wav = AF.resample(wav, sr, target_sr)\n\n    # доводим до ровно WIN_SAMPLES\n    if wav.numel() < WIN_SAMPLES:\n        wav = F.pad(wav, (0, WIN_SAMPLES - wav.numel()))\n    else:\n        wav = wav[:WIN_SAMPLES]\n    return wav\n\ndef load_audio_mono_nocache(path: Path, target_sr=CFG.sr) -> torch.Tensor:\n    wav, sr = torchaudio.load(str(path))\n    wav = wav.mean(dim=0)\n    if sr != target_sr:\n        wav = AF.resample(wav, sr, target_sr)\n    return wav\n\n\ndef crop_or_pad(wav: torch.Tensor, start_s: float, win_samples: int) -> torch.Tensor:\n    start = int(round(start_s * CFG.sr))\n    end = start + win_samples\n    if start < 0:\n        start = 0\n        end = win_samples\n\n    if end <= wav.numel():\n        return wav[start:end]\n\n    out = torch.zeros(win_samples, dtype=wav.dtype)\n    if start < wav.numel():\n        seg = wav[start:]\n        out[:seg.numel()] = seg\n    return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:05.902439Z","iopub.execute_input":"2025-12-27T04:10:05.902681Z","iopub.status.idle":"2025-12-27T04:10:05.917059Z","shell.execute_reply.started":"2025-12-27T04:10:05.902656Z","shell.execute_reply":"2025-12-27T04:10:05.916475Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Разметка событий + labels для окна\n","metadata":{}},{"cell_type":"code","source":"# события TP по записи\ntp_by_rec = {}\nfor r in tp.itertuples(index=False):\n    tp_by_rec.setdefault(r.recording_id, []).append((float(r.t_min), float(r.t_max), int(r.species_id)))\n\ndef labels_for_window(events, t0: float, t1: float, num_classes: int) -> np.ndarray:\n    y = np.zeros(num_classes, dtype=np.float32)\n    for a, b, sid in events:\n        if a < t1 and b > t0:  # overlap\n            if 0 <= sid < num_classes:\n                y[sid] = 1.0\n    return y\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:05.934436Z","iopub.execute_input":"2025-12-27T04:10:05.934668Z","iopub.status.idle":"2025-12-27T04:10:05.943011Z","shell.execute_reply.started":"2025-12-27T04:10:05.934626Z","shell.execute_reply":"2025-12-27T04:10:05.942440Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Сэмплинг train окон: TP + FP + background\n","metadata":{}},{"cell_type":"code","source":"def build_samples(tp: pd.DataFrame, fp: pd.DataFrame,\n                  fp_ratio=CFG.fp_ratio, bg_ratio=CFG.bg_ratio) -> pd.DataFrame:\n    rows = []\n\n    # --- TP окна: вокруг центра события\n    for r in tp.itertuples(index=False):\n        center = 0.5*(float(r.t_min) + float(r.t_max))\n        t0 = center - CFG.win_sec/2\n        t1 = t0 + CFG.win_sec\n        rows.append((r.recording_id, t0, t1, \"tp\", -1))\n\n    # --- FP hard negatives: берём подвыборку\n    n_fp = int(len(tp) * fp_ratio)\n    fp_s = fp.sample(n=min(n_fp, len(fp)), random_state=CFG.seed)\n    for r in fp_s.itertuples(index=False):\n        center = 0.5*(float(r.t_min) + float(r.t_max))\n        t0 = center - CFG.win_sec/2\n        t1 = t0 + CFG.win_sec\n        rows.append((r.recording_id, t0, t1, \"fp\", int(r.species_id)))\n\n    # --- Background: случайные окна из тех же записей (учим “шум/насекомых”)\n    # RFCx записи обычно ~60с, поэтому для простоты выбираем t0 в [0, 60-win]\n    n_bg = int(len(tp) * bg_ratio)\n    recs = tp[\"recording_id\"].unique()\n    for _ in range(n_bg):\n        rid = recs[np.random.randint(0, len(recs))]\n        max_t0 = max(0.0, 60.0 - CFG.win_sec)\n        t0 = float(np.random.uniform(0.0, max_t0))\n        t1 = t0 + CFG.win_sec\n        rows.append((rid, t0, t1, \"bg\", -1))\n\n    df = pd.DataFrame(rows, columns=[\"recording_id\",\"t0\",\"t1\",\"kind\",\"fp_species_id\"])\n    return df.sample(frac=1.0, random_state=CFG.seed).reset_index(drop=True)\n\nsamples = build_samples(tp, fp)\nprint(\"samples:\", samples.shape)\nprint(samples[\"kind\"].value_counts())\nsamples.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:05.944629Z","iopub.execute_input":"2025-12-27T04:10:05.945041Z","iopub.status.idle":"2025-12-27T04:10:05.990044Z","shell.execute_reply.started":"2025-12-27T04:10:05.945022Z","shell.execute_reply":"2025-12-27T04:10:05.989505Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Train/Valid split без утечки (по recording_id)","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import GroupShuffleSplit\n\ngss = GroupShuffleSplit(n_splits=1, test_size=0.15, random_state=CFG.seed)\ntr_idx, va_idx = next(gss.split(samples, groups=samples[\"recording_id\"]))\n\ntrain_s = samples.iloc[tr_idx].reset_index(drop=True)\nvalid_s = samples.iloc[va_idx].reset_index(drop=True)\n\nprint(\"train:\", len(train_s), \"| valid:\", len(valid_s))\nprint(\"unique rec train/valid:\", train_s.recording_id.nunique(), \"/\", valid_s.recording_id.nunique())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:05.990783Z","iopub.execute_input":"2025-12-27T04:10:05.990954Z","iopub.status.idle":"2025-12-27T04:10:06.002107Z","shell.execute_reply.started":"2025-12-27T04:10:05.990939Z","shell.execute_reply":"2025-12-27T04:10:06.001480Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Mel на GPU + SpecAugment (важно для скорости и качества)","metadata":{}},{"cell_type":"code","source":"import torchaudio\nimport torch\n\n# Создай один раз (вне цикла), чтобы не пересоздавать каждый батч\n_mel = torchaudio.transforms.MelSpectrogram(\n    sample_rate=CFG.sr,\n    n_fft=CFG.n_fft,\n    hop_length=CFG.hop_length,\n    n_mels=CFG.n_mels,\n    f_min=CFG.fmin,\n    f_max=CFG.fmax,\n    power=2.0,\n).to(CFG.device)\n\n_db = torchaudio.transforms.AmplitudeToDB(stype=\"power\", top_db=80.0).to(CFG.device)\n\ndef wav_to_logmel_torch(wav: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    wav: [B, WIN] float32 на GPU\n    return: [B, 1, M, T] float32 на GPU\n    \"\"\"\n    # MelSpectrogram ожидает [B, T]\n    S = _mel(wav)                 # [B, M, T] power\n    S = _db(S)                    # [B, M, T] dB, без -inf\n\n    # z-norm per-sample (очень важно eps)\n    mean = S.mean(dim=(1,2), keepdim=True)\n    std  = S.std(dim=(1,2), keepdim=True).clamp_min(1e-4)\n    S = (S - mean) / std\n\n    # финальная защита\n    S = torch.nan_to_num(S, nan=0.0, posinf=0.0, neginf=0.0)\n\n    return S.unsqueeze(1)         # [B,1,M,T]\n\n\ndef spec_augment_torch(x: torch.Tensor, p=0.5, max_mask_pct=0.10, num_masks=2) -> torch.Tensor:\n    \"\"\"\n    x: [B,1,M,T] на GPU\n    \"\"\"\n    if p <= 0 or (torch.rand(1).item() > p):\n        return x\n    B, C, M, T = x.shape\n    x = x.clone()\n    for _ in range(num_masks):\n        if torch.rand(1).item() < 0.5:\n            # freq mask\n            f = max(1, int(max_mask_pct * M))\n            f0 = torch.randint(0, max(1, M - f + 1), (1,)).item()\n            x[:, :, f0:f0+f, :] = 0\n        else:\n            # time mask\n            t = max(1, int(max_mask_pct * T))\n            t0 = torch.randint(0, max(1, T - t + 1), (1,)).item()\n            x[:, :, :, t0:t0+t] = 0\n    return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:06.002923Z","iopub.execute_input":"2025-12-27T04:10:06.003110Z","iopub.status.idle":"2025-12-27T04:10:06.020756Z","shell.execute_reply.started":"2025-12-27T04:10:06.003093Z","shell.execute_reply":"2025-12-27T04:10:06.020138Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Dataset: возвращаем waveform + y + mask/weights (без librosa, быстро)\n","metadata":{}},{"cell_type":"code","source":"class RFCXWindowDataset(Dataset):\n    def __init__(self, df: pd.DataFrame, train: bool):\n        self.df = df.reset_index(drop=True)\n        self.train = train\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        row = self.df.iloc[i]\n        rid = row.recording_id\n        kind = row.kind\n\n        wav_full = load_audio_mono_nocache(audio_path(TRAIN_DIR, rid))       # torch [T] CPU\n        wav = crop_or_pad(wav_full, float(row.t0), WIN_SAMPLES)      # torch [WIN] CPU\n\n        # label multi-hot по пересечению с TP\n        events = tp_by_rec.get(rid, [])\n        y = labels_for_window(events, float(row.t0), float(row.t1), num_classes)  # np [C]\n\n        # weights: по умолчанию 1\n        w = np.ones(num_classes, dtype=np.float32)\n\n        if kind == \"fp\":\n            # hard negative: усиливаем штраф только по fp_species_id\n            sid = int(row.fp_species_id)\n            y[sid] = 0.0\n            w[:] = 0.0\n            w[sid] = CFG.fp_focus\n\n        # возвращаем waveform (CPU), y/w (CPU) -> на GPU перенесем батчем\n        return wav.numpy().astype(np.float32), y, w\n        \nCFG.num_workers = 0\n\ntrain_loader = DataLoader(RFCXWindowDataset(train_s, True),\n                          batch_size=CFG.train_bs, shuffle=True,\n                          num_workers=CFG.num_workers, pin_memory=True,\n                          drop_last=True, persistent_workers=False)\n\nvalid_loader = DataLoader(RFCXWindowDataset(valid_s, False),\n                          batch_size=CFG.valid_bs, shuffle=False,\n                          num_workers=CFG.num_workers, pin_memory=True,\n                          persistent_workers=False)\n\nxb0, yb0, wb0 = next(iter(train_loader))\nprint(\"wave batch:\", xb0.shape, \"| y:\", yb0.shape, \"| w:\", wb0.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:06.022497Z","iopub.execute_input":"2025-12-27T04:10:06.022877Z","iopub.status.idle":"2025-12-27T04:10:10.472363Z","shell.execute_reply.started":"2025-12-27T04:10:06.022858Z","shell.execute_reply":"2025-12-27T04:10:10.470364Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Модель: EfficientNet (timm) под 1-канальную mel","metadata":{}},{"cell_type":"code","source":"class SpecNet(nn.Module):\n    def __init__(self, num_classes: int):\n        super().__init__()\n        self.backbone = timm.create_model(\n            CFG.backbone,\n            pretrained=CFG.pretrained,  # интернет нужен только здесь\n            in_chans=1,\n            num_classes=num_classes\n        )\n\n    def forward(self, x):\n        return self.backbone(x)\n\nmodel = SpecNet(num_classes).to(CFG.device)\nprint(\"params:\", sum(p.numel() for p in model.parameters())/1e6, \"M\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:10.474188Z","iopub.execute_input":"2025-12-27T04:10:10.474546Z","iopub.status.idle":"2025-12-27T04:10:10.796832Z","shell.execute_reply.started":"2025-12-27T04:10:10.474497Z","shell.execute_reply":"2025-12-27T04:10:10.795516Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Метрика lwlrap","metadata":{}},{"cell_type":"code","source":"def _one_sample_lwlrap(truth, scores):\n    pos = np.where(truth > 0)[0]\n    if len(pos) == 0:\n        return (None, None)\n    rank = scores.argsort()[::-1]\n    prec = []\n    hit = 0\n    for i, k in enumerate(rank, start=1):\n        if truth[k] > 0:\n            hit += 1\n            prec.append(hit / i)\n    return (pos, np.array(prec, dtype=np.float32))\n\ndef lwlrap(truth, scores):\n    C = truth.shape[1]\n    per_class_prec = [[] for _ in range(C)]\n    for t, s in zip(truth, scores):\n        res = _one_sample_lwlrap(t, s)\n        if res[0] is None:\n            continue\n        pos, prec = res\n        for cls, p in zip(pos, prec):\n            per_class_prec[cls].append(p)\n\n    per_class_lwlrap = np.array([np.mean(v) if len(v) else 0.0 for v in per_class_prec], dtype=np.float32)\n    weights = truth.sum(axis=0)\n    weights = weights / (weights.sum() + 1e-12)\n    return float((per_class_lwlrap * weights).sum())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:10.798477Z","iopub.execute_input":"2025-12-27T04:10:10.798914Z","iopub.status.idle":"2025-12-27T04:10:10.814563Z","shell.execute_reply.started":"2025-12-27T04:10:10.798874Z","shell.execute_reply":"2025-12-27T04:10:10.813723Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Loss/оптимизация: BCE + class pos_weight + FP weights + AMP","metadata":{}},{"cell_type":"code","source":"# pos_weight по train_s (примерно)\ndef estimate_pos_weight(df: pd.DataFrame) -> torch.Tensor:\n    Ys = []\n    for r in df.sample(n=min(1200, len(df)), random_state=CFG.seed).itertuples(index=False):\n        events = tp_by_rec.get(r.recording_id, [])\n        y = labels_for_window(events, float(r.t0), float(r.t1), num_classes)\n        Ys.append(y)\n    Y = np.stack(Ys)\n    pos = Y.sum(axis=0)\n    neg = len(Y) - pos\n    pw = (neg + 1e-6) / (pos + 1e-6)\n    return torch.tensor(pw, device=CFG.device).float()\n\npos_weight = estimate_pos_weight(train_s)\nbce = nn.BCEWithLogitsLoss(reduction=\"none\", pos_weight=pos_weight)\n\nscaler = torch.amp.GradScaler('cuda', enabled=(CFG.device==\"cuda\"))\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.wd)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:10.816212Z","iopub.execute_input":"2025-12-27T04:10:10.816945Z","iopub.status.idle":"2025-12-27T04:10:10.847207Z","shell.execute_reply.started":"2025-12-27T04:10:10.816902Z","shell.execute_reply":"2025-12-27T04:10:10.846550Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Train/Valid цикл","metadata":{}},{"cell_type":"code","source":"def to_tensor(x, device):\n    if torch.is_tensor(x):\n        return x.to(device, non_blocking=True)\n    return torch.from_numpy(x).to(device, non_blocking=True)\n\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\n\nbce = torch.nn.BCEWithLogitsLoss(reduction=\"none\")\n\ndef run_epoch(loader, train: bool):\n    model.train(train)\n    losses = []\n    all_y, all_p = [], []\n\n    for wav, y, w in loader:\n        # wav: [B, WIN]  y: [B,C]  w: [B,C]\n        wav = wav.to(CFG.device, non_blocking=True).float()\n        y   = y.to(CFG.device, non_blocking=True).float()\n        w   = w.to(CFG.device, non_blocking=True).float()\n\n        if train:\n            optimizer.zero_grad(set_to_none=True)\n\n        with torch.amp.autocast(device_type=\"cuda\", enabled=(CFG.device == \"cuda\")):\n            x = wav_to_logmel_torch(wav)          # [B,1,M,T] GPU\n            if train:\n                x = spec_augment_torch(x, p=CFG.specaug_p)\n\n            logits = model(x)                     # [B,C]\n\n            # БАЗОВЫЙ СТАБИЛЬНЫЙ ЛОСС (без бутстрэпа) — сначала доведи до finite!\n            loss_mat = bce(logits, y)             # [B,C]\n            denom = w.sum().clamp_min(1.0)        # защита от 0\n            loss = (loss_mat * w).sum() / denom\n\n        # защита от NaN/Inf\n        if not torch.isfinite(loss):\n            print(\"⚠️ non-finite loss, skip batch\")\n            continue\n\n        if train:\n            scaler.scale(loss).backward()\n            # клиппинг до step (важно)\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n\n            scaler.step(optimizer)\n            scaler.update()\n\n        losses.append(loss.item())\n        all_y.append(y.detach().cpu().numpy())\n        all_p.append(logits.detach().cpu().numpy())\n\n    all_y = np.concatenate(all_y) if len(all_y) else np.zeros((0, num_classes), np.float32)\n    all_p = np.concatenate(all_p) if len(all_p) else np.zeros((0, num_classes), np.float32)\n\n    return float(np.mean(losses)) if losses else float(\"nan\"), lwlrap(all_y, all_p)\n\n\n\nbest = -1.0\nfor e in range(1, CFG.epochs + 1):\n    tr_loss, tr_lwl = run_epoch(train_loader, True)\n    va_loss, va_lwl = run_epoch(valid_loader, False)\n    scheduler.step()\n\n    print(f\"Epoch {e:02d} | train loss {tr_loss:.4f} lwlrap {tr_lwl:.4f} | valid loss {va_loss:.4f} lwlrap {va_lwl:.4f}\")\n\n    if va_lwl > best:\n        best = va_lwl\n        torch.save({\"model\": model.state_dict(), \"species_cols\": species_cols}, \"best_model.pt\")\n\nprint(\"Best valid lwlrap:\", best)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:10:10.848208Z","iopub.execute_input":"2025-12-27T04:10:10.848483Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Инференс + submission.csv (sliding-window + max-pool)","metadata":{}},{"cell_type":"code","source":"import gc, torch\n\n_load_audio_cached.cache_clear()\ndel train_loader, valid_loader\ngc.collect()\ntorch.cuda.empty_cache()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.inference_mode()\ndef predict_recording(path: Path) -> np.ndarray:\n    wav = load_audio_mono_nocache(path)\n    T = wav.numel()\n\n    # нарезаем окна с hop_sec_test\n    starts = list(range(0, max(1, T - WIN_SAMPLES + 1), HOP_SAMPLES_TEST))\n    if len(starts) == 0:\n        starts = [0]\n\n    probs_max = torch.zeros(num_classes, device=CFG.device)\n\n    bs = 32\n    for i in range(0, len(starts), bs):\n        batch_starts = starts[i:i+bs]\n        batch = []\n        for s in batch_starts:\n            seg = wav[s:s+WIN_SAMPLES]\n            if seg.numel() < WIN_SAMPLES:\n                pad = torch.zeros(WIN_SAMPLES - seg.numel())\n                seg = torch.cat([seg, pad], dim=0)\n            batch.append(seg)\n\n        batch = torch.stack(batch, dim=0).to(CFG.device)  # [B, WIN]\n        x = wav_to_logmel_torch(batch)                    # [B,1,M,T]\n        logits = model(x)\n        probs = torch.sigmoid(logits).max(dim=0).values   # max-pool по окнам\n        probs_max = torch.maximum(probs_max, probs)\n\n    return probs_max.detach().cpu().numpy()\n\n# загружаем лучший чекпойнт\nckpt = torch.load(\"best_model.pt\", map_location=\"cpu\")\nmodel.load_state_dict(ckpt[\"model\"], strict=True)\nmodel.eval()\n\nrows = []\nfor rid in sub[\"recording_id\"].tolist():\n    p = audio_path(TEST_DIR, rid)\n    probs = predict_recording(p)\n    rows.append([rid] + probs.tolist())\n\nout = pd.DataFrame(rows, columns=[\"recording_id\"] + species_cols)\nout.to_csv(\"submission.csv\", index=False)\n\nprint(\"Saved submission.csv:\", out.shape)\nout.head()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}