{"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":"import matplotlib.pyplot as plt\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\nfrom torch.utils.data import Dataset, DataLoader, SubsetRandomSampler, Subset\nfrom torchvision import datasets, transforms\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, accuracy_score, average_precision_score, precision_recall_curve, roc_curve, precision_score, recall_score\nwarnings.filterwarnings(\"ignore\")\nfrom xgboost import XGBClassifier\n\n# ------------------------------\n# 1. Setup\n# ------------------------------\n\nclass CFG:\n    COMP          = \"mayo-clinic-strip-ai\"\n    TILE_SIZE     = 1024\n    BATCH_SIZE    = 16\n    EPOCHS        = 1\n    LR            = 1e-4\n    IMG_MEAN      = (0.485,0.456,0.406)\n    IMG_STD       = (0.229,0.224,0.225)\n    DEBUG         = True          # False for full dataset\n    DEBUG_FRAC    = 0.5          # Percentage of dataset to run on DEBUG\n    NUM_WORKERS   = 0\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\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\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# ------------------------------\n# 2. Image Preprocessing\n# ------------------------------\n\ntransform = A.Compose([\n    # A.RandomResizedCrop(size=(cfg.TILE_SIZE, cfg.TILE_SIZE),\n    #                     scale=(0.8, 1.0)),\n    # A.HorizontalFlip(), A.VerticalFlip(), A.RandomRotate90(),\n    A.Resize(height=cfg.TILE_SIZE, width=cfg.TILE_SIZE),\n    A.Normalize(cfg.IMG_MEAN, cfg.IMG_STD), albumentations.pytorch.ToTensorV2(),\n])\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\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\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 = row.label_id if \"label_id\" in row else -1\n        return img.float(), torch.tensor(label).long()\n\n# ------------------------------\n# 3. Split Before Feature Extraction\n# ------------------------------\n\nsgkf = StratifiedGroupKFold(n_splits=5, shuffle=True, random_state=cfg.SEED)\n\n# ------------------------------\n# 4. Feature Extraction Function\n# ------------------------------\n\n\n# test_ds = TileDataset(test_csv, transform, split=\"test\")\n# test_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# )\n\nfor i, (train_idx, val_idx) in enumerate(sgkf.split(train_csv, y=train_csv[\"label\"], groups=train_csv[\"patient_id\"])):\n    if i == 0:  # second split (0-based indexing)\n        break\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\nmodel = timm.create_model(\"resnet18\", pretrained=True, num_classes=num_classes)\nmodel = model.to(cfg.DEVICE)\ncriterion = nn.CrossEntropyLoss()\n\ndef extract_features(loader):\n    features, labels = [], []\n    with torch.no_grad():\n        for x, y in loader:\n            x, y = x.to(cfg.DEVICE), y.to(cfg.DEVICE)\n            out = model(x);  loss = criterion(out,y)\n            features.append(out.cpu().numpy())\n            labels.extend(y.numpy())\n    return np.concatenate(features), np.array(labels)\n\n\ntrain_dl = DataLoader(\n    TileDataset(train_df, transform, 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, transform, split=\"train\"),\n    batch_size=cfg.BATCH_SIZE,\n    shuffle=False,\n    num_workers=cfg.NUM_WORKERS,\n    pin_memory=True,\n)\n\nsample_id = train_df.image_id.iloc[0]\nprint((data_dir / \"train\" / f\"{sample_id}.tif\").exists())  # 应该 True\n\nX_train, y_train = extract_features(train_dl)\nX_val, y_val = extract_features(val_dl)\n\n\n# ------------------------------\n# 5. Train XGBoost Classifier\n# ------------------------------\n\ndepth = 4\nn = 50\nrate = 0.05\nalpha = 1\nlambd = 1\n\nxgb = XGBClassifier(n_estimators=n,\n                    max_depth=depth,\n                    use_label_encoder=False,\n                    eval_metric='mlogloss',\n                    learning_rate=rate,\n                    reg_alpha=alpha,       # L1 regularization\n                    reg_lambda=lambd)      # L2 regularization)\nxgb.fit(X_train, y_train)\n\n\n# ------------------------------\n# 6. Evaluation\n# ------------------------------\n\ny_pred_train = xgb.predict(X_train)\ny_pred_val = xgb.predict(X_val)\ny_prob_train = xgb.predict_proba(X_train)[:,label_map[\"LAA\"]]\ny_prob_val = xgb.predict_proba(X_val)[:,label_map[\"LAA\"]]\nauc_val = roc_auc_score(y_val, y_prob_val)\nacc_train = accuracy_score(y_train, y_pred_train)\nacc_val = accuracy_score(y_val, y_pred_val)\nprc_val = average_precision_score(y_val, y_prob_val)\nprec = precision_score(y_val, y_pred_val)\nrec = recall_score(y_tval, y_pred_val)\n\nprint(f\"max_depth={depth} n_estimators={n} learning_rate={rate} reg_alpha={alpha} reg_lambda={lambd}\")\nprint(f\"train-acc={acc_train:.4f} val-acc={acc_val:.4f} val-AUC={auc_val:.4f} val-PRC={prc_val:.4f} precision={prec:.4f} recall={rec:.4f}\")\n\n\n# --- ROC Curve ---\nfpr, tpr, _ = roc_curve(y_val, y_prob_val)\n\nplt.figure()\nplt.plot(fpr, tpr, label=f'AUROC = {auc_val:.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 ---\nprecision, recall, _ = precision_recall_curve(y_val, y_prob_val)\n\n\nplt.figure()\nplt.plot(recall, precision, label=f'AUPRC = {prc_val:.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\ntorch.save(model.state_dict(), \"cnn.pth\")\nxgb.save_model(\"xgb.json\")\nxgb.save_model(\"model.xgb\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-04T04:41:22.932559Z","iopub.execute_input":"2025-08-04T04:41:22.933648Z","iopub.status.idle":"2025-08-04T04:52:17.728499Z","shell.execute_reply.started":"2025-08-04T04:41:22.933483Z","shell.execute_reply":"2025-08-04T04:52:17.727584Z"}},"outputs":[],"execution_count":null}]}