{"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.5\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.5: ближе к удачному v0.3, но с честной валидацией.\nIMG_SIZE = 352\nSLICES_PER_SERIES = 12\nN_TRAIN_STUDIES = 768\nEPOCHS = 3\nBATCH_SIZE = 8\nLR = 2e-4\nNUM_WORKERS = 0\nAMP = True\nGRAD_CLIP = 1.0\nVAL_FRAC = 0.25\nUSE_POS_WEIGHT = True\nPOSW_MIN, POSW_MAX = 0.7, 2.0\nEARLY_STOP_PATIENCE = 1\n\nprint(\"COMP:\", COMP, flush=True)\nprint(\"train/test:\", train.shape, test.shape, flush=True)\nprint(\"train_series/test_series:\", train_series.shape, test_series.shape, flush=True)\nprint(\"CONFIG v0.5:\", {\n    \"IMG_SIZE\": IMG_SIZE,\n    \"SLICES_PER_SERIES\": SLICES_PER_SERIES,\n    \"N_TRAIN_STUDIES\": N_TRAIN_STUDIES,\n    \"EPOCHS\": EPOCHS,\n    \"BATCH_SIZE\": BATCH_SIZE,\n    \"LR\": LR,\n    \"VAL_FRAC\": VAL_FRAC,\n    \"USE_POS_WEIGHT\": USE_POS_WEIGHT,\n    \"POSW_MIN\": POSW_MIN,\n    \"POSW_MAX\": POSW_MAX,\n    \"EARLY_STOP_PATIENCE\": EARLY_STOP_PATIENCE,\n}, flush=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ячейка 2 — weak labels v1.2 + label audit файлы\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    \"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    \"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_2.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)\nprint(stats.round(4), flush=True)\n\n# Label audit: по каждой метке сохраняем позитивные/негативные сниппеты.\ndef snippet(text, n=260):\n    t = re.sub(r\"\\s+\", \" \", str(text)).strip()\n    return t[:n]\n\naudit_rows = []\nrng = np.random.default_rng(SEED)\nfor lab in LABELS:\n    pos_idx = weak.index[weak[lab] == 1.0].to_numpy()\n    neg_idx = weak.index[weak[lab] == 0.0].to_numpy()\n\n    for value, idxs in [(1.0, pos_idx), (0.0, neg_idx)]:\n        if len(idxs) == 0:\n            continue\n        take = rng.choice(idxs, size=min(20, len(idxs)), replace=False)\n        for i in take:\n            audit_rows.append({\n                \"label\": lab,\n                \"weak_value\": value,\n                \"row\": int(i),\n                \"StudyInstanceUID\": train.loc[i, \"StudyInstanceUID\"],\n                \"snippet\": snippet(train.loc[i, \"Report\"]),\n            })\n\naudit = pd.DataFrame(audit_rows)\naudit.to_csv(\"/kaggle/working/label_audit_v0_5.csv\", index=False)\nprint(\"saved: /kaggle/working/label_audit_v0_5.csv\", audit.shape, flush=True)\n\n# Печатаем по 2 примера на метку, чтобы быстро увидеть враньё правил.\nfor lab in LABELS:\n    print(\"\\n###\", lab, flush=True)\n    sub = audit[audit[\"label\"] == lab]\n    for value in [1.0, 0.0]:\n        ss = sub[sub[\"weak_value\"] == value].head(2)\n        for _, r in ss.iterrows():\n            print(f\"[{int(value)}] {r['snippet'][:220]}\", flush=True)","metadata":{"trusted":true},"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)} | 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_5.csv\", index=False)\ntest_df.to_csv(\"/kaggle/working/test_selected_series_v0_5.csv\", index=False)\nprint(\"saved selected series csv\", flush=True)","metadata":{"trusted":true},"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\nif \"SEED\" not in globals():\n    SEED = 42\nif \"IMG_SIZE\" not in globals():\n    IMG_SIZE = 352\nif \"SLICES_PER_SERIES\" not in globals():\n    SLICES_PER_SERIES = 12\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)\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        if apply_voi_lut is not None:\n            try:\n                arr = apply_voi_lut(arr, ds)\n            except Exception:\n                pass\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        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    files = sorted(Path(path).glob(\"*.dcm\"))\n    if len(files) == 0:\n        return zero\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)\n    arr = (arr - 0.5) / 0.5\n    return torch.from_numpy(arr).float().unsqueeze(1)\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        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)\n        plane_mask = torch.tensor(pm, dtype=torch.float32)\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        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.dropout = nn.Dropout(0.10)\n        self.head = nn.Linear(feat_dim, n_labels)\n\n    def forward(self, x, plane_mask):\n        B, P, S, C, H, W = x.shape\n        slice_valid = (x.abs().sum(dim=(3, 4, 5)) > 0).float()\n        z = x.reshape(B * P * S, C, H, W)\n        f = self.backbone(z).reshape(B, P, S, -1)\n        w = slice_valid.unsqueeze(-1)\n        series_feat = (f * w).sum(dim=2) / w.sum(dim=2).clamp_min(1.0)\n        pm = plane_mask.unsqueeze(-1)\n        study_feat = (series_feat * pm).sum(dim=1) / pm.sum(dim=1).clamp_min(1.0)\n        return self.head(self.dropout(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(logits, y2, reduction=\"none\", pos_weight=pos_weight)\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, weight_decay=1e-2)\nprint(\"model ready, params:\", round(sum(p.numel() for p in model.parameters()) / 1e6, 3), \"M\", flush=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ячейка 5 — train v0.5: честная валидация, reliability-флаги, early stop\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\nif \"SEED\" not in globals(): SEED = 42\nif \"N_TRAIN_STUDIES\" not in globals(): N_TRAIN_STUDIES = 768\nif \"EPOCHS\" not in globals(): EPOCHS = 3\nif \"BATCH_SIZE\" not in globals(): BATCH_SIZE = 8\nif \"NUM_WORKERS\" not in globals(): NUM_WORKERS = 0\nif \"AMP\" not in globals(): AMP = True\nif \"GRAD_CLIP\" not in globals(): GRAD_CLIP = 1.0\nif \"VAL_FRAC\" not in globals(): VAL_FRAC = 0.25\nif \"USE_POS_WEIGHT\" not in globals(): USE_POS_WEIGHT = True\nif \"POSW_MIN\" not in globals(): POSW_MIN = 0.7\nif \"POSW_MAX\" not in globals(): POSW_MAX = 2.0\nif \"EARLY_STOP_PATIENCE\" not in globals(): EARLY_STOP_PATIENCE = 1\nif \"LABELS\" not in globals():\n    LABELS = [\"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\"]\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_5.csv\"\n    if Path(p).exists():\n        train_df = pd.read_csv(p)\n    else:\n        raise NameError(\"Нет train_df. Выполни Ячейки 1–3.\")\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\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\nif USE_POS_WEIGHT:\n    pos_w = []\n    for 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, POSW_MIN, POSW_MAX)) if np.isfinite(w) else 1.0\n        pos_w.append(w)\n    pos_weight = torch.tensor(pos_w, dtype=torch.float32, device=DEVICE)\nelse:\n    pos_weight = None\nprint(\"pos_weight:\", None if pos_weight is None else {c: round(float(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 evaluate_probs(logits_np, y_np):\n    prob = 1 / (1 + np.exp(-logits_np))\n    rows = []\n    for j, c in enumerate(LABELS):\n        y = y_np[:, j]\n        m = ~np.isnan(y)\n        known = int(m.sum())\n        pos = int((y[m] == 1).sum()) if known else 0\n        neg = int((y[m] == 0).sum()) if known else 0\n        auc = np.nan\n        reliable = False\n        if known >= 50 and pos >= 20 and neg >= 20 and len(np.unique(y[m])) == 2:\n            try:\n                auc = roc_auc_score(y[m], prob[m, j])\n                reliable = True\n            except Exception:\n                pass\n        elif known > 0 and len(np.unique(y[m])) == 2:\n            try:\n                auc = roc_auc_score(y[m], prob[m, j])\n            except Exception:\n                pass\n        rows.append({\"label\": c, \"known\": known, \"pos\": pos, \"neg\": neg, \"auc\": auc, \"reliable\": reliable})\n    rep = pd.DataFrame(rows)\n    macro_all = float(np.nanmean(rep[\"auc\"])) if rep[\"auc\"].notna().any() else np.nan\n    macro_rel = float(np.nanmean(rep.loc[rep[\"reliable\"], \"auc\"])) if rep.loc[rep[\"reliable\"], \"auc\"].notna().any() else np.nan\n    return rep, macro_all, macro_rel\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    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(f\"[{desc}] {done}/{total_batches} loss={float(loss.item()):.4f} avg={total_loss / max(total_n, 1):.4f} elapsed={fmt_eta(elapsed)} eta={fmt_eta(eta)}\", flush=True)\n\n    logits = torch.cat(all_logits).numpy()\n    y = np.concatenate([t.numpy() for t in all_y], axis=0)\n    rep, macro_all, macro_rel = evaluate_probs(logits, y)\n    print(f\"[{desc}] done | avg={total_loss / max(total_n, 1):.4f} macro_all={macro_all:.4f} macro_reliable={macro_rel:.4f} time={fmt_eta(time.time() - t0)}\", flush=True)\n    return total_loss / max(total_n, 1), rep, macro_all, macro_rel\n\nbest_metric = -1.0\nbad_epochs = 0\n\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, rep, macro_all, macro_rel = run_epoch(va_loader, False, log_every=25, desc=f\"ep{ep}/val\")\n\n    metric = macro_rel if np.isfinite(macro_rel) else macro_all\n    print(f\"\\nepoch {ep}: train_loss={tr_loss:.4f} val_loss={va_loss:.4f} macro_all={macro_all:.4f} macro_reliable={macro_rel:.4f}\", flush=True)\n    print(rep[[\"label\", \"known\", \"pos\", \"neg\", \"auc\", \"reliable\"]].to_string(index=False), flush=True)\n\n    if np.isfinite(metric) and metric > best_metric:\n        best_metric = metric\n        bad_epochs = 0\n        torch.save(model.state_dict(), \"/kaggle/working/milnet_v0_5_best.pt\")\n        print(f\"saved best: /kaggle/working/milnet_v0_5_best.pt | metric={best_metric:.4f}\", flush=True)\n    else:\n        bad_epochs += 1\n        print(f\"no improve, bad_epochs={bad_epochs}\", flush=True)\n        if ep >= 1 and bad_epochs >= EARLY_STOP_PATIENCE:\n            print(\"early stop\", flush=True)\n            break\n\ntorch.save(model.state_dict(), \"/kaggle/working/milnet_v0_5_last.pt\")\nprint(\"saved last: /kaggle/working/milnet_v0_5_last.pt\", flush=True)\nprint(\"Ячейка 5 завершена.\", flush=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ячейка 6 — inference и submission.csv\n\nif \"BATCH_SIZE\" not in globals(): BATCH_SIZE = 8\nif \"NUM_WORKERS\" not in globals(): NUM_WORKERS = 0\nif \"AMP\" not in globals(): AMP = True\nif \"LABELS\" not in globals():\n    LABELS = [\"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\"]\n\nneed = [\"KneeStudyDS\", \"model\", \"DEVICE\", \"test_df\", \"sample_sub\"]\nmissing = [x for x in need if x not in globals()]\nif missing:\n    raise NameError(f\"Не хватает {missing}. Выполни предыдущие ячейки.\")\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\nmodel_path = None\nfor p in [\"/kaggle/working/milnet_v0_5_best.pt\", \"/kaggle/working/milnet_v0_5_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            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    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},"outputs":[],"execution_count":null}]}