{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":37333,"databundleVersionId":3949526,"sourceType":"competition"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ========================================\n# ⚙️  Section 0 — Configs\n# ----------------------------------------\nimport os, random, gc, json, zipfile, math, warnings, hashlib, time, sys\nfrom pathlib import Path\nimport numpy as np, pandas as pd\nimport torch, torch.nn as nn, torch.optim as optim, torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, SubsetRandomSampler\nimport albumentations as A\nimport albumentations.pytorch\nimport tifffile, cv2\nimport timm\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import roc_auc_score, average_precision_score, precision_score, recall_score, roc_curve, precision_recall_curve, accuracy_score\nimport matplotlib.pyplot as plt\nwarnings.filterwarnings(\"ignore\")\n\nclass CFG:\n    COMP          = \"mayo-clinic-strip-ai\"\n    TILE_SIZE     = 224          # 读入后再缩放\n    BATCH_SIZE    = 16\n    EPOCHS        = 4             # DEBUG 模式下先跑 1 个 epoch\n    LR            = 1e-4\n    IMG_MEAN      = (0.485,0.456,0.406)\n    IMG_STD       = (0.229,0.224,0.225)\n    DEBUG         = False          # 🚀 改成 False 训练全量\n    DEBUG_FRAC    = 1         # 只用 5 % 样本\n    NUM_WORKERS   = 1\n    SEED          = 42\n    DEVICE        = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ncfg = CFG()\n\ndef set_seed(seed=42):\n    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)\nset_seed(cfg.SEED)\n\n# ----------------------------------------\n# Section 1 — Initialization\n# ----------------------------------------\ndata_dir = Path(\"/kaggle/input/mayo-clinic-strip-ai\")\ntrain_csv = pd.read_csv(data_dir / \"train.csv\")\ntest_csv  = pd.read_csv(data_dir / \"test.csv\")\n\n# ----------------------------------------\n# Section 2 — Configuration\n# ----------------------------------------\nif cfg.DEBUG:\n    train_csv = (train_csv.groupby(\"label\", group_keys=False)    # ← 改这里\n                           .apply(lambda x: x.sample(frac=cfg.DEBUG_FRAC,\n                                                     random_state=cfg.SEED))\n                           .reset_index(drop=True))\n\n# === Section 3 — patient-level split  ===\nsgkf = StratifiedGroupKFold(n_splits=5, shuffle=True, random_state=cfg.SEED)\ntrain_idx, val_idx = next(sgkf.split(train_csv,\n                                    y=train_csv[\"label\"],\n                                    groups=train_csv[\"patient_id\"]))\n\ntrain_df = train_csv.iloc[train_idx].reset_index(drop=True)\nval_df   = train_csv.iloc[val_idx].reset_index(drop=True)\n\nlabel_map = {lbl: i for i, lbl in enumerate(sorted(train_csv[\"label\"].unique()))}\ntrain_df[\"label_id\"] = train_df[\"label\"].map(label_map)\nval_df  [\"label_id\"] = val_df[\"label\"].map(label_map)\nnum_classes = len(label_map)\n\n\n# ----------------------------------------\n# Section 4 — Dataset & Loading Mechanism\n# ----------------------------------------\ntfm_train = A.Compose([\n    # A.RandomResizedCrop(size=(cfg.TILE_SIZE, cfg.TILE_SIZE),  # ✅ 用 size\n    #                     scale=(0.8, 1.0)),\n    A.Resize(height=cfg.TILE_SIZE, width=cfg.TILE_SIZE),\n    A.HorizontalFlip(), A.VerticalFlip(), A.RandomRotate90(),\n    A.Normalize(cfg.IMG_MEAN, cfg.IMG_STD), albumentations.pytorch.ToTensorV2(),\n])\n\ntfm_val = A.Compose([\n    A.Resize(height=cfg.TILE_SIZE, width=cfg.TILE_SIZE),      # ✅ 用 height/width\n    A.Normalize(cfg.IMG_MEAN, cfg.IMG_STD), albumentations.pytorch.ToTensorV2(),\n])\n\nclass TileDataset(Dataset):\n    def __init__(self, df, transforms=None, split=\"train\"):\n        self.df = df\n        self.tfm = transforms\n        self.split = split            # \"train\" / \"test\"\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        # ⇢ 路径\n        folder = \"train\" if self.split == \"train\" else \"test\"\n        img_path = data_dir / folder / f\"{row.image_id}.tif\"\n        img = tifffile.imread(img_path)\n        if img.ndim == 2:\n            img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n        if self.tfm:\n            img = self.tfm(image=img)[\"image\"]\n\n        # ⇢ 标签：测试集没有 label_id，给个哑元 -1\n        label = row.label_id if \"label_id\" in row else -1\n        return img.float(), torch.tensor(label).long()\n\ntrain_dl = DataLoader(\n    TileDataset(train_df, tfm_train, split=\"train\"),\n    batch_size=cfg.BATCH_SIZE,\n    shuffle=True,\n    num_workers=cfg.NUM_WORKERS,\n    pin_memory=True,\n)\n\nval_dl = DataLoader(\n    TileDataset(val_df, tfm_val, split=\"train\"),\n    batch_size=cfg.BATCH_SIZE,\n    shuffle=False,\n    num_workers=cfg.NUM_WORKERS,\n    pin_memory=True,\n)\n\ntest_ds = TileDataset(test_csv, tfm_val, split=\"test\")\ntest_dl = DataLoader(\n    test_ds,\n    batch_size=cfg.BATCH_SIZE,\n    shuffle=False,\n    num_workers=cfg.NUM_WORKERS,\n    pin_memory=True,\n)\nsample_id = train_df.image_id.iloc[0]\nprint((data_dir / \"train\" / f\"{sample_id}.tif\").exists())  # True\n\n# ----------------------------------------\n# Section 5 Model Creation & Initialization\n# ---------------------------------------\n\n# Adapted from https://github.com/rishikksh20/CrossViT-pytorch/blob/master/module.py\n\nclass AttentionMILPooling(nn.Module):\n    \n    def __init__(self, in_dim, hidden_dim=128):\n        super().__init__()\n        self.attention_fc1 = nn.Linear(in_dim, hidden_dim)\n        self.attention_fc2 = nn.Linear(hidden_dim, 1)\n\n    def forward(self, x):\n        attn_weights = self.attention_fc2(torch.tanh(self.attention_fc1(x))) \n        attn_weights = F.softmax(attn_weights, dim=1) \n        weighted_sum = torch.sum(attn_weights * x, dim=1)\n        return weighted_sum, attn_weights\n\nclass ViT_MIL_Model(nn.Module):\n    \n    def __init__(self, model_name, num_classes, attn_hidden_dim=64, pretrained=True):\n        super().__init__()\n        self.vit = timm.create_model(model_name, pretrained=pretrained)\n        self.vit.reset_classifier(0)  # remove default classifier head\n        self.mil_pool = AttentionMILPooling(self.vit.num_features, attn_hidden_dim)\n        self.classifier = nn.Linear(self.vit.num_features, num_classes)\n\n    def forward(self, x):\n        # Extract patch embeddings + pooling\n        patch_embeddings = self.vit.forward_features(x)  \n        pooled, attn_weights = self.mil_pool(patch_embeddings)  \n        \n        # Classifier on the pooled features\n        logits = self.classifier(pooled)  \n        return logits, attn_weights\n\n# chose the 224 base for simplicity\nmodel = ViT_MIL_Model(\"vit_base_patch16_224\", num_classes=num_classes).to(cfg.DEVICE)\n\n# model.vit.patch_embed.img_size = (1024, 1024)\n# model.vit.pos_embed = torch.nn.Parameter(\n#     F.interpolate(\n#         model.vit.pos_embed.reshape(1, 65, 768).transpose(1, 2).reshape(1, 768, 8, 8),\n#         size=(64, 64),\n#         mode='bicubic',\n#         align_corners=False\n#     ).reshape(1, 768, -1).transpose(1, 2)\n# )\n\nmodel = model.to(cfg.DEVICE)\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.AdamW(model.parameters(), lr=cfg.LR)\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer,\n                                                 T_max=cfg.EPOCHS)\n\n# ----------------------------------------\n# Section 6 — Training\n# ----------------------------------------\ndef train_one_epoch(dl):\n    model.train()\n    total = 0\n    correct = 0\n    for x, y in dl:\n        x, y = x.to(cfg.DEVICE), y.to(cfg.DEVICE)\n        optimizer.zero_grad()\n        output = model(x)\n        out = output[0] if isinstance(output, tuple) else output\n        loss = criterion(out, y)\n        loss.backward()\n        optimizer.step()\n        preds = out.argmax(1)\n        total += y.size(0)\n        correct += (preds == y).sum().item()\n    return correct / total\n\n@torch.no_grad()\ndef validate(dl):\n    model.eval()\n    total = 0\n    correct = 0\n    y_true = []\n    y_prob = []\n    y_pred = []\n    \n    for x, y in dl:\n        x, y = x.to(cfg.DEVICE), y.to(cfg.DEVICE)\n        output = model(x)\n        out = output[0] if isinstance(output, tuple) else output\n        prob = out.softmax(1)[:, 1]\n        preds = out.argmax(1)\n        y_true.append(y.cpu().numpy())\n        y_prob.append(prob.cpu().numpy())\n        y_pred.append(preds.cpu().numpy())\n        total += y.size(0)\n        correct += (preds == y).sum().item()\n\n    y_true = np.concatenate(y_true)\n    y_prob = np.concatenate(y_prob)\n    y_pred = np.concatenate(y_pred)\n\n    # calcualte metrics\n    auc = roc_auc_score(y_true, y_prob)\n    acc = correct / total\n    average_precision = average_precision_score(y_true, y_prob)\n    fpr, tpr, _ = roc_curve(y_true, y_prob)\n    precision, recall, _ = precision_recall_curve(y_true, y_prob)\n    \n    return auc, acc, average_precision, fpr, tpr, precision, recall\n    \nbest_auc = 0.0\nfor epoch in range(cfg.EPOCHS):\n    acc  = train_one_epoch(train_dl)\n    val_auc, val_acc, average_precision, fpr, tpr, precision, recall  = validate(val_dl)\n    scheduler.step()\n    print(f\"Epoch {epoch+1}/{cfg.EPOCHS}  train-acc={acc:.4f} val-acc={val_acc:.4f} val-AUC={val_auc:.4f} avg_precision={average_precision:.4f}\")\n    if val_auc > best_auc:\n        # save the better model\n        # print(\"saving\")\n        best_auc = val_auc\n        torch.save(model.state_dict(), \"best_model.pth\")\n        \n# ----------------------------------------\n# Section 7 — Test\n# ----------------------------------------\ntest_ds = TileDataset(test_csv, tfm_val)\ntest_dl = DataLoader(test_ds, batch_size=cfg.BATCH_SIZE,\n                     shuffle=False, num_workers=cfg.NUM_WORKERS,\n                     pin_memory=True)\n\n# --- ROC Curve based on best model --- \n\nplt.figure()\nplt.plot(fpr, tpr, label=f'AUROC = {val_auc:.4f}')\nplt.plot([0, 1], [0, 1], linestyle='--', color='gray')\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\nplt.title(\"ROC Curve\")\nplt.legend()\nplt.grid()\nplt.show()\n\n# Save to file\nplt.savefig(\"ROC.png\")\n\n# --- Precision-Recall (PR) Curve based on best model ---\n\nplt.figure()\nplt.plot(recall, precision, label=f'AUPRC = {average_precision:.4f}')\nplt.xlabel(\"Recall\")\nplt.ylabel(\"Precision\")\nplt.title(\"Precision-Recall Curve\")\nplt.legend()\nplt.grid()\nplt.show()\n\n# Save to file\nplt.savefig(\"PR.png\")\n\nmodel.load_state_dict(torch.load(\"best_model.pth\")); model.eval()\n\nprobs=[]\nfor x,_ in test_dl:\n    x = x.to(cfg.DEVICE)\n    logits, _ = model(x)\n    prob = logits.softmax(1)[:, label_map[\"CE\"]]\n    probs.append(prob.detach().cpu().numpy())\n\nsub = pd.DataFrame({\n    \"image_id\": test_csv.image_id,\n    \"CE_prob\":  np.concatenate(probs)\n})\nsub.to_csv(\"submission.csv\", index=False)\nprint(\"✅ submission.csv saved!\\n\", sub.head())","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-08-03T03:17:40.009749Z","iopub.execute_input":"2025-08-03T03:17:40.010168Z","iopub.status.idle":"2025-08-03T03:38:08.3769Z","shell.execute_reply.started":"2025-08-03T03:17:40.010139Z","shell.execute_reply":"2025-08-03T03:38:08.375039Z"}},"outputs":[],"execution_count":null}]}