{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Ячейка 1 — конфиг и загрузка v0.4\n\nfrom pathlib import Path\nimport os, re, json, random, math, time, warnings\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nwarnings.filterwarnings(\"ignore\")\npd.set_option(\"display.max_columns\", 100)\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\nBASE = Path(\"/kaggle/input\")\nCOMP = None\nfor p in BASE.rglob(\"sample_submission.csv\"):\n    COMP = p.parent\n    break\nif COMP is None:\n    raise FileNotFoundError(\"Не нашёл sample_submission.csv. Проверь Add Data.\")\n\ntrain = pd.read_csv(COMP / \"train.csv\")\ntest = pd.read_csv(COMP / \"test.csv\")\ntrain_series = pd.read_csv(COMP / \"train_series.csv\")\ntest_series = pd.read_csv(COMP / \"test_series.csv\")\nsample_sub = pd.read_csv(COMP / \"sample_submission.csv\")\n\nLABELS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n    \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n]\n\n# v0.4 config: уже не совсем smoke, но ещё контролируемо.\nIMG_SIZE = 384","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-14T21:30:27.826423Z","iopub.execute_input":"2026-08-14T21:30:27.826804Z","iopub.status.idle":"2026-08-14T21:30:28.448804Z","shell.execute_reply.started":"2026-08-14T21:30:27.826766Z","shell.execute_reply":"2026-08-14T21:30:28.447934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ячейка 2 — weak labels v1.1: negation-aware, чуть лучше OA и строже Synovitis\n\nNEG_RE = re.compile(\n    r\"(\\bno\\b|\\bnot\\b|\\bwithout\\b|\\bsin\\b|\\bno hay\\b|\\bkein\\b|\\bkeine\\b|\\baucun\\b|\\baucune\\b|\"\n    r\"\\bgeen\\b|\\bniet\\b|\\bsem\\b|\\bnon\\b|\\bнет\\b|\\bбез\\b|\\bintact\\b|\\bnormal\\b|\\bwithin normal\\b|\"\n    r\"\\bunremarkable\\b|\\bconservad)\",\n    re.I\n)\n\nRULES = {\n    \"ACL\": [r\"\\bacl\\b\", r\"\\blca\\b\", r\"cruzado anterior\", r\"croisé antérieur\", r\"vorderes kreuzband\", r\"крестообраз\"],\n    \"MCL\": [r\"\\bmcl\\b\", r\"colateral medial\", r\"collatéral médial\", r\"mediales kollateral\", r\"медиальн.{0,20}коллатерал\"],\n    \"Medial Meniscus\": [r\"medial meniscus\", r\"meniscus medial\", r\"menisco medial\", r\"menisco interno\", r\"innenmeniskus\", r\"медиальн.{0,20}мениск\"],\n    \"Lateral Meniscus\": [r\"lateral meniscus\", r\"meniscus lateral\", r\"menisco lateral\", r\"menisco externo\", r\"außenmeniskus\", r\"латеральн.{0,20}мениск\"],\n\n    # OA: расширили, но всё ещё стараемся ловить compartment/medial/lateral.\n    \"Medial OA\": [\n        r\"medial.{0,40}osteoarth\", r\"medial.{0,40}arthrosis\", r\"medial.{0,40}arthrose\",\n        r\"osteoarth.{0,40}medial\", r\"arthrosis.{0,40}medial\", r\"arthrose.{0,40}medial\",\n        r\"artrosis.{0,40}medial\", r\"femorotibial medial\", r\"medial.{0,30}gonarthrosis\",\n        r\"медиальн.{0,40}остеоартр\"\n    ],\n    \"Lateral OA\": [\n        r\"lateral.{0,40}osteoarth\", r\"lateral.{0,40}arthrosis\", r\"lateral.{0,40}arthrose\",\n        r\"osteoarth.{0,40}lateral\", r\"arthrosis.{0,40}lateral\", r\"arthrose.{0,40}lateral\",\n        r\"artrosis.{0,40}lateral\", r\"femorotibial lateral\", r\"lateral.{0,30}gonarthrosis\",\n        r\"латеральн.{0,40}остеоартр\"\n    ],\n    \"PF OA\": [\n        r\"patellofemoral.{0,40}(osteoarth|arthrosis|arthrose|chondrop|chondros)\",\n        r\"(osteoarth|arthrosis|arthrose|chondrop|chondros).{0,40}patellofemoral\",\n        r\"femoropatelar\", r\"rétropatellaire\", r\"пателлофеморал\"\n    ],\n\n    \"Effusion\": [r\"effusion\", r\"derrame\", r\"erguss\", r\"épanchement\", r\"versamento\", r\"joint fluid\", r\"выпот\"],\n    \"Synovitis\": [r\"\\bsynovitis\\b\", r\"\\bsinovitis\\b\", r\"синовит\"],\n    \"Baker's\": [r\"baker\", r\"popliteal cyst\", r\"quiste popl\", r\"poplitea\", r\"беккер\", r\"бейкер\"],\n    \"Contusion\": [r\"contusion\", r\"contusión\", r\"bone bruise\", r\"bone marrow edema\", r\"edema óseo\", r\"костномозгов\"],\n    \"Fracture\": [r\"fracture\", r\"fractura\", r\"fraktur\", r\"перелом\"],\n}\n\ndef split_sentences(text):\n    return [s.strip() for s in re.split(r\"[\\.\\!\\?\\n\\r]+\", str(text).lower()) if s.strip()]\n\ndef weak_one(text):\n    hits = {lab: {\"pos\": 0, \"neg\": 0} for lab in LABELS}\n    for s in split_sentences(text):\n        neg = bool(NEG_RE.search(s))\n        for lab in LABELS:\n            if any(re.search(p, s, flags=re.I) for p in RULES[lab]):\n                hits[lab][\"neg\" if neg else \"pos\"] += 1\n\n    out = {}\n    for lab in LABELS:\n        if hits[lab][\"pos\"] > 0:\n            out[lab] = 1.0\n        elif hits[lab][\"neg\"] > 0:\n            out[lab] = 0.0\n        else:\n            out[lab] = np.nan\n    return out\n\nprint(\"Извлекаем weak labels...\", flush=True)\nt0 = time.time()\nweak = pd.DataFrame([weak_one(x) for x in train[\"Report\"].fillna(\"\")])\nweak.insert(0, \"StudyInstanceUID\", train[\"StudyInstanceUID\"].values)\nweak.to_csv(\"/kaggle/working/train_weak_v1_1.csv\", index=False)\nprint(f\"weak labels done in {time.time() - t0:.1f}s\", flush=True)\n\nstats = pd.DataFrame({\n    \"known_frac\": weak[LABELS].notna().mean(),\n    \"pos_rate_among_known\": weak[LABELS].mean(skipna=True),\n}).sort_values(\"known_frac\", ascending=False)\n\nprint(stats.round(4), flush=True)\nprint(\"\\nknown labels per study:\", flush=True)\nprint(weak[LABELS].notna().sum(axis=1).describe(), flush=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-14T21:30:28.450489Z","iopub.execute_input":"2026-08-14T21:30:28.451088Z","iopub.status.idle":"2026-08-14T21:30:38.876475Z","shell.execute_reply.started":"2026-08-14T21:30:28.451059Z","shell.execute_reply":"2026-08-14T21:30:38.875492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ячейка 3 — выбор серий: одна серия на плоскость\n\nPLANES = [\"Axial\", \"Sagittal\", \"Coronal\"]\n\ndef select_series(sdf, split, desc=\"\"):\n    rows = []\n    studies = list(sdf.groupby(\"StudyInstanceUID\"))\n    t0 = time.time()\n\n    for i, (study, x) in enumerate(studies):\n        rec = {\"study\": study, \"split\": split}\n        for plane in PLANES:\n            xp = x[x[\"Anatomical_Plane\"] == plane].copy()\n            if len(xp) == 0:\n                rec[plane] = \"\"\n                continue\n            xp[\"score\"] = xp[\"Fluid_Sensitive\"].fillna(0) * 2 + xp[\"Fat_Suppression\"].fillna(0)\n            r = xp.sort_values(\"score\", ascending=False).iloc[0]\n            rec[plane] = str(COMP / f\"{split}_series\" / str(study) / str(r[\"SeriesInstanceUID\"]))\n        rows.append(rec)\n\n        if desc and (i % 1000 == 0 or i == len(studies) - 1):\n            print(f\"[{desc}] {i + 1}/{len(studies)} studies | elapsed={time.time() - t0:.1f}s\", flush=True)\n\n    return pd.DataFrame(rows)\n\ntrain_sel = select_series(train_series, \"train\", \"select train\")\ntest_sel = select_series(test_series, \"test\", \"select test\")\n\ntrain_df = train_sel.merge(weak, left_on=\"study\", right_on=\"StudyInstanceUID\", how=\"left\").drop(columns=[\"StudyInstanceUID\"])\ntest_df = test_sel.copy()\n\nprint(\"train_df:\", train_df.shape, \"test_df:\", test_df.shape, flush=True)\ndisplay(train_df.head(3))\n\ntrain_df.to_csv(\"/kaggle/working/train_selected_series_v0_4.csv\", index=False)\ntest_df.to_csv(\"/kaggle/working/test_selected_series_v0_4.csv\", index=False)\nprint(\"saved selected series csv\", flush=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-14T21:30:38.877713Z","iopub.execute_input":"2026-08-14T21:30:38.878077Z","iopub.status.idle":"2026-08-14T21:30:53.346119Z","shell.execute_reply.started":"2026-08-14T21:30:38.878042Z","shell.execute_reply":"2026-08-14T21:30:53.344656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ячейка 4 — dataset и MIL-модель, self-contained fallback\n\nfrom pathlib import Path\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\nfrom torchvision import models\nfrom PIL import Image\nimport pydicom\n\n# fallback-конфиг, если Ячейка 1 не была выполнена в этом kernel\nif \"SEED\" not in globals():\n    SEED = 42\nif \"IMG_SIZE\" not in globals():\n    IMG_SIZE = 384\nif \"SLICES_PER_SERIES\" not in globals():\n    SLICES_PER_SERIES = 14\nif \"AMP\" not in globals():\n    AMP = True\nif \"LABELS\" not in globals():\n    LABELS = [\n        \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n        \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n        \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n    ]\nif \"PLANES\" not in globals():\n    PLANES = [\"Axial\", \"Sagittal\", \"Coronal\"]\nif \"COMP\" not in globals():\n    BASE = Path(\"/kaggle/input\")\n    COMP = None\n    for p in BASE.rglob(\"sample_submission.csv\"):\n        COMP = p.parent\n        break\n    if COMP is None:\n        raise FileNotFoundError(\"Не нашёл sample_submission.csv. Проверь Add Data.\")\n\ntorch.manual_seed(SEED)\ntorch.backends.cudnn.benchmark = True\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nUSE_AMP = bool(AMP and DEVICE == \"cuda\")\nprint(\"DEVICE:\", DEVICE, \"USE_AMP:\", USE_AMP, flush=True)\nprint(\"fallback config:\", {\n    \"IMG_SIZE\": IMG_SIZE,\n    \"SLICES_PER_SERIES\": SLICES_PER_SERIES,\n    \"AMP\": AMP,\n    \"SEED\": SEED,\n}, flush=True)\n\ntry:\n    from pydicom.pixel_data_handlers.util import apply_voi_lut\nexcept Exception:\n    apply_voi_lut = None\n\nRESAMPLE = getattr(getattr(Image, \"Resampling\", Image), \"BILINEAR\")\n\ndef read_dicom(path):\n    try:\n        ds = pydicom.dcmread(str(path))\n        arr = ds.pixel_array\n\n        if apply_voi_lut is not None:\n            try:\n                arr = apply_voi_lut(arr, ds)\n            except Exception:\n                pass\n\n        arr = arr.astype(np.float32)\n\n        if arr.ndim == 3:\n            arr = arr[..., 0] if arr.shape[-1] in [3, 4] else arr[0]\n\n        if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n            arr = arr.max() - arr\n\n        arr = (arr - arr.min()) / (arr.max() - arr.min() + 1e-6)\n\n        h, w = arr.shape\n        m = max(h, w)\n        canvas = np.zeros((m, m), dtype=np.float32)\n        y0 = (m - h) // 2\n        x0 = (m - w) // 2\n        canvas[y0:y0 + h, x0:x0 + w] = arr\n\n        img = Image.fromarray((canvas * 255).astype(np.uint8)).resize((IMG_SIZE, IMG_SIZE), RESAMPLE)\n        return np.asarray(img).astype(np.float32) / 255.0\n    except Exception:\n        return np.zeros((IMG_SIZE, IMG_SIZE), dtype=np.float32)\n\ndef series_tensor(path):\n    zero = torch.zeros((SLICES_PER_SERIES, 1, IMG_SIZE, IMG_SIZE), dtype=torch.float32)\n    if not path:\n        return zero\n\n    files = sorted(Path(path).glob(\"*.dcm\"))\n    if len(files) == 0:\n        return zero\n\n    idx = np.linspace(0, len(files) - 1, SLICES_PER_SERIES).round().astype(int)\n    imgs = [read_dicom(files[i]) for i in idx]\n    arr = np.stack(imgs, axis=0)  # S,H,W\n    arr = (arr - 0.5) / 0.5\n    return torch.from_numpy(arr).float().unsqueeze(1)  # S,1,H,W\n\nclass KneeStudyDS(Dataset):\n    def __init__(self, df, has_labels=True):\n        self.df = df.reset_index(drop=True)\n        self.has_labels = has_labels\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        r = self.df.iloc[i]\n        xs, pm = [], []\n\n        for plane in PLANES:\n            path = r.get(plane, \"\")\n            if isinstance(path, str) and len(path) > 0 and Path(path).exists():\n                xs.append(series_tensor(path))\n                pm.append(1.0)\n            else:\n                xs.append(torch.zeros((SLICES_PER_SERIES, 1, IMG_SIZE, IMG_SIZE), dtype=torch.float32))\n                pm.append(0.0)\n\n        x = torch.stack(xs, dim=0)  # 3,S,1,H,W\n        plane_mask = torch.tensor(pm, dtype=torch.float32)\n\n        if self.has_labels:\n            y = torch.tensor([r.get(c, np.nan) for c in LABELS], dtype=torch.float32)\n        else:\n            y = torch.full((len(LABELS),), float(\"nan\"))\n\n        return x, plane_mask, y\n\nclass MILNet(nn.Module):\n    def __init__(self, n_labels=12):\n        super().__init__()\n        self.backbone = models.resnet18(weights=None)\n        self.backbone.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        feat_dim = self.backbone.fc.in_features\n        self.backbone.fc = nn.Identity()\n        self.head = nn.Linear(feat_dim, n_labels)\n\n    def forward(self, x, plane_mask):\n        # x: B,3,S,1,H,W\n        B, P, S, C, H, W = x.shape\n        slice_valid = (x.abs().sum(dim=(3, 4, 5)) > 0).float()  # B,P,S\n\n        z = x.reshape(B * P * S, C, H, W)\n        f = self.backbone(z).reshape(B, P, S, -1)\n\n        w = slice_valid.unsqueeze(-1)  # B,P,S,1\n        series_feat = (f * w).sum(dim=2) / w.sum(dim=2).clamp_min(1.0)  # B,P,512\n\n        pm = plane_mask.unsqueeze(-1)  # B,P,1\n        study_feat = (series_feat * pm).sum(dim=1) / pm.sum(dim=1).clamp_min(1.0)\n        return self.head(study_feat)\n\ndef masked_bce(logits, y, pos_weight=None):\n    m = ~torch.isnan(y)\n    y2 = torch.nan_to_num(y, nan=0.0)\n    loss = F.binary_cross_entropy_with_logits(\n        logits, y2, reduction=\"none\", pos_weight=pos_weight\n    )\n    return (loss * m).sum() / m.sum().clamp_min(1.0)\n\nmodel = MILNet(len(LABELS)).to(DEVICE)\nopt = torch.optim.AdamW(model.parameters(), lr=2e-4)\nprint(\"model ready, params:\", round(sum(p.numel() for p in model.parameters()) / 1e6, 3), \"M\", flush=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-14T21:30:53.34694Z","iopub.status.idle":"2026-08-14T21:30:53.347283Z","shell.execute_reply.started":"2026-08-14T21:30:53.347113Z","shell.execute_reply":"2026-08-14T21:30:53.347132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ячейка 5 — train v0.4 с pos_weight, логированием и ETA, self-contained config fallback\n\nfrom sklearn.metrics import roc_auc_score\nimport time\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\n\n# fallback-конфиг, если Ячейка 1 не выполнялась в этом kernel\nif \"SEED\" not in globals():\n    SEED = 42\nif \"N_TRAIN_STUDIES\" not in globals():\n    N_TRAIN_STUDIES = 1024\nif \"EPOCHS\" not in globals():\n    EPOCHS = 4\nif \"BATCH_SIZE\" not in globals():\n    BATCH_SIZE = 8\nif \"NUM_WORKERS\" not in globals():\n    NUM_WORKERS = 0\nif \"AMP\" not in globals():\n    AMP = True\nif \"GRAD_CLIP\" not in globals():\n    GRAD_CLIP = 1.0\nif \"VAL_FRAC\" not in globals():\n    VAL_FRAC = 0.10\nif \"LABELS\" not in globals():\n    LABELS = [\n        \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n        \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n        \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n    ]\n\nneed_from_cell4 = [\"KneeStudyDS\", \"masked_bce\", \"model\", \"opt\", \"DEVICE\"]\nmissing = [x for x in need_from_cell4 if x not in globals()]\nif missing:\n    raise NameError(f\"Не хватает объектов из Ячейки 4: {missing}. Сначала выполни Ячейку 4.\")\n\nif \"train_df\" not in globals():\n    p = \"/kaggle/working/train_selected_series_v0_4.csv\"\n    if Path(p).exists():\n        train_df = pd.read_csv(p)\n        print(\"loaded train_df from\", p, train_df.shape, flush=True)\n    else:\n        raise NameError(\"Нет train_df. Выполни Ячейки 1–3, чтобы собрать train_selected_series_v0_4.csv.\")\n\nUSE_AMP = bool(AMP and DEVICE == \"cuda\")\n\ndef fmt_eta(seconds):\n    seconds = int(max(0, seconds))\n    return f\"{seconds // 60:02d}:{seconds % 60:02d}\"\n\nprint(\"config fallback:\", {\n    \"N_TRAIN_STUDIES\": N_TRAIN_STUDIES,\n    \"EPOCHS\": EPOCHS,\n    \"BATCH_SIZE\": BATCH_SIZE,\n    \"NUM_WORKERS\": NUM_WORKERS,\n    \"AMP\": AMP,\n    \"USE_AMP\": USE_AMP,\n    \"GRAD_CLIP\": GRAD_CLIP,\n    \"VAL_FRAC\": VAL_FRAC,\n}, flush=True)\n\nprint(\"Готовим train/val split...\", flush=True)\nknown_cnt = train_df[LABELS].notna().sum(axis=1)\npool = train_df[known_cnt > 0].copy()\nprint(\"studies with any weak label:\", len(pool), flush=True)\n\nif len(pool) > N_TRAIN_STUDIES:\n    pool = pool.sample(N_TRAIN_STUDIES, random_state=SEED).reset_index(drop=True)\n    print(\"sampled for train:\", len(pool), flush=True)\n\npool = pool.sample(frac=1.0, random_state=SEED).reset_index(drop=True)\ncut = int(len(pool) * (1.0 - VAL_FRAC))\ntr_df = pool.iloc[:cut].copy()\nva_df = pool.iloc[cut:].copy()\nprint(\"train/val:\", len(tr_df), len(va_df), flush=True)\n\n# pos_weight по weak labels: negatives / positives, с защитой от крайностей\npos_w = []\nfor c in LABELS:\n    y = tr_df[c]\n    pos = float((y == 1).sum())\n    neg = float((y == 0).sum())\n    w = neg / max(pos, 1.0)\n    w = float(np.clip(w, 0.5, 10.0)) if np.isfinite(w) else 1.0\n    pos_w.append(w)\n\npos_weight = torch.tensor(pos_w, dtype=torch.float32, device=DEVICE)\nprint(\"pos_weight:\", {c: round(w, 3) for c, w in zip(LABELS, pos_w)}, flush=True)\n\ntr_loader = DataLoader(KneeStudyDS(tr_df, True), batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)\nva_loader = DataLoader(KneeStudyDS(va_df, True), batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\nprint(\"train batches:\", len(tr_loader), \"val batches:\", len(va_loader), flush=True)\n\nscaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n\ndef run_epoch(loader, train_mode=True, log_every=10, desc=\"train\"):\n    model.train() if train_mode else model.eval()\n    total_loss, total_n = 0.0, 0\n    all_logits, all_y = [], []\n\n    total_batches = len(loader)\n    t0 = time.time()\n    print(f\"\\n[{desc}] start | batches={total_batches} | train_mode={train_mode}\", flush=True)\n\n    with torch.set_grad_enabled(train_mode):\n        for bi, (x, pm, y) in enumerate(loader):\n            bt0 = time.time()\n            x = x.to(DEVICE, non_blocking=True)\n            pm = pm.to(DEVICE, non_blocking=True)\n            y = y.to(DEVICE, non_blocking=True)\n\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                logits = model(x, pm)\n                loss = masked_bce(logits, y, pos_weight=pos_weight)\n\n            if train_mode:\n                opt.zero_grad(set_to_none=True)\n                if USE_AMP:\n                    scaler.scale(loss).backward()\n                    if GRAD_CLIP:\n                        scaler.unscale_(opt)\n                        torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                    scaler.step(opt)\n                    scaler.update()\n                else:\n                    loss.backward()\n                    if GRAD_CLIP:\n                        torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                    opt.step()\n\n            bs = len(y)\n            total_loss += float(loss.item()) * bs\n            total_n += bs\n            all_logits.append(logits.detach().float().cpu())\n            all_y.append(y.detach().cpu())\n\n            if (bi % log_every == 0) or (bi == total_batches - 1):\n                elapsed = time.time() - t0\n                done = bi + 1\n                eta = elapsed / max(done, 1) * (total_batches - done)\n                print(\n                    f\"[{desc}] {done}/{total_batches} ({100.0 * done / total_batches:.1f}%) \"\n                    f\"loss={float(loss.item()):.4f} avg={total_loss / max(total_n, 1):.4f} \"\n                    f\"batch={time.time() - bt0:.2f}s elapsed={fmt_eta(elapsed)} eta={fmt_eta(eta)}\",\n                    flush=True\n                )\n\n    logits = torch.cat(all_logits).numpy()\n    y = np.concatenate([t.numpy() for t in all_y], axis=0)\n    prob = 1 / (1 + np.exp(-logits))\n\n    aucs = []\n    for j, c in enumerate(LABELS):\n        m = ~np.isnan(y[:, j])\n        if m.sum() > 0 and len(np.unique(y[m, j])) == 2:\n            try:\n                aucs.append((c, roc_auc_score(y[m, j], prob[m, j])))\n            except Exception:\n                pass\n\n    macro = float(np.mean([a for _, a in aucs])) if aucs else float(\"nan\")\n    print(f\"[{desc}] done | avg={total_loss / max(total_n, 1):.4f} macroAUC={macro:.4f} time={fmt_eta(time.time() - t0)}\", flush=True)\n    return total_loss / max(total_n, 1), macro, aucs\n\nbest_auc = -1.0\nfor ep in range(EPOCHS):\n    print(\"\\n\" + \"=\" * 90, flush=True)\n    print(f\"EPOCH {ep + 1}/{EPOCHS}\", flush=True)\n\n    tr_loss, _, _ = run_epoch(tr_loader, True, log_every=10, desc=f\"ep{ep}/train\")\n    va_loss, va_auc, va_aucs = run_epoch(va_loader, False, log_every=25, desc=f\"ep{ep}/val\")\n\n    print(f\"\\nepoch {ep}: train_loss={tr_loss:.4f} val_loss={va_loss:.4f} val_macroAUC={va_auc:.4f}\", flush=True)\n    print(\"val AUCs:\", [(c, round(float(a), 3)) for c, a in va_aucs], flush=True)\n\n    if np.isfinite(va_auc) and va_auc > best_auc:\n        best_auc = va_auc\n        torch.save(model.state_dict(), \"/kaggle/working/milnet_v0_4_best.pt\")\n        print(f\"saved best: /kaggle/working/milnet_v0_4_best.pt | best_auc={best_auc:.4f}\", flush=True)\n\ntorch.save(model.state_dict(), \"/kaggle/working/milnet_v0_4_last.pt\")\nprint(\"saved last: /kaggle/working/milnet_v0_4_last.pt\", flush=True)\nprint(\"Ячейка 5 завершена.\", flush=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-14T21:30:53.348824Z","iopub.status.idle":"2026-08-14T21:30:53.349128Z","shell.execute_reply.started":"2026-08-14T21:30:53.348986Z","shell.execute_reply":"2026-08-14T21:30:53.349004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ячейка 6 — inference и submission.csv\n\nmodel_path = None\nfor p in [\"/kaggle/working/milnet_v0_4_best.pt\", \"/kaggle/working/milnet_v0_4_last.pt\"]:\n    if Path(p).exists():\n        model_path = p\n        break\n\nuse_model = model_path is not None\nprint(\"use_model:\", use_model, \"| model_path:\", model_path, flush=True)\n\nif use_model:\n    try:\n        model.load_state_dict(torch.load(model_path, map_location=DEVICE))\n        model.eval()\n    except Exception as e:\n        print(\"load failed, fallback to 0.5:\", repr(e), flush=True)\n        use_model = False\n\ntest_loader = DataLoader(KneeStudyDS(test_df, False), batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\nprint(\"test batches:\", len(test_loader), flush=True)\n\npreds = []\nif use_model:\n    t0 = time.time()\n    with torch.no_grad():\n        for bi, (x, pm, _) in enumerate(test_loader):\n            x = x.to(DEVICE)\n            pm = pm.to(DEVICE)\n            with torch.cuda.amp.autocast(enabled=USE_AMP):\n                logits = model(x, pm)\n            preds.append(torch.sigmoid(logits).float().cpu().numpy())\n\n            if (bi % 20 == 0) or (bi == len(test_loader) - 1):\n                elapsed = time.time() - t0\n                done = bi + 1\n                eta = elapsed / max(done, 1) * (len(test_loader) - done)\n                print(f\"[inference] {done}/{len(test_loader)} elapsed={fmt_eta(elapsed)} eta={fmt_eta(eta)}\", flush=True)\n\n    preds = np.vstack(preds) if len(preds) else np.full((len(test_df), len(LABELS)), 0.5)\nelse:\n    preds = np.full((len(test_df), len(LABELS)), 0.5)\n\nsub = sample_sub.copy().set_index(\"StudyInstanceUID\")\npred_df = pd.DataFrame(preds, columns=LABELS)\npred_df.insert(0, \"StudyInstanceUID\", test_df[\"study\"].values)\n\nfor _, r in pred_df.iterrows():\n    if r[\"StudyInstanceUID\"] in sub.index:\n        sub.loc[r[\"StudyInstanceUID\"], LABELS] = r[LABELS].values\n\nsub = sub.reset_index()[[\"StudyInstanceUID\"] + LABELS]\nsub[LABELS] = sub[LABELS].fillna(0.5).clip(0.0, 1.0)\nsub.to_csv(\"/kaggle/working/submission.csv\", index=False)\n\nprint(\"saved: /kaggle/working/submission.csv\", sub.shape, flush=True)\ndisplay(sub.head())\n\nassert list(sub.columns) == list(sample_sub.columns)\nassert sub[LABELS].isna().sum().sum() == 0\nprint(\"submission OK\", flush=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-14T21:30:53.350856Z","iopub.status.idle":"2026-08-14T21:30:53.351137Z","shell.execute_reply.started":"2026-08-14T21:30:53.351006Z","shell.execute_reply":"2026-08-14T21:30:53.351021Z"}},"outputs":[],"execution_count":null}]}