{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":22990,"datasetId":1136396,"databundleVersionId":2048213},{"sourceType":"datasetVersion","sourceId":2010866,"datasetId":979056,"databundleVersionId":2050389},{"sourceType":"datasetVersion","sourceId":1663886,"datasetId":982136,"databundleVersionId":1700216},{"sourceType":"modelInstanceVersion","sourceId":658703,"databundleVersionId":14607972,"modelInstanceId":498040}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport pathlib, sys, os, random, time\nimport numba, cv2, gc","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-15T11:57:29.480355Z","iopub.execute_input":"2025-12-15T11:57:29.480636Z","iopub.status.idle":"2025-12-15T11:57:29.484535Z","shell.execute_reply.started":"2025-12-15T11:57:29.480613Z","shell.execute_reply":"2025-12-15T11:57:29.483758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install \"numpy<2.0\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install albumentations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T04:46:37.210175Z","iopub.execute_input":"2025-12-15T04:46:37.210451Z","iopub.status.idle":"2025-12-15T04:46:40.461254Z","shell.execute_reply.started":"2025-12-15T04:46:37.210430Z","shell.execute_reply":"2025-12-15T04:46:40.460508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T04:48:08.161323Z","iopub.execute_input":"2025-12-15T04:48:08.162164Z","iopub.status.idle":"2025-12-15T04:48:20.849464Z","shell.execute_reply.started":"2025-12-15T04:48:08.162130Z","shell.execute_reply":"2025-12-15T04:48:20.848706Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n%matplotlib inline\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\nfrom tqdm.notebook import tqdm\n# from albumentations import *\n# import albumentations as A\n# import rasterio\n# from rasterio.windows import Window\n# from albumentations import *\n# from albumentations.pytorch import ToTensor\n# from albumentations.pytorch import ToTensorV2 as ToTensor\nimport cv2\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport segmentation_models_pytorch as smp\n# from sklearn.model_selection import KFold\nimport tifffile as tiff\nimport torch\nimport torch.backends.cudnn as cudnn\nimport torch.nn as nn\nfrom torch.nn import functional as F\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset, sampler\nfrom tqdm import tqdm_notebook as tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T04:48:26.156405Z","iopub.execute_input":"2025-12-15T04:48:26.157146Z","iopub.status.idle":"2025-12-15T04:48:36.044975Z","shell.execute_reply.started":"2025-12-15T04:48:26.157116Z","shell.execute_reply":"2025-12-15T04:48:36.044348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.data as D\nimport torchvision\nfrom torchvision import transforms as T","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:34:13.001651Z","iopub.execute_input":"2025-12-14T08:34:13.001975Z","iopub.status.idle":"2025-12-14T08:34:13.006245Z","shell.execute_reply.started":"2025-12-14T08:34:13.001950Z","shell.execute_reply":"2025-12-14T08:34:13.005647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_seed(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic=True\n    \nset_seed()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:34:15.742508Z","iopub.execute_input":"2025-12-14T08:34:15.743090Z","iopub.status.idle":"2025-12-14T08:34:15.752682Z","shell.execute_reply.started":"2025-12-14T08:34:15.743065Z","shell.execute_reply":"2025-12-14T08:34:15.751903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def seed_everything(seed=2**3):\n#     torch.manual_seed(seed)\n#     torch.cuda.manual_seed(seed)\n#     np.random.seed(seed)\n#     random.seed(seed)\n#     torch.backends.cudnn.deterministic = True\n# seed_everything(121)\n\n# fold = 0\n# nfolds = 5\nreduce = 4\nsz = 256\n\nBATCH_SIZE = 16\nDEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\nNUM_WORKERS = 4\nNUM_EPOCHS = 5\nSEED = 2020\nTH = 0.39\n\nDEVICE = 'cuda' if torch.cuda.is_available() else'cpu'\nDATA = '../input/hubmap-kidney-segmentation/test/'\nLABELS = '../input/hubmap-kidney-segmentation/train.csv'\nMASKS = '../input/hubmap-256x256/masks'\nTRAIN = '../input/hubmap-256x256/train'\ndf_sample = pd.read_csv('../input/hubmap-kidney-segmentation/sample_submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T08:10:45.707026Z","iopub.execute_input":"2025-12-15T08:10:45.707348Z","iopub.status.idle":"2025-12-15T08:10:45.714886Z","shell.execute_reply.started":"2025-12-15T08:10:45.707325Z","shell.execute_reply":"2025-12-15T08:10:45.714172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rle_encode_less_memory(img):\n    pixels=img.T.flatten()\n    pixels[0]=0\n    pixels[-1]=0\n    runs = np.where(pixels[1:] != pixels[:-1])[0]+2\n    runs[1::2]-=runs[::2]\n    return ' '.join(str(x) for x in runs)\n\n\nmean = np.array([0.65459856, 0.48386562, 0.69428385])\nstd = np.array([0.15167958, 0.23584107, 0.13146145])\n\n\ndef img2tensor(img, dtype: np.dtype = np.float32):\n    if img.ndim == 2:\n        img = np.expand_dims(img, 2)\n    img = np.transpose(img, (2, 0, 1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\n\nclass HuBMAPDataset(Dataset):\n    # 这里的 ids 参数：是你在这个 Dataset 里想要包含的那些病人/大图 ID 的列表\n    def __init__(self, ids, tfms=None):\n        self.ids = ids\n        self.tfms = tfms\n        \n        # 核心逻辑：只加载文件名中包含指定 ID 的图片\n        # 假设 TRAIN 是你的图片文件夹路径\n        self.fnames = [\n            fname for fname in os.listdir(TRAIN) \n            if fname.split('_')[0] in self.ids\n        ]\n\n    def __len__(self):\n        return len(self.fnames)\n\n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        \n        # 读取图片和 Mask (保持你原来的逻辑)\n        imgs = cv2.cvtColor(cv2.imread(os.path.join(TRAIN, fname)), cv2.COLOR_BGR2RGB)\n        masks = cv2.imread(os.path.join(MASKS, fname), cv2.IMREAD_GRAYSCALE)\n        \n        if self.tfms is not None:\n            augmented = self.tfms(image=imgs, mask=masks)\n            imgs, masks = augmented['image'], augmented['mask']\n            \n        # 这里的 img2tensor, mean, std 应该是你在外部定义的全局变量或函数，保持不变\n        return img2tensor((imgs / 255.0 - mean) / std), img2tensor(masks)\n\nfrom albumentations import *\n\ndef get_augmentation(p=1.0):\n    return Compose([\n        HorizontalFlip(),\n        VerticalFlip(),\n        RandomRotate90(),\n        ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, border_mode=cv2.BORDER_REFLECT),\n        OneOf([\n            OpticalDistortion(p=0.3),\n            GridDistortion(p=.1),\n            # IAAPiecewiseAffine(p=0.3),\n            PiecewiseAffine(p=0.3, scale=(0.03, 0.05)),\n        ], p=0.3),\n        OneOf([\n            HueSaturationValue(10, 15, 10),\n            CLAHE(clip_limit=2),\n            RandomBrightnessContrast(),\n        ], p=0.3),\n    ], p=p)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T08:10:48.788397Z","iopub.execute_input":"2025-12-15T08:10:48.788685Z","iopub.status.idle":"2025-12-15T08:10:48.800673Z","shell.execute_reply.started":"2025-12-15T08:10:48.788663Z","shell.execute_reply":"2025-12-15T08:10:48.799903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\n# 确保导入了 Dataset 需要的库，如 torch, cv2, Dataset 等\n\n# --- 第一步：获取所有唯一的 ID ---\n# 这里直接扫描文件夹文件名来获取 ID，比读 CSV 更稳，因为能确保文件存在\nall_files = os.listdir(TRAIN)\nall_ids = set([f.split('_')[0] for f in all_files]) # 提取 ID 并去重\nunique_ids = list(all_ids)\n\n# --- 第二步：按 7:2:1 划分 ID ---\nrandom.seed(42) # 固定随机种子\nrandom.shuffle(unique_ids)\n\nn_total = len(unique_ids)\nn_train = int(n_total * 0.7)\nn_valid = int(n_total * 0.2)\n\ntrain_ids = unique_ids[:n_train]\nvalid_ids = unique_ids[n_train : n_train + n_valid]\ntest_ids  = unique_ids[n_train + n_valid:]\n\nprint(f\"ID 划分情况 -> Train: {len(train_ids)}, Valid: {len(valid_ids)}, Test: {len(test_ids)}\")\n\n# --- 第三步：实例化三个 Dataset ---\n# 注意：train_ds 传入增强 (get_augmentation)，valid 和 test 不传 (None)\n\n# 训练集：传入 70% 的 ID，开启增强\ntrain_ds = HuBMAPDataset(ids=train_ids, tfms=get_augmentation())\n\n# 验证集：传入 20% 的 ID，无增强\nvalid_ds = HuBMAPDataset(ids=valid_ids, tfms=None)\n\n# 测试集：传入 10% 的 ID，无增强\ntest_ds  = HuBMAPDataset(ids=test_ids, tfms=None)\n\n# --- 第四步：创建 DataLoader ---\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=4)\nvalid_loader = DataLoader(valid_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=4)\ntest_loader  = DataLoader(test_ds,  batch_size=BATCH_SIZE, shuffle=False, num_workers=4)\n\nprint(f\"最终切片(Slice)数量 -> Train: {len(train_ds)}, Valid: {len(valid_ds)}, Test: {len(test_ds)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T08:10:54.013608Z","iopub.execute_input":"2025-12-15T08:10:54.014161Z","iopub.status.idle":"2025-12-15T08:10:54.055064Z","shell.execute_reply.started":"2025-12-15T08:10:54.014136Z","shell.execute_reply":"2025-12-15T08:10:54.054261Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(valid_loader), len(train_loader),len(test_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T08:11:09.262818Z","iopub.execute_input":"2025-12-15T08:11:09.263438Z","iopub.status.idle":"2025-12-15T08:11:09.268627Z","shell.execute_reply.started":"2025-12-15T08:11:09.263413Z","shell.execute_reply":"2025-12-15T08:11:09.267916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"imgs, masks = next(iter(test_loader))\nplt.figure(figsize=(16, 16))\nfor i, (img, mask) in enumerate(zip(imgs, masks)):\n    img = ((img.permute(1, 2, 0)*std + mean) * 255.0).numpy().astype(np.uint8)\n    plt.subplot(4, 4, i+1)\n    plt.imshow(img, vmin=0, vmax=255)\n    plt.imshow(mask.squeeze().numpy(), alpha=0.6)\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:37:43.867883Z","iopub.execute_input":"2025-12-14T08:37:43.868630Z","iopub.status.idle":"2025-12-14T08:37:45.758324Z","shell.execute_reply.started":"2025-12-14T08:37:43.868590Z","shell.execute_reply":"2025-12-14T08:37:45.757225Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def get_model():\n#     model = torchvision.models.segmentation.fcn_resnet50(True)\n#     model.classifier[4] = nn.Conv2d(512, 1, kernel_size=(1, 1), stride=(1, 1))\n#     model.aux_classifier[4] = nn.Conv2d(256, 1, kernel_size=(1, 1), stride=(1, 1))\n#     return model\n\ndef get_model(num_classes=1):\n    model = torchvision.models.segmentation.fcn_resnet50(weights=\"FCN_ResNet50_Weights.COCO_WITH_VOC_LABELS_V1\")\n    model.classifier[4] = nn.Conv2d(512, 1, kernel_size=1)\n    return model\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:39:34.673719Z","iopub.execute_input":"2025-12-14T08:39:34.674316Z","iopub.status.idle":"2025-12-14T08:39:34.678837Z","shell.execute_reply.started":"2025-12-14T08:39:34.674283Z","shell.execute_reply":"2025-12-14T08:39:34.677974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SoftDiceLoss(nn.Module):\n    def __init__(self, smooth=1., dims=(-2,-1)):\n\n        super(SoftDiceLoss, self).__init__()\n        self.smooth = smooth\n        self.dims = dims\n\n    def forward(self, x, y):\n\n        tp = (x * y).sum(self.dims)\n        fp = (x * (1 - y)).sum(self.dims)\n        fn = ((1 - x) * y).sum(self.dims)\n\n        dc = (2 * tp + self.smooth) / (2 * tp + fp + fn + self.smooth)\n        dc = dc.mean()\n\n        return 1 - dc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:39:44.793396Z","iopub.execute_input":"2025-12-14T08:39:44.793692Z","iopub.status.idle":"2025-12-14T08:39:44.799007Z","shell.execute_reply.started":"2025-12-14T08:39:44.793670Z","shell.execute_reply":"2025-12-14T08:39:44.798439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_test(model, loader, loss_fn):\n    model.eval()\n    losses = []\n    dices = []\n    ious = []\n\n    with torch.no_grad():\n        for image, target in loader:\n            image = image.to(DEVICE)\n            target = target.float().to(DEVICE)\n\n            output = model(image)['out']\n            loss = loss_fn(output, target)\n            losses.append(loss.item())\n\n            # sigmoid 后再算准确率\n            probs = output.sigmoid()\n\n            # dice score\n            y_pred = (probs > 0.5).float()\n            intersection = (y_pred * target).sum()\n            union = y_pred.sum() + target.sum()\n            dice = (2 * intersection + 1e-7) / (union + 1e-7)\n            dices.append(dice.item())\n\n            # iou score\n            union_iou = y_pred.sum() + target.sum() - intersection\n            iou = (intersection + 1e-7) / (union_iou + 1e-7)\n            ious.append(iou.item())\n\n    # 返回三个值 !!!\n    return np.mean(losses), np.mean(dices), np.mean(ious)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:39:51.227498Z","iopub.execute_input":"2025-12-14T08:39:51.227794Z","iopub.status.idle":"2025-12-14T08:39:51.234464Z","shell.execute_reply.started":"2025-12-14T08:39:51.227773Z","shell.execute_reply":"2025-12-14T08:39:51.233817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport time\n\n# --------------------------\n# Loss functions\n# --------------------------\nbce_fn = nn.BCEWithLogitsLoss()\ndice_fn = SoftDiceLoss()\n\ndef loss_fn(y_pred, y_true):\n    bce = bce_fn(y_pred, y_true)\n    dice = dice_fn(y_pred.sigmoid(), y_true)\n    return 0.8*bce + 0.2*dice\n\n\n# --------------------------\n# 实验配置\n# --------------------------\noptimizers = {\n    \"AdamW\": torch.optim.AdamW,\n    # \"SGD\": torch.optim.SGD\n}\nlearning_rates = [1e-4]\n#[1e-3, 1e-4, 5e-5]\n# learning_rates = [0.1, 0.05, 0.01]\n\nEPOCHES = 5\n\n# 记录结果\nresults = {}   # { \"AdamW_lr1e-4\": { \"train\":[], \"val\":[], \"dice\":[], ... } }\n\n\n# --------------------------\n# ⭕ 主循环：不同 Optimizer × LR\n# --------------------------\nfor opt_name, opt_class in optimizers.items():\n    for lr in learning_rates:\n\n        tag = f\"{opt_name}_lr{lr}\"\n        print(\"\\n\" + \"=\"*70)\n        print(f\"▶ Running experiment: {tag}\")\n        print(\"=\"*70)\n\n        # 记录\n        results[tag] = {\"train\": [], \"val\": [], \"dice\": [], \"iou\": []}\n\n        # 初始化模型 & 优化器\n        model = get_model().to(DEVICE)\n        optimizer = opt_class(\n            model.parameters(),\n            lr=lr,\n            weight_decay=1e-3 if opt_name==\"AdamW\" else 0\n        )\n\n        best_loss = 999\n\n        # --------------------------\n        # 🔥 训练循环\n        # --------------------------\n        for epoch in range(1, EPOCHES+1):\n            model.train()\n            train_losses = []\n            start = time.time()\n\n            for image, target in train_loader:\n                image = image.to(DEVICE)\n                target = target.float().to(DEVICE)\n\n                optimizer.zero_grad()\n\n                output = model(image)['out']\n                loss = loss_fn(output, target)\n                loss.backward()\n                optimizer.step()\n\n                train_losses.append(loss.item())\n\n            # 🔥 验证\n            vloss, vdice, viou = valid_test(model, valid_loader, loss_fn)\n\n            # 保存\n            results[tag][\"train\"].append(np.mean(train_losses))\n            results[tag][\"val\"].append(vloss)\n            results[tag][\"dice\"].append(vdice)\n            results[tag][\"iou\"].append(viou)\n\n            print(f\"Epoch {epoch:02d} | \"\n                  f\"Train {np.mean(train_losses):.4f} | \"\n                  f\"Val {vloss:.4f} | \"\n                  f\"Dice {vdice:.4f} | \"\n                  f\"IoU  {viou:.4f} | \"\n                  f\"{(time.time()-start)/60:.2f} min\")\n\n            # 🔥 保存最佳模型\n            if vloss < best_loss:\n                best_loss = vloss\n                torch.save(model.state_dict(), f\"best_{tag}.pth\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import openpyxl\n\ndef save_results_to_excel(results, filename=\"training_results.xlsx\"):\n    wb = openpyxl.Workbook()\n\n    # 删除默认Sheet\n    default_sheet = wb.active\n    wb.remove(default_sheet)\n\n    for tag, metrics in results.items():\n        ws = wb.create_sheet(title=tag[:30])  # Excel sheet名最多31字符\n\n        # 写表头\n        ws.append([\"Epoch\", \"Train Loss\", \"Val Loss\", \"Dice\", \"IoU\"])\n\n        # 假设每项都是长度 EPOCHES 的 list\n        E = len(metrics[\"train\"])\n\n        for i in range(E):\n            ws.append([\n                i+1,\n                metrics[\"train\"][i],\n                metrics[\"val\"][i],\n                metrics[\"dice\"][i],\n                metrics[\"iou\"][i]\n            ])\n\n    wb.save(filename)\n    print(f\"✔ Results saved to {filename}\")\n\nsave_results_to_excel(results, \"optimizer_SGD_lr_comparison.xlsx\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T09:56:01.312551Z","iopub.execute_input":"2025-12-14T09:56:01.313121Z","iopub.status.idle":"2025-12-14T09:56:01.988595Z","shell.execute_reply.started":"2025-12-14T09:56:01.313090Z","shell.execute_reply":"2025-12-14T09:56:01.987841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(16, 4))\n\n# --------------------------\n# Loss\n# --------------------------\nplt.subplot(1, 3, 1)\nfor tag in results:\n    plt.plot(results[tag][\"train\"], label=f\"{tag}_train\")\n    plt.plot(results[tag][\"val\"], label=f\"{tag}_val\", linestyle=\"--\")\nplt.title(\"Training & Validation Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\n\n# --------------------------\n# Dice\n# --------------------------\nplt.subplot(1, 3, 2)\nfor tag in results:\n    plt.plot(results[tag][\"dice\"], label=tag)\nplt.title(\"Validation Dice\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Dice\")\nplt.legend()\n\n# --------------------------\n# IoU\n# --------------------------\nplt.subplot(1, 3, 3)\nfor tag in results:\n    plt.plot(results[tag][\"iou\"], label=tag)\nplt.title(\"Validation IoU\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"IoU\")\nplt.legend()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T09:56:05.333234Z","iopub.execute_input":"2025-12-14T09:56:05.334025Z","iopub.status.idle":"2025-12-14T09:56:05.828143Z","shell.execute_reply.started":"2025-12-14T09:56:05.333999Z","shell.execute_reply":"2025-12-14T09:56:05.827507Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nprint(f\"\\n🚀 Start Test Set: best_{tag}.pth ...\")\n\n# 1.load model\nbest_model = get_model().to(DEVICE)\n\n# 2. load best weight\nweight_path = f\"best_{tag}.pth\"\nbest_model.load_state_dict(torch.load(weight_path))\nbest_model.eval() \n\ntest_loss, test_dice, test_iou = valid_test(best_model, test_loader, loss_fn)\n\nprint(\"\\n\" + \"=\"*40)\nprint(f\"🎉 (Final Test Results)\")\nprint(\"=\"*40)\nprint(f\"Best Config: {tag}\")\nprint(f\"Test Loss: {test_loss:.4f}\")\nprint(f\"Test Dice: {test_dice:.4f}\")\nprint(f\"Test IoU : {test_iou:.4f}\")\nprint(\"=\"*40)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T09:56:10.570349Z","iopub.execute_input":"2025-12-14T09:56:10.570779Z","iopub.status.idle":"2025-12-14T09:56:21.892333Z","shell.execute_reply.started":"2025-12-14T09:56:10.570757Z","shell.execute_reply":"2025-12-14T09:56:21.891412Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n# ==========================================\n# 🎨 可视化单张测试结果 (Auto-Compatible)\n# ==========================================\ndef plot_random_test_result(model, loader, device):\n    model.eval()\n    \n    # 1. 从 dataset 中随机抽取一张\n    # loader.dataset 支持直接索引访问\n    dataset = loader.dataset\n    total_idx = len(dataset)\n    random_idx = 518\n    # random.randint(0, total_idx - 1)\n    \n    print(f\"\\n🎨 正在绘制第 [{random_idx}/{total_idx}] 张测试样本的可视化结果...\")\n    \n    # 2. 获取数据 (image: 3xHxW, mask: HxW)\n    image, mask = dataset[random_idx]\n    \n    # 增加 Batch 维度 (1, 3, H, W) 以送入模型\n    image_input = image.unsqueeze(0).to(device)\n    \n    # 3. 模型预测\n    with torch.no_grad():\n        output = model(image_input)\n        \n        # --- 兼容性处理 (这是关键) ---\n        # 情况A: Torchvision模型 (返回字典 {'out': ...})\n        if isinstance(output, dict) and 'out' in output:\n            output = output['out']\n        # 情况B: TransUNet (可能返回列表/元组)\n        elif isinstance(output, (tuple, list)):\n            output = output[0]\n        # 情况C: SMP U-Net (直接返回 Tensor)\n        # 不需要做额外处理\n        \n        prob = output.sigmoid()\n        pred = (prob > 0.5).float()\n\n    # 4. 数据转换 (Tensor -> Numpy) 用于绘图\n    # 原图反归一化 (简单 Min-Max 归一化，保证显示正常)\n    img_np = image.permute(1, 2, 0).numpy()\n    img_show = (img_np - img_np.min()) / (img_np.max() - img_np.min())\n    \n    mask_np = mask.squeeze().numpy()\n    pred_np = pred.squeeze().cpu().numpy()\n    \n    # 5. 绘图 (1行3列)\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n    \n    # --- 图1: 原图 ---\n    axes[0].imshow(img_show)\n    axes[0].set_title(f\"Image Index: {random_idx}\", fontsize=12, color='navy')\n    axes[0].axis('off')\n\n    # --- 图2: 真值 (Ground Truth) ---\n    axes[1].imshow(img_show)\n    # 背景透明化，只显示前景\n    masked_gt = np.ma.masked_where(mask_np == 0, mask_np)\n    axes[1].imshow(masked_gt, alpha=0.7, cmap='spring') # 亮粉色表示真值\n    axes[1].set_title(\"Ground Truth (Spring)\", fontsize=12)\n    axes[1].axis('off')\n\n    # --- 图3: 预测 (Prediction) ---\n    axes[2].imshow(img_show)\n    masked_pred = np.ma.masked_where(pred_np == 0, pred_np)\n    \n    # 顺便算一下这张图的 Dice 给你参考\n    dice_score = (2 * (pred_np * mask_np).sum()) / (pred_np.sum() + mask_np.sum() + 1e-7)\n    \n    axes[2].imshow(masked_pred, alpha=0.7, cmap='jet') # 彩虹色表示预测\n    axes[2].set_title(f\"Prediction | Dice: {dice_score:.2f}\", fontsize=12)\n    axes[2].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n# ==========================================\n# ▶️ 执行绘图\n# ==========================================\n# 直接传入 best_model 和 test_loader 即可\nplot_random_test_result(best_model, test_loader, DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T10:29:20.957459Z","iopub.execute_input":"2025-12-14T10:29:20.957734Z","iopub.status.idle":"2025-12-14T10:29:21.561044Z","shell.execute_reply.started":"2025-12-14T10:29:20.957711Z","shell.execute_reply":"2025-12-14T10:29:21.560267Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## UNet Start","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef validation(model, loader, loss_fn):\n    model.eval()\n    losses = []\n    dices = []\n    ious = []\n\n    for image, target in loader:\n        image = image.to(DEVICE)\n        target = target.float().to(DEVICE)\n\n        # -------------------------------------------------\n        # ❌ 错误写法 (针对旧模型): output = model(image)['out']\n        # ✅ 正确写法 (针对 SMP ): output = model(image)\n        # -------------------------------------------------\n        output = model(image)\n        \n        # 为了兼容性，也可以写成这样防守型代码：\n        # if isinstance(output, dict):\n        #     output = output['out']\n\n        loss = loss_fn(output, target)\n        losses.append(loss.item())\n\n        # sigmoid 后再算准确率\n        probs = output.sigmoid()\n\n        # dice score\n        y_pred = (probs > 0.5).float()\n        intersection = (y_pred * target).sum()\n        union = y_pred.sum() + target.sum()\n        dice = (2 * intersection + 1e-7) / (union + 1e-7)\n        dices.append(dice.item())\n\n        # iou score\n        union_iou = y_pred.sum() + target.sum() - intersection\n        iou = (intersection + 1e-7) / (union_iou + 1e-7)\n        ious.append(iou.item())\n\n    return np.mean(losses), np.mean(dices), np.mean(ious)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T10:29:32.610088Z","iopub.execute_input":"2025-12-14T10:29:32.610868Z","iopub.status.idle":"2025-12-14T10:29:32.617073Z","shell.execute_reply.started":"2025-12-14T10:29:32.610842Z","shell.execute_reply":"2025-12-14T10:29:32.616210Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport time\nimport numpy as np\nimport segmentation_models_pytorch as smp\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\n# ============================\n# 1. 准备数据 (复用之前的 7:2:1 逻辑)\n# ============================\n# 假设 train_loader, valid_loader, test_loader 已经在上一步创建好了\n# 如果没有，请先运行上一步生成 Loader 的代码\n\n# ============================\n# 2. 定义模型 (SMP U-Net + Attention)\n# ============================\nmodel = smp.Unet(\n    encoder_name=\"mobilenet_v2\",    # 轻量级 encoder，适合快速实验\n    encoder_weights=\"imagenet\",     # 使用预训练权重加速收敛\n    in_channels=3,\n    classes=1,\n    # --- 进阶参数: scSE 注意力机制 ---\n    decoder_attention_type='scse',\n)\nmodel.to(DEVICE)\n\n# ============================\n# 3. 定义 Loss, 优化器, Scheduler\n# ============================\n# 混合 Loss: BCE + Dice\nbce_fn = nn.BCEWithLogitsLoss()\ndice_fn = smp.losses.DiceLoss(mode='binary', from_logits=True)\n\ndef loss_fn(y_pred, y_true):\n    return 0.5 * bce_fn(y_pred, y_true) + 0.5 * dice_fn(y_pred, y_true)\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-3)\n\n# 监控 val_loss，如果 5 个 epoch 不下降，学习率减半\nlr_step = ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=5, verbose=True)\n\n# ============================\n# 4. 辅助函数: 单个 Epoch 训练逻辑\n# ============================\ndef train_one_epoch(model, loader, loss_fn, optimizer):\n    model.train()\n    running_loss = 0.0\n    for image, target in loader:\n        image, target = image.to(DEVICE), target.float().to(DEVICE)\n        \n        optimizer.zero_grad()\n        output = model(image) # SMP 直接返回 Tensor\n        loss = loss_fn(output, target)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n    return running_loss / len(loader)\n\n# ============================\n# 5. 主训练循环\n# ============================\nEPOCHES = 5\nbest_dice = 0.0\nsave_path = \"best_model_unet_scse.pth\"\n\n# 表头格式化\nheader = r'''\nEpoch | Train Loss |  Val Loss  |  Val Dice  |  Val IoU   |  Time (m)\n'''\nprint(header)\n# 格式说明: {:6d}整数 | {:10.4f}浮点数 ...\nraw_line = '{:6d} | {:10.4f} | {:10.4f} | {:10.4f} | {:10.4f} | {:8.2f}'\n\nfor epoch in range(1, EPOCHES + 1):\n    start_time = time.time()\n    \n    # --- 训练 ---\n    train_loss = train_one_epoch(model, train_loader, loss_fn, optimizer)\n    \n    # --- 验证 (复用之前的 validation 函数) ---\n    val_loss, val_dice, val_iou = validation(model, valid_loader, loss_fn) # 注意这里用验证集 vloader\n    \n    # --- 更新学习率 ---\n    # ReduceLROnPlateau 需要一个指标，这里我们监控 val_loss\n    lr_step.step(val_loss)\n\n    # --- 保存最佳模型 ---\n    if val_dice > best_dice:\n        best_dice = val_dice\n        torch.save(model.state_dict(), save_path)\n        save_msg = \"--> Saved\"\n    else:\n        save_msg = \"\"\n\n    # --- 打印日志 ---\n    duration = (time.time() - start_time) / 60\n    print(raw_line.format(epoch, train_loss, val_loss, val_dice, val_iou, duration) + save_msg)\n\nprint(f\"\\n训练结束！最佳模型已保存至: {save_path}, 最佳 Dice: {best_dice:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T10:32:45.311790Z","iopub.execute_input":"2025-12-14T10:32:45.312677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport segmentation_models_pytorch as smp\nimport numpy as np\n\n# ============================\n# 1. 重新定义模型结构\n# ============================\n# 必须与训练时的配置完全一致 (encoder, attention 等)\nbest_model = smp.Unet(\n    encoder_name=\"mobilenet_v2\",\n    encoder_weights=None,           # 测试时不需要下载预训练权重，因为我们要加载自己的\n    in_channels=3,\n    classes=1,\n    decoder_attention_type='scse',  # 别忘了这个！\n)\n\n# ============================\n# 2. 加载你保存的权重\n# ============================\nweight_path = \"best_model_unet_scse.pth\"\n# map_location确保即使在只有CPU的机器上也能加载\nbest_model.load_state_dict(torch.load(weight_path, map_location=DEVICE))\n\n# 转移到设备并开启评估模式\nbest_model.to(DEVICE)\nbest_model.eval() \n\nprint(f\"✅ 已成功加载模型: {weight_path}\")\n\n# ============================\n# 3. 在测试集上跑分\n# ============================\n# 使用你之前定义好的 validation 函数 (确保是去掉了 ['out'] 的那个版本)\ntest_loss, test_dice, test_iou = validation(best_model, test_loader, loss_fn)\n\nprint(\"\\n\" + \"=\"*50)\nprint(f\"🏆 最终测试集成绩 (Final Test Results)\")\nprint(\"=\"*50)\nprint(f\"Test Loss : {test_loss:.4f}\")\nprint(f\"Test Dice : {test_dice:.4f}\")\nprint(f\"Test IoU  : {test_iou:.4f}\")\nprint(\"=\"*50)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport matplotlib.pyplot as plt\nimport numpy as np\n\ndef plot_random_one(model, dataset, device):\n    \"\"\"\n    随机从数据集中抽取一张并显示\n    \"\"\"\n    model.eval()\n    \n    # 1. 随机生成一个序号 (Index)\n    total_idx = len(dataset)\n    random_idx = 238\n    # random.randint(0, total_idx - 1)\n    \n    print(f\"🎲 正在随机抽取第 [{random_idx}/{total_idx}] 张图片进行测试...\")\n    \n    # 2. 获取单张数据 (注意：这里没有 Batch 维度)\n    image, mask = dataset[random_idx] \n    # image shape: (3, 256, 256)\n    # mask shape:  (1, 256, 256) 或 (256, 256)\n\n    # 3. 增加 Batch 维度: (3, H, W) -> (1, 3, H, W)\n    image_input = image.unsqueeze(0).to(device)\n    \n    # 4. 模型预测\n    with torch.no_grad():\n        output = model(image_input)       # SMP 直接返回 Tensor\n        prob = output.sigmoid()           # 转概率\n        pred = (prob > 0.5).float()       # 转 0/1\n\n    # 5. 数据转换用于画图 (Tensor -> Numpy)\n    # --- 原图 ---\n    img_np = image.permute(1, 2, 0).numpy()\n    # 简单的 Min-Max 反归一化，确保能看清\n    img_show = (img_np - img_np.min()) / (img_np.max() - img_np.min())\n    \n    # --- Mask ---\n    mask_np = mask.squeeze().numpy()\n    pred_np = pred.squeeze().cpu().numpy()\n    \n    # 6. 开始画图 (1行3列)\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n    \n    # --- 图1: 原图 ---\n    axes[0].imshow(img_show)\n    axes[0].set_title(f\"Image Index: {random_idx}\", fontsize=14, color='blue')\n    axes[0].axis('off')\n\n    # --- 图2: 真值 (Ground Truth) ---\n    axes[1].imshow(img_show)\n    # 隐藏背景(0)，只显示前景\n    masked_gt = np.ma.masked_where(mask_np == 0, mask_np)\n    axes[1].imshow(masked_gt, alpha=0.7, cmap='spring') # spring 是亮粉/绿色，对比度高\n    axes[1].set_title(\"Ground Truth\", fontsize=14)\n    axes[1].axis('off')\n\n    # --- 图3: 预测 (Prediction) ---\n    axes[2].imshow(img_show)\n    masked_pred = np.ma.masked_where(pred_np == 0, pred_np)\n    \n    # 计算这张图的 Dice 只是为了展示\n    dice_score = (2 * (pred_np * mask_np).sum()) / (pred_np.sum() + mask_np.sum() + 1e-7)\n    \n    axes[2].imshow(masked_pred, alpha=0.7, cmap='jet') # jet 是彩虹色，很显眼\n    axes[2].set_title(f\"Prediction | Dice: {dice_score:.2f}\", fontsize=14)\n    axes[2].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n# ========================\n# 运行命令\n# ========================\n# 注意：这里传入的是 test_ds (Dataset)，不是 test_loader\nplot_random_one(best_model, test_ds, DEVICE)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ViT start","metadata":{}},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport pathlib, sys, os, random, time\nimport numba, cv2, gc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:14:55.729827Z","iopub.execute_input":"2025-12-15T14:14:55.730639Z","iopub.status.idle":"2025-12-15T14:14:58.173644Z","shell.execute_reply.started":"2025-12-15T14:14:55.730611Z","shell.execute_reply":"2025-12-15T14:14:58.172994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torchvision\nfrom torch import nn\nfrom torch.optim import Adam\nfrom torchvision import transforms\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data.sampler import SequentialSampler, RandomSampler\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, ReduceLROnPlateau, CosineAnnealingWarmRestarts\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom scipy.ndimage.interpolation import zoom\nfrom albumentations.pytorch import ToTensorV2\nfrom PIL import Image\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:14:58.175087Z","iopub.execute_input":"2025-12-15T14:14:58.175494Z","iopub.status.idle":"2025-12-15T14:15:09.626762Z","shell.execute_reply.started":"2025-12-15T14:14:58.175467Z","shell.execute_reply":"2025-12-15T14:15:09.625955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install albumentations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:09.627599Z","iopub.execute_input":"2025-12-15T14:15:09.628054Z","iopub.status.idle":"2025-12-15T14:15:19.785667Z","shell.execute_reply.started":"2025-12-15T14:15:09.628033Z","shell.execute_reply":"2025-12-15T14:15:19.784695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!wget https://storage.googleapis.com/vit_models/imagenet21k/R50%2BViT-B_16.npz\n#! pip install self-attention-cv\n!pip install einops\n!pip install ml_collections\ntu_path = '../input/transunet/TransUNet-main'\nsys.path.append(tu_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:19.786702Z","iopub.execute_input":"2025-12-15T14:15:19.787003Z","iopub.status.idle":"2025-12-15T14:15:28.349535Z","shell.execute_reply.started":"2025-12-15T14:15:19.786963Z","shell.execute_reply":"2025-12-15T14:15:28.348784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#整个训练流程的大脑：所有训练、模型、优化器、scheduler、数据增强的设置都在这里统一管理。\nclass CFG:\n    data = 256 #512 # 数据集的尺寸或 patch 尺寸（你使用 256×256 图像）。\n    debug=False\n    apex=False\n    print_freq=100\n    num_workers=4\n    img_size=256 # appropriate input size for encoder \n    # 使用一种余弦退火 + 重启的学习率调度器：\n    scheduler='CosineAnnealingWarmRestarts' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts']\n   \n    epoch=5 # Change epochs\n    \n    # 训练使用 Lovasz Loss（专门为 IoU 优化的 loss）。\n    criterion= 'Lovasz' #'DiceBCELoss' # ['DiceLoss', 'Hausdorff', 'Lovasz']\n    base_model='Unet' # ['Unet']\n    encoder = 'vit' # ['attention','efficientnet-b5'] or other encoders from smp\n    lr=1e-4\n    min_lr=1e-6\n    batch_size=16\n    weight_decay=1e-6\n    gradient_accumulation_steps=1\n    seed=2021\n    n_fold=5\n    trn_fold= 0 #[0, 1, 2, 3, 4]\n    train=True\n    inference=False\n    optimizer = 'Adam'\n    T_0=10\n    \n    #RandAugment 参数\n    N=5 \n    M=9\n    \n    T_max=10\n    #factor=0.2\n    #patience=4\n    #eps=1e-6\n    smoothing=1\n    in_channels=3\n    \n    #Vision Transformer 的 transformer block 数量\n    vit_blocks=12 #[8, 12]\n    \n    vit_linear=1024 #1024\n    classes=1\n    MODEL_NAME = 'R50-ViT-B_16'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:28.352126Z","iopub.execute_input":"2025-12-15T14:15:28.352368Z","iopub.status.idle":"2025-12-15T14:15:28.358626Z","shell.execute_reply.started":"2025-12-15T14:15:28.352344Z","shell.execute_reply":"2025-12-15T14:15:28.357997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#函数通过设置 Python、NumPy、PyTorch、CUDA 和 CuDNN 的随机种子，尝试让每次训练结果尽可能一致（可复现）\ndef seed_torch(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    \n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    #the following line gives ~10% speedup\n    #but may lead to some stochasticity in the results \n    torch.backends.cudnn.benchmark = True # from @Iafoss comment\n\n#为整个训练程序设定种子\nseed_torch(seed=CFG.seed)\nprint(f\"Seeding set to: {CFG.seed}\")\n\n# deterministic=True → 追求可复现\n# benchmark=True → 追求速度，可能带来一些随机性\n# 结果大多数时候非常接近，但不保证 100% 完全一样。\n# 深度学习本身有噪声 → 你只需要“趋势稳定”","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:28.359358Z","iopub.execute_input":"2025-12-15T14:15:28.359604Z","iopub.status.idle":"2025-12-15T14:15:28.649580Z","shell.execute_reply.started":"2025-12-15T14:15:28.359578Z","shell.execute_reply":"2025-12-15T14:15:28.648896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#定义 5 套不同的数据增强，将会得到5组数据\n# base / weak / strong 都是对原图做不同程度的 transformation\n'''\n| mode   | 用途     | 特点                            |\n| ------ | ------ | ----------------------------- |\n| base   | 常用训练增强 | 翻转、旋转、变形、颜色增强                 |\n| rand   | 强增强    | RandAugment + 颜色 + 几何增强       |\n| strong | 最强增强   | elastic + noise + distortions |\n| weak   | 较弱增强   | 小幅度旋转 + 翻转 + resize           |\n| valid  | 验证     | 无增强，只 resize                  |\n\n\n为什么要这么多模式？\n\n因为：\n\n某些模型早期适合 weak\n\n中期适合 base\n\n提升泛化能力适合 strong / rand\n\n验证必须 valid\n\n'''\ndef get_transform(mode='base'):\n    if mode == 'base':\n        base_transform = A.Compose([\n            A.Resize(CFG.img_size, CFG.img_size, p = 1.0),\n            A.HorizontalFlip(),\n            A.VerticalFlip(),\n            A.RandomRotate90(),\n            A.ShiftScaleRotate(\n                shift_limit = 0.0625,\n                scale_limit = 0.2,\n                rotate_limit = 20,\n                p = 0.4,\n                border_mode = cv2.BORDER_REFLECT\n            ),\n            A.OneOf([\n                A.OpticalDistortion(p=0.4),\n                A.GridDistortion(p = 0.1),\n                A.PiecewiseAffine(p=0.4)\n            ], p=0.3),\n            \n            A.OneOf([\n                A.HueSaturationValue(10,15,10),\n                A.CLAHE(clip_limit = 3),\n                A.RandomBrightnessContrast(),\n            ], p = 0.4),\n            ToTensorV2()\n        ],p = 1.0)\n        \n        return base_transform\n    \n    elif mode == 'rand':\n        rand_transform = A.Compose([\n            RandAugment(CFG.N, CFG.M),\n            A.Transpose(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=0.85),\n            A.Resize(CFG.img_size, CFG.img_size, p=1.0),\n            A.Normalize(),\n            ToTensorV2()\n        ])\n        return rand_transform\n    \n    elif mode == 'strong':\n        strong_transform = A.Compose([\n            A.Transpose(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.ElasticTransform(alpha=120, sigma=120 * 0.05, alpha_affine=120 * 0.03, p=0.5),\n            A.OneOf([\n                A.RandomGamma(),\n                A.GaussNoise()           \n            ], p=0.5),\n            A.OneOf([\n                A.OpticalDistortion(p=0.4),\n                A.GridDistortion(p=0.2),\n                A.PiecewiseAffine(p=0.4),\n            ], p=0.5),\n            A.OneOf([\n                A.HueSaturationValue(10,15,10),\n                A.CLAHE(clip_limit=4),\n                A.RandomBrightnessContrast(),            \n            ], p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=0.85),\n            A.Resize(CFG.img_size, CFG.img_size, p=1.0),\n            ToTensorV2()\n        ])\n    \n        return strong_transform\n    \n    elif mode == 'weak':\n        weak_transform = A.Compose([\n            A.Resize(CFG.img_size, CFG.img_size, p=0.5),\n            A.HorizontalFlip(),\n            A.VerticalFlip(),\n            A.RandomRotate90(),\n            A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.4, \n                               border_mode=cv2.BORDER_REFLECT),\n            ToTensorV2()\n            ], p=1.0)\n        \n        return weak_transform\n    \n    elif mode == 'valid':\n        val_transform = A.Compose([\n                A.Resize(CFG.img_size, CFG.img_size, p=1.0),\n                ToTensorV2()\n            ], p=1.0)\n        return val_transform\n    \n    else:\n        print(\"Mode Unknown!\")\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:28.650494Z","iopub.execute_input":"2025-12-15T14:15:28.650806Z","iopub.status.idle":"2025-12-15T14:15:28.662971Z","shell.execute_reply.started":"2025-12-15T14:15:28.650781Z","shell.execute_reply":"2025-12-15T14:15:28.662367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 根据 train.csv 的 id，从硬盘上筛选对应的 patch 文件名，构建训练数据集。\n# train.csv 决定了「哪些原图有标注（mask）」以及「哪些原图属于训练集」。\n# 没有在 train.csv 里的图，就没有 mask，不能用来训练。\n\n'''\n① 看 train.csv 中哪些原图有标签\n② 从 train 目录中找到所有这些原图切出来的 patch\n③ 用这些 patch 和它们的 mask 来训练模型\n\n'''\nimport albumentations as A\n# class HuBMAPDataset(Dataset):\n#     def __init__(self, main_dir, df, train = True, transform = None):\n#         self.ids = df.id.values\n#         self.fnames = [fname for fname in os.listdir(train_dir) if fname.split('_')[0] in self.ids]\n        \n#         self.main_dir = main_dir\n#         self.df = df\n#         self.train = train\n#         self.transform = transform\n        \n#     def __len__(self):\n#         return len(self.fnames)\n    \n#     def __getitem__(self, idx):\n#         fname = self.fnames[idx]\n        \n#         img = cv2.cvtColor(cv2.imread(os.path.join(main_dir, 'train', fname)), cv2.COLOR_BGR2RGB)\n#         mask = cv2.imread(os.path.join(main_dir, 'masks', fname), cv2.IMREAD_GRAYSCALE)\n        \n#         if self.transform is not None:\n#             aug = self.transform(image = img, mask = mask)\n#             img, mask = aug['image'], aug['mask']\n            \n#         img = img.type('torch.FloatTensor')\n#         img = img/255\n#         mask = mask.type('torch.FloatTensor')\n        \n#         return img, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:28.664052Z","iopub.execute_input":"2025-12-15T14:15:28.664349Z","iopub.status.idle":"2025-12-15T14:15:28.688538Z","shell.execute_reply.started":"2025-12-15T14:15:28.664325Z","shell.execute_reply":"2025-12-15T14:15:28.687963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\nPath('/kaggle/working/networks').mkdir(parents=True, exist_ok=True)\n\nimport requests\n\nif Path(\"/kaggle/working/networks/vit_seg_modeling.py\").is_file():\n    print(\"Already Exist!\")\nelse:\n    print(\"Downloading `vit_seg_modeling.py` ...\")\n    request = requests.get(\"https://raw.githubusercontent.com/Beckschen/TransUNet/main/networks/vit_seg_modeling.py\")\n    with open(\"/kaggle/working/networks/vit_seg_modeling.py\",\"wb\") as f:\n        f.write(request.content)\n    print(\"Completed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:28.689342Z","iopub.execute_input":"2025-12-15T14:15:28.689504Z","iopub.status.idle":"2025-12-15T14:15:28.935928Z","shell.execute_reply.started":"2025-12-15T14:15:28.689491Z","shell.execute_reply":"2025-12-15T14:15:28.935327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport requests\n\nif Path(\"/kaggle/working/networks/vit_seg_configs.py\").is_file():\n    print(\"Already Exist!\")\nelse:\n    print(\"Downloading `vit_seg_configs.py` ...\")\n    request = requests.get(\"https://raw.githubusercontent.com/Beckschen/TransUNet/main/networks/vit_seg_configs.py\")\n    with open(\"/kaggle/working/networks/vit_seg_configs.py\",\"wb\") as f:\n        f.write(request.content)\n    print(\"Completed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:28.936688Z","iopub.execute_input":"2025-12-15T14:15:28.937022Z","iopub.status.idle":"2025-12-15T14:15:29.171838Z","shell.execute_reply.started":"2025-12-15T14:15:28.936983Z","shell.execute_reply":"2025-12-15T14:15:29.171119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import requests\n\nif Path(\"/kaggle/working/networks/vit_seg_modeling_resnet_skip.py\").is_file():\n    print(\"Already Exist!\")\nelse:\n    print(\"Downloading `vit_seg_configs.py` ...\")\n    request = requests.get(\"https://raw.githubusercontent.com/Beckschen/TransUNet/main/networks/vit_seg_modeling_resnet_skip.py\")\n    with open(\"/kaggle/working/networks/vit_seg_modeling_resnet_skip.py\",\"wb\") as f:\n        f.write(request.content)\n    print(\"Completed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:29.172703Z","iopub.execute_input":"2025-12-15T14:15:29.172978Z","iopub.status.idle":"2025-12-15T14:15:29.377071Z","shell.execute_reply.started":"2025-12-15T14:15:29.172923Z","shell.execute_reply":"2025-12-15T14:15:29.376110Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n#导入 ViT 模型与配置字典\n\nfrom networks.vit_seg_modeling import VisionTransformer as ViT_seg\nfrom networks.vit_seg_modeling import CONFIGS as CONFIGS_ViT_seg\n\n#选择模型配置\nconfig_vit = CONFIGS_ViT_seg[CFG.MODEL_NAME]\n\n#修改配置（非常重要）\n\n#设置输出类别数（0=背景，1=肾小球）\nconfig_vit.n_classes = 1\n# 设置 skip connections 数量\nconfig_vit.n_skip = 3\n#设置 ViT 预训练权重\nconfig_vit.pretrained_path = './R50+ViT-B_16.npz'\n# 设置 transformer dropout\nconfig_vit.transformer.dropout_rate = 0.2\n# 设置 MLP 隐藏层维度（ViT-B/16 的默认 mlp_dim 就是：768*4 = 3072）\nconfig_vit.transformer.mlp_dim = 3072\n#设置 attention heads 数量\nconfig_vit.transformer.num_heads = 4\n# 设置 transformer 层数 （原始 ViT-B/16 = 12 layers， 这里减少到 8，是为了：\n# 降低显存，减少训练时间，仍保持较好性能\nconfig_vit.transformer.num_layers = 8\n\nconfig_vit","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:29.377910Z","iopub.execute_input":"2025-12-15T14:15:29.378176Z","iopub.status.idle":"2025-12-15T14:15:29.468232Z","shell.execute_reply.started":"2025-12-15T14:15:29.378158Z","shell.execute_reply":"2025-12-15T14:15:29.467420Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ViTHuBMAP(nn.Module):\n    # 构造函数，设置默认配置\n    def __init__(self, configs = config_vit):\n        super(ViTHuBMAP,self).__init__()\n        \n        # 构建 TransUNet 模型（核心！）\n        self.model = ViT_seg(configs, \n                             img_size = CFG.img_size, \n                             num_classes = CFG.classes)\n        # 加载预训练权重（不需要再次预训练）\n        self.model.load_from(weights = np.load(configs.pretrained_path))\n        \n    # forward：定义前向传播（模型怎么跑）   \n    def forward(self, x):\n        img_segs = self.model(x)\n        return img_segs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:29.468931Z","iopub.execute_input":"2025-12-15T14:15:29.469164Z","iopub.status.idle":"2025-12-15T14:15:29.474589Z","shell.execute_reply.started":"2025-12-15T14:15:29.469147Z","shell.execute_reply":"2025-12-15T14:15:29.473864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#专门创建一个文件夹，用来存放所有自定义的 loss 函数文件\n\nPath('/kaggle/working/losses_pytorch').mkdir(parents=True, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:29.477763Z","iopub.execute_input":"2025-12-15T14:15:29.478410Z","iopub.status.idle":"2025-12-15T14:15:29.500773Z","shell.execute_reply.started":"2025-12-15T14:15:29.478389Z","shell.execute_reply":"2025-12-15T14:15:29.499930Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 检查是否有 Hausdorff Loss这个文件\nif Path(\"/kaggle/working/losses_pytorch/hausdorff.py\").is_file():\n    print(\"Already Exist!\")\nelse:\n    print(\"Downloading `hausdorff.py` ...\")\n    request = requests.get(\"https://raw.githubusercontent.com/JunMa11/SegLossOdyssey/master/losses_pytorch/hausdorff.py\")\n    with open(\"/kaggle/working/losses_pytorch/hausdorff.py\",\"wb\") as f:\n        f.write(request.content)\n    print(\"Completed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:29.501711Z","iopub.execute_input":"2025-12-15T14:15:29.502001Z","iopub.status.idle":"2025-12-15T14:15:29.741493Z","shell.execute_reply.started":"2025-12-15T14:15:29.501978Z","shell.execute_reply":"2025-12-15T14:15:29.740831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 检查Lovasz Loss\nif Path(\"/kaggle/working/losses_pytorch/lovasz_loss.py\").is_file():\n    print(\"Already Exist!\")\nelse:\n    print(\"Downloading `lovasz_loss.py` ...\")\n    request = requests.get(\"https://raw.githubusercontent.com/JunMa11/SegLossOdyssey/master/losses_pytorch/lovasz_loss.py\")\n    with open(\"/kaggle/working/losses_pytorch/lovasz_loss.py\",\"wb\") as f:\n        f.write(request.content)\n    print(\"Completed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:29.742268Z","iopub.execute_input":"2025-12-15T14:15:29.742596Z","iopub.status.idle":"2025-12-15T14:15:29.969796Z","shell.execute_reply.started":"2025-12-15T14:15:29.742577Z","shell.execute_reply":"2025-12-15T14:15:29.968913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 检查Focal_loss\n\nif Path(\"/kaggle/working/losses_pytorch/focal_loss.py\").is_file():\n    print(\"Already Exist!\")\nelse:\n    print(\"Downloading `focal_loss.py` ...\")\n    request = requests.get(\"https://raw.githubusercontent.com/JunMa11/SegLossOdyssey/master/losses_pytorch/focal_loss.py\")\n    with open(\"/kaggle/working/losses_pytorch/focal_loss.py\",\"wb\") as f:\n        f.write(request.content)\n    print(\"Completed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:29.970928Z","iopub.execute_input":"2025-12-15T14:15:29.971327Z","iopub.status.idle":"2025-12-15T14:15:30.210047Z","shell.execute_reply.started":"2025-12-15T14:15:29.971299Z","shell.execute_reply":"2025-12-15T14:15:30.209381Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from losses_pytorch.hausdorff import HausdorffDTLoss\nfrom losses_pytorch.lovasz_loss import LovaszSoftmax\nfrom losses_pytorch.focal_loss import FocalLoss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:30.210986Z","iopub.execute_input":"2025-12-15T14:15:30.211246Z","iopub.status.idle":"2025-12-15T14:15:30.220112Z","shell.execute_reply.started":"2025-12-15T14:15:30.211222Z","shell.execute_reply":"2025-12-15T14:15:30.219327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#定义计算Diceloss = 1-Dice Score\n\nclass DiceLoss(nn.Module):\n    def __init__(self, weight = None, size_average = True):\n        super(DiceLoss, self).__init__()\n   \n    \n    def forward(self, inputs, targets, smooth = CFG.smoothing):\n        # #首先对模型输出做 sigmoid（Sigmoid 把它变成概率（0～1））\n        inputs = F.sigmoid(inputs)\n\n        #展平为 1D 向量\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        # 计算交集 intersection\n        # 乘积表示：“预测为正类的概率 × Ground truth 正类像素” \n        # 全部相加就是预测与真实的“重叠面积”。\n        #也就是dice公式中的：∣P∩G∣\n        intersection = (inputs * targets).sum()\n\n        #计算 Dice Score = 2*∣P∩G∣/（∣P｜+｜G∣）\n        dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)\n        \n        return dice\n\n# 这个实现实际上返回的是 Dice Score（越大越好），不是 Loss（越小越好）\n# Dice Loss = 1 - Dice Score。\n# Dice越接近1预测越准。\n\n\n# '''\n# 为什么加smooth？\n# 1：避免分母为 0（最重要）\n# 2.使 gradient 更稳定（尤其是小目标）\n# 3.避免 Dice 在前景面积极小时出现极端高/低值\n# '''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:30.220970Z","iopub.execute_input":"2025-12-15T14:15:30.221273Z","iopub.status.idle":"2025-12-15T14:15:30.234083Z","shell.execute_reply.started":"2025-12-15T14:15:30.221249Z","shell.execute_reply":"2025-12-15T14:15:30.233500Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#DiceBCELoss = Dice Loss + Binary Cross Entropy (BCE)\n# '''\n# Dice 强在哪里？\n# → 区域重叠，解决前景小的问题\n\n# BCE 强在哪里？\n# → 每像素优化，收敛快，不会震荡\n\n# 解决类别不平衡\n# Dice 解决 foreground 少的问题\n\n\n# '''\n\nclass DiceBCELoss(nn.Module):\n    \n    # Formula Given Above\n    def __init__(self, weight = None, size_average = True):\n        super(DiceBCELoss, self).__init__()\n        \n    def forward(self, inputs, targets, smooth = CFG.smoothing):\n        #sigmoid 把 logits 转成概率\n        inputs = F.sigmoid(inputs)\n\n        #展平成一维\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        #计算 intersection（Dice numerator）\n        intersection = (inputs * targets).mean()\n        dice_loss = 1 - (2.*intersection + smooth)/(inputs.mean() + targets.mean() + smooth)\n        BCE = F.binary_cross_entropy(inputs, targets, reduction = 'mean')\n        \n        Dice_BCE = BCE + dice_loss\n        \n        return Dice_BCE  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:30.234779Z","iopub.execute_input":"2025-12-15T14:15:30.235032Z","iopub.status.idle":"2025-12-15T14:15:30.251893Z","shell.execute_reply.started":"2025-12-15T14:15:30.235007Z","shell.execute_reply":"2025-12-15T14:15:30.251193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Hausdorff_loss(nn.Module):\n    def __init__(self):\n        super(Hausdorff_loss, self).__init__()\n        \n    def forward(self, inputs, targets):\n        return HausdorffDTLoss()(inputs, targets)\n    \nclass FocalDLoss(nn.Module):\n    def __init__(self):\n        super(FocalDLoss, self).__init__()\n        \n    def forward(self, inputs, targets):\n        return FocalLoss()(inputs, targets)\n    \n    \nclass Lovasz_loss(nn.Module):\n    def __init__(self):\n        super(Lovasz_loss, self).__init__()\n        \n    def forward(self, inputs, targets):\n        return LovaszSoftmax()(inputs, targets)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:30.252630Z","iopub.execute_input":"2025-12-15T14:15:30.252837Z","iopub.status.idle":"2025-12-15T14:15:30.274714Z","shell.execute_reply.started":"2025-12-15T14:15:30.252821Z","shell.execute_reply":"2025-12-15T14:15:30.273855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if CFG.criterion == 'DiceBCELoss':\n    criterion = DiceBCELoss()\nelif CFG.criterion == 'DiceLoss':\n    criterion = DiceLoss()\nelif CFG.criterion == 'FocalLoss':\n    criterion = FocalDLoss()\nelif CFG.criterion == 'Hausdorff':\n    criterion = Hausdorff_loss()\nelif CFG.criterion == 'Lovasz':\n    criterion = Lovasz_loss()\n\n# criterion","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:30.275430Z","iopub.execute_input":"2025-12-15T14:15:30.275661Z","iopub.status.idle":"2025-12-15T14:15:30.291914Z","shell.execute_reply.started":"2025-12-15T14:15:30.275645Z","shell.execute_reply":"2025-12-15T14:15:30.291360Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def HuBMAPLoss(images, targets, model, device, loss_func = criterion):\n#     model.to(device)\n#     images = images.to(device)\n#     targets = targets.to(device)\n\n#     outputs = model(images)\n#     loss = loss_func(outputs, targets)\n    \n#     return loss, outputs\ndef HuBMAPLoss(images, targets, model, device, loss_func=criterion):\n    # model.to(device) 其实建议放在循环外面做一次即可，放在这里也可以但效率稍低\n    # model.to(device) \n    \n    # -----------------------------------------------------------\n    # 🛠️ 修复 1: 图片转浮点数并归一化 (解决 ByteTensor 报错)\n    # -----------------------------------------------------------\n    # 原始数据是 uint8 (0-255)，模型需要 float32 (0.0-1.0)\n    images = images.to(device).float() / 255.0\n    \n    # -----------------------------------------------------------\n    # 🛠️ 修复 2: 标签转浮点数并增加通道维度 (解决维度不匹配)\n    # -----------------------------------------------------------\n    # (Batch, 256, 256) -> (Batch, 1, 256, 256)\n    targets = targets.to(device).float().unsqueeze(1)\n\n    # 前向传播\n    outputs = model(images)\n    \n    # -----------------------------------------------------------\n    # 🛠️ 修复 3: 兼容性处理 (防止 ViT 返回 tuple)\n    # -----------------------------------------------------------\n    if isinstance(outputs, (tuple, list)):\n        outputs = outputs[0]\n        \n    # 计算 Loss\n    loss = loss_func(outputs, targets)\n    \n    return loss, outputs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:30.292746Z","iopub.execute_input":"2025-12-15T14:15:30.292927Z","iopub.status.idle":"2025-12-15T14:15:30.308687Z","shell.execute_reply.started":"2025-12-15T14:15:30.292908Z","shell.execute_reply":"2025-12-15T14:15:30.308149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nimport pandas as pd\nimport numpy as np\nfrom torch.utils.data import Dataset\n\nclass HuBMAPDataset(Dataset):\n    def __init__(self, main_dir, df, train=True, transform=None):\n        self.main_dir = main_dir\n        self.train = train\n        self.transform = transform\n        \n        # -----------------------------------------------------------\n        # 🛠️ 修复 1: 兼容 DataFrame 和 List\n        # -----------------------------------------------------------\n        # 如果传入的是 DataFrame，提取 id 列；如果传入的是 list/array，直接使用\n        if isinstance(df, pd.DataFrame):\n            # 假设 DataFrame 肯定有一列叫 'id' 或 'image_id'\n            if 'id' in df.columns:\n                self.ids = df['id'].values.astype(str)\n            else:\n                self.ids = df.iloc[:, 0].values.astype(str) # 取第一列作为ID\n        else:\n            # 如果已经是列表或 numpy array\n            self.ids = np.array(df).astype(str)\n\n        # -----------------------------------------------------------\n        # 🛠️ 修复 2: 定义 train_dir 并优化文件筛选\n        # -----------------------------------------------------------\n        # 你原来的代码里直接用了 train_dir 但没定义它\n        self.train_dir = os.path.join(self.main_dir, 'train')\n        self.mask_dir = os.path.join(self.main_dir, 'masks')\n        \n        # 获取所有图片文件名\n        # 注意：为了加快速度，这里建议把 ids 转为 set 进行查找\n        id_set = set(self.ids)\n        \n        # 遍历 train 文件夹，只保留那些 ID 在我们列表里的图片\n        # 文件名格式通常是: \"id_slice.png\" -> split('_')[0] 得到 id\n        self.fnames = [\n            fname for fname in os.listdir(self.train_dir) \n            if fname.split('_')[0] in id_set\n        ]\n\n    def __len__(self):\n        return len(self.fnames)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        \n        # 读取图片\n        img_path = os.path.join(self.train_dir, fname)\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        # 读取 Mask\n        mask_path = os.path.join(self.mask_dir, fname)\n        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n        \n        # 应用数据增强 (Albumentations)\n        if self.transform is not None:\n            aug = self.transform(image=img, mask=mask)\n            img = aug['image']\n            mask = aug['mask']\n            \n        # -----------------------------------------------------------\n        # 🛠️ 修复 3: 标准化 Torch Tensor 转换写法\n        # -----------------------------------------------------------\n        # 如果 transform 里已经有 ToTensorV2，这里就不需要手动转 tensor 了\n        # 如果 transform 只有几何变换，这里需要手动转：\n        \n        # 确保是 Tensor 格式 (C, H, W)\n        if not isinstance(img, torch.Tensor):\n            img = torch.from_numpy(img).permute(2, 0, 1).float()\n            img = img / 255.0\n        \n        if not isinstance(mask, torch.Tensor):\n            mask = torch.from_numpy(mask).float()\n            # Mask 通常不需要除以 255，因为它应该是 0/1 或者是类别索引\n            # 只有当 mask 是 0-255 的灰度图且代表概率时才除以 255\n            # 假设你的 mask 是 0和255 (二值化)，这里归一化到 0-1\n            mask = mask / 255.0 \n            # 增加 channel 维度 (H, W) -> (1, H, W)\n            mask = mask.unsqueeze(0) \n            \n        return img, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:30.309525Z","iopub.execute_input":"2025-12-15T14:15:30.309758Z","iopub.status.idle":"2025-12-15T14:15:30.330396Z","shell.execute_reply.started":"2025-12-15T14:15:30.309738Z","shell.execute_reply":"2025-12-15T14:15:30.329601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"main_dir = '../input/hubmap-256x256/'\ntrain_dir = '../input/hubmap-256x256//train/'\nmasks_dir = '../input/hubmap-256x256//masks/'\ndirectory_list = os.listdir('../input/hubmap-256x256/train')\ntrain_df = pd.read_csv('../input/hubmap-kidney-segmentation/train.csv')\n    \ndirectory_list = [fnames.split('_')[0] for fnames in directory_list]\ndir_df = pd.DataFrame(directory_list, columns=['id'])\ndir_df\n# #前面已经划分过\n# train_ds = HuBMAPDataset(main_dir, train_ids, train = True, transform = get_transform('base'))\n# valid_ds = HuBMAPDataset(main_dir, valid_ids, train = True, transform = get_transform('valid'))\n# test_ds = HuBMAPDataset(main_dir, test_ids, train = True, transform = get_transform('valid'))\n\n# len(train_ds), len(valid_ds),len(test_ds)  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:30.331143Z","iopub.execute_input":"2025-12-15T14:15:30.331382Z","iopub.status.idle":"2025-12-15T14:15:31.164901Z","shell.execute_reply.started":"2025-12-15T14:15:30.331366Z","shell.execute_reply":"2025-12-15T14:15:31.164154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\n# 确保导入了 Dataset 需要的库，如 torch, cv2, Dataset 等\n\n# --- 第一步：获取所有唯一的 ID ---\n# 这里直接扫描文件夹文件名来获取 ID，比读 CSV 更稳，因为能确保文件存在\nall_files = os.listdir(train_dir)\nall_ids = set([f.split('_')[0] for f in all_files]) # 提取 ID 并去重\nunique_ids = list(all_ids)\n\n# --- 第二步：按 7:2:1 划分 ID ---\nrandom.seed(42) # 固定随机种子\nrandom.shuffle(unique_ids)\n\nn_total = len(unique_ids)\nn_train = int(n_total * 0.7)\nn_valid = int(n_total * 0.2)\n\ntrain_ids = unique_ids[:n_train]\nvalid_ids = unique_ids[n_train : n_train + n_valid]\ntest_ids  = unique_ids[n_train + n_valid:]\n\nprint(f\"ID 划分情况 -> Train: {len(train_ids)}, Valid: {len(valid_ids)}, Test: {len(test_ids)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:31.165873Z","iopub.execute_input":"2025-12-15T14:15:31.166488Z","iopub.status.idle":"2025-12-15T14:15:31.176764Z","shell.execute_reply.started":"2025-12-15T14:15:31.166469Z","shell.execute_reply":"2025-12-15T14:15:31.176138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = HuBMAPDataset(main_dir, train_ids, train = True, transform = get_transform('base'))\nvalid_ds = HuBMAPDataset(main_dir, valid_ids, train = True, transform = get_transform('valid'))\ntest_ds = HuBMAPDataset(main_dir, test_ids, train = False)\n\ntrain_loader = DataLoader(train_ds, batch_size = CFG.batch_size, pin_memory = True, shuffle = True, num_workers=CFG.num_workers)\nvalid_loader = DataLoader(valid_ds, batch_size = CFG.batch_size, pin_memory = True, shuffle = False, num_workers= CFG.num_workers)\ntest_loader = DataLoader(test_ds, batch_size = CFG.batch_size, pin_memory = True, shuffle = False, num_workers= CFG.num_workers)\n\nprint(f\"最终切片(Slice)数量 -> Train: {len(train_ds)}, Valid: {len(valid_ds)}, Test: {len(test_ds)}\")\nprint(f\"loader -> Train: {len(train_loader)}, Valid: {len(valid_loader)}, Test: {len(test_loader)}\")   ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:31.177589Z","iopub.execute_input":"2025-12-15T14:15:31.177977Z","iopub.status.idle":"2025-12-15T14:15:31.226033Z","shell.execute_reply.started":"2025-12-15T14:15:31.177937Z","shell.execute_reply":"2025-12-15T14:15:31.225255Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(epoch, model, device, optimizer, scheduler, trainloader):\n    model.train()\n    t = time.time()\n    total_loss = 0 \n    \n    for step, (images, targets) in enumerate(trainloader):\n        loss, outputs = HuBMAPLoss(images, targets, model, device)\n        loss.backward()\n        if ((step+1)%4==0 or (step+1)==len(trainloader)):\n            optimizer.step()\n            scheduler.step()\n            optimizer.zero_grad()\n        loss = loss.detach().item()\n        total_loss += loss\n        \n        if ((step+1)%10==0 or (step+1)==len(trainloader)):\n            print(\n                    f'epoch {epoch} train step {step+1}/{len(trainloader)}, ' + \\\n                    f'loss: {total_loss/len(trainloader):.4f}, ' + \\\n                    f'time: {(time.time() - t):.4f}', end= '\\r' if (step + 1) != len(trainloader) else '\\n'\n                )\n\n            \n        \ndef valid_one_epoch(epoch, model, device, optimizer, scheduler, validloader):\n    model.eval()\n    t = time.time()\n    total_loss = 0\n    \n    for step, (images, targets) in enumerate(validloader):\n        loss, outputs = HuBMAPLoss(images, targets, model, device)\n        loss = loss.detach().item()\n        total_loss += loss\n        \n        if ((step+1)%4==0 or (step+1)==len(validloader)):\n            scheduler.step(total_loss/len(validloader))\n        \n        if ((step+1)%10==0 or (step+1)==len(validloader)):\n            print(\n                    f'**epoch {epoch} trainz step {step+1}/{len(validloader)}, ' + \\\n                    f'loss: {total_loss/len(validloader):.4f}, ' + \\\n                    f'time: {(time.time() - t):.4f}', end= '\\r' if (step + 1) != len(validloader) else '\\n'\n                )\n    \n\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:31.226845Z","iopub.execute_input":"2025-12-15T14:15:31.227193Z","iopub.status.idle":"2025-12-15T14:15:31.234849Z","shell.execute_reply.started":"2025-12-15T14:15:31.227169Z","shell.execute_reply":"2025-12-15T14:15:31.234243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel = ViTHuBMAP().to(device)\noptimizer = Adam(model.parameters(), lr = CFG.lr, weight_decay= CFG.weight_decay, amsgrad = False)\n\n# scheduler setting\nif CFG.scheduler == 'CosineAnnealingWarmRestarts':\n    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, T_mult=1, eta_min=CFG.min_lr, last_epoch=-1)\nelif CFG.scheduler == 'ReduceLROnPlateau':\n    scheduler = ReduceLROnPlateauReduceLROnPlateau(optimizer, mode='min', factor=CFG.factor, patience=CFG.patience, verbose=True, eps=CFG.eps)\nelif CFG.scheduler == 'CosineAnnealingLR':\n    scheduler = CosineAnnealingLR(optimizer, T_max=CFG.T_max, eta_min=CFG.min_lr, last_epoch=-1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:31.235710Z","iopub.execute_input":"2025-12-15T14:15:31.235888Z","iopub.status.idle":"2025-12-15T14:15:33.179854Z","shell.execute_reply.started":"2025-12-15T14:15:31.235875Z","shell.execute_reply":"2025-12-15T14:15:33.179003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(f'Training Loop...')\n# # for fold, (tr_idx, val_idx) in enumerate(gkf.split(dir_df, groups=dir_df[dir_df.columns[0]].values)):\n# #     if fold != CFG.trn_fold: # Train only one fold\n# #         continue \n\n# #     trainloader, validloader = prepare_train_valid_dataloader(dir_df, [fold])\n\n# for epoch in range(CFG.epoch):\n#     train_one_epoch(epoch, model, device, optimizer, scheduler, train_loader)\n#     with torch.no_grad():\n#         valid_one_epoch(epoch, model, device, optimizer, scheduler, valid_loader)\n        \n#         #torch.save(model.state_dict(),f'FOLD-{fold}-EPOCH-{epoch}-model.pth')\n        \n# torch.save(model.state_dict(),f'bestvit-model.pth')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:33.180737Z","iopub.execute_input":"2025-12-15T14:15:33.180971Z","iopub.status.idle":"2025-12-15T14:15:33.184603Z","shell.execute_reply.started":"2025-12-15T14:15:33.180929Z","shell.execute_reply":"2025-12-15T14:15:33.183827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def valid_epoch(epoch, model, device, optimizer, scheduler, validloader):\n    model.eval()\n    t = time.time()\n    total_loss = 0\n    total_dice = 0 # 1. 新增：记录总 Dice\n    \n    for step, (images, targets) in enumerate(validloader):\n        loss, outputs = HuBMAPLoss(images, targets, model, device)\n        loss = loss.detach().item()\n        total_loss += loss\n        \n        # 2. 新增：计算 Dice 指标 (这对分割任务很重要)\n        # ---------------------------------------------------\n        # 准备 targets (加维度以匹配输出)\n        targets_dice = targets.to(device).float().unsqueeze(1)\n        \n        # 计算预测 (Sigmoid -> 二值化)\n        prob = outputs.sigmoid()\n        pred = (prob > 0.5).float()\n        \n        # 计算交并比\n        intersection = (pred * targets_dice).sum()\n        union = pred.sum() + targets_dice.sum()\n        dice = (2. * intersection + 1e-7) / (union + 1e-7)\n        total_dice += dice.item()\n        # ---------------------------------------------------\n        \n        # 3. 修正：打印进度\n        if ((step+1)%10==0 or (step+1)==len(validloader)):\n            # 注意：这里的 loss 应该是 total_loss / (step+1) 才是当前的平均值\n            # 原代码除以 len(validloader) 会导致刚开始显示的 loss 非常小\n            print(\n                f'**epoch {epoch} valid step {step+1}/{len(validloader)}, ' + \\\n                f'loss: {total_loss/(step+1):.4f}, ' + \\\n                f'dice: {total_dice/(step+1):.4f}, ' + \\\n                f'time: {(time.time() - t):.4f}', end= '\\r' if (step + 1) != len(validloader) else '\\n'\n            )\n            \n    # 计算整个 Epoch 的平均 Loss\n    avg_loss = total_loss / len(validloader)\n    \n    # 4. 修正：Scheduler 更新移到循环外 (Epoch 结束时更新)\n    if scheduler is not None:\n        # 如果是 ReduceLROnPlateau，需要传入 loss\n        if isinstance(scheduler, torch.optim.lr_scheduler.ReduceLROnPlateau):\n            scheduler.step(avg_loss)\n        else:\n            # 其他 scheduler (如 CosineAnnealing) 通常不需要传参或只传 epoch\n            # 如果你的 scheduler 不需要在这里更新，可以注释掉\n            pass \n\n    # 5. 关键：必须返回 loss，供外部保存模型使用\n    return avg_loss\n\n# 初始化最小 loss\nbest_valid_loss = float('inf')\n\nfor epoch in range(CFG.epoch):\n    train_one_epoch(epoch, model, device, optimizer, scheduler, train_loader)\n    \n    with torch.no_grad():\n        # ✅ 现在这里有返回值了\n        val_loss = valid_epoch(epoch, model, device, optimizer, scheduler, valid_loader)\n        \n    # ✅ 可以比较并保存了\n    if val_loss < best_valid_loss:\n        print(f\"🔥 Loss Improved: {best_valid_loss:.4f} -> {val_loss:.4f}\")\n        best_valid_loss = val_loss\n        torch.save(model.state_dict(), 'bestvit-model-b16.pth')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:15:33.185550Z","iopub.execute_input":"2025-12-15T14:15:33.185887Z","iopub.status.idle":"2025-12-15T15:16:46.364759Z","shell.execute_reply.started":"2025-12-15T14:15:33.185862Z","shell.execute_reply":"2025-12-15T15:16:46.363300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import FileLink\n\n# 将 'vit-model.pth' 替换成你实际保存的文件名\nFileLink('bestvit-model-b16.pth')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T15:17:17.313917Z","iopub.execute_input":"2025-12-15T15:17:17.314653Z","iopub.status.idle":"2025-12-15T15:17:17.319163Z","shell.execute_reply.started":"2025-12-15T15:17:17.314626Z","shell.execute_reply":"2025-12-15T15:17:17.318431Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Step 1: 重新实例化模型结构 (必须先造一个空的“壳”)\n# 这里的 config_vit 必须和你训练时用的配置完全一致\nmodel = ViTHuBMAP(config_vit) \n\n# Step 2: 加载训练好的权重 (把训练好的“肉”装进去)\n# 'vit-model.pth' 必须是你之前 torch.save() 时保存的文件名\n# map_location=DEVICE 确保权重被加载到正确的设备(CPU/GPU)上\n# weight_path = '/kaggle/input/vit/pytorch/default/1/FOLD-0-model.pth'  # <--- 请确认你的文件名是这个，还是 FOLD-0-model.pth？\nweight_path = '/kaggle/working/bestvit-model-b16.pth'\nmodel.load_state_dict(torch.load(weight_path, map_location=device))\n\n# Step 3: 把完整模型移到 GPU (如果还没移的话)\nmodel.to(device)\nmodel.eval() # 极其重要！关闭 Dropout 和 BatchNormal 的训练行为\n\nprint(f\"✅ 模型已加载完毕: {weight_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T15:18:55.103775Z","iopub.execute_input":"2025-12-15T15:18:55.104048Z","iopub.status.idle":"2025-12-15T15:18:56.748934Z","shell.execute_reply.started":"2025-12-15T15:18:55.104030Z","shell.execute_reply":"2025-12-15T15:18:56.748338Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport time\nimport numpy as np\nDEVICE = device\n@torch.no_grad()\ndef inference_test(model, test_loader, device):\n    \"\"\"\n    在测试集上评估模型表现\n    返回: 平均 Loss, 平均 Dice, 平均 IoU\n    \"\"\"\n    model.eval()\n    t = time.time()\n    \n    total_loss = 0\n    total_dice = 0\n    total_iou = 0\n    \n    # 打印开始提示\n    print(f\"🚀 Start Testing on {len(test_loader.dataset)} images...\")\n\n    for step, (images, targets) in enumerate(test_loader):\n        # 1. 计算 Loss (复用你现有的 HuBMAPLoss)\n        # 注意: HuBMAPLoss 内部已经包含了 images.float()/255.0 和 targets.unsqueeze(1) 的处理\n        loss, outputs = HuBMAPLoss(images, targets, model, device)\n        \n        # 记录 Loss\n        loss_val = loss.detach().item()\n        total_loss += loss_val\n        \n        # 2. 计算 Dice 和 IoU 指标 (核心评估部分)\n        # 需要手动处理 targets 以匹配 Dice 计算的维度\n        targets_dice = targets.to(device).float().unsqueeze(1)\n        \n        prob = outputs.sigmoid()      # 转为 0-1 概率\n        pred = (prob > 0.5).float()   # 二值化预测 (0 或 1)\n        \n        intersection = (pred * targets_dice).sum()\n        union = pred.sum() + targets_dice.sum()\n        \n        # Dice Score\n        dice = (2. * intersection + 1e-7) / (union + 1e-7)\n        total_dice += dice.item()\n        \n        # IoU Score\n        iou = (intersection + 1e-7) / (union - intersection + 1e-7)\n        total_iou += iou.item()\n\n        # 3. 打印进度 (类似你的 valid 函数)\n        if ((step + 1) % 10 == 0) or ((step + 1) == len(test_loader)):\n            print(\n                f'Test Step {step+1}/{len(test_loader)} | '\n                f'Loss: {total_loss / (step+1):.4f} | '\n                f'Dice: {total_dice / (step+1):.4f} | '\n                f'IoU:  {total_iou  / (step+1):.4f} | '\n                f'Time: {(time.time() - t):.2f}s', \n                end='\\r' if (step + 1) != len(test_loader) else '\\n'\n            )\n            \n    # 计算最终平均分\n    avg_loss = total_loss / len(test_loader)\n    avg_dice = total_dice / len(test_loader)\n    avg_iou  = total_iou  / len(test_loader)\n    \n    print(\"\\n\" + \"=\"*40)\n    print(f\"🏆 Final Test Results\")\n    print(\"=\"*40)\n    print(f\"Avg Loss : {avg_loss:.4f}\")\n    print(f\"Avg Dice : {avg_dice:.4f}\")\n    print(f\"Avg IoU  : {avg_iou:.4f}\")\n    print(\"=\"*40)\n    \n    return avg_loss, avg_dice, avg_iou\n\ntest_loss, test_dice, test_iou = inference_test(model, test_loader, DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T15:19:40.748850Z","iopub.execute_input":"2025-12-15T15:19:40.749401Z","iopub.status.idle":"2025-12-15T15:19:50.199613Z","shell.execute_reply.started":"2025-12-15T15:19:40.749381Z","shell.execute_reply":"2025-12-15T15:19:50.198822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Metrices\nimport torch\nimport numpy as np\n\ndef dice_coefficient(y_true, y_pred):\n    smooth = 1\n    inputs = y_true.view(-1)\n    targets = y_pred.view(-1)\n        \n    intersection = (inputs * targets).sum()\n    dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)\n        \n    return dice\n\ndef iou(y_true, y_pred):\n    smooth = 1\n    intersection = torch.sum(y_true * y_pred)\n    union = torch.sum(y_true) + torch.sum(y_pred) - intersection\n    iou = (intersection + smooth) / (union + smooth)\n    return iou\n\ndef accuracy(y_true, y_pred):\n    correct = torch.sum((y_pred == y_true).float())\n    total = y_true.numel()\n    acc = correct / total\n    return acc\n\ndef precision(y_true, y_pred):\n    smooth = 1\n    true_positive = torch.sum(y_true * y_pred)\n    predicted_positive = torch.sum(y_pred)\n    precision = (true_positive + smooth) / (predicted_positive + smooth)\n    return precision\n\ndef recall(y_true, y_pred):\n    smooth = 1\n    true_positive = torch.sum(y_true * y_pred)\n    actual_positive = torch.sum(y_true)\n    recall = (true_positive + smooth) / (actual_positive + smooth)\n    return recall","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T15:20:37.953675Z","iopub.execute_input":"2025-12-15T15:20:37.953985Z","iopub.status.idle":"2025-12-15T15:20:37.961694Z","shell.execute_reply.started":"2025-12-15T15:20:37.953935Z","shell.execute_reply":"2025-12-15T15:20:37.960966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef calculate_metrices(validloader, model, device):\n    total_dice = 0.0\n    total_iou = 0.0\n    total_acc = 0.0\n    total_prec = 0.0\n    total_rec = 0.0\n    num_batches = 0\n\n    model.eval()\n    with torch.no_grad():\n        for img, mask in validloader:\n            img = img.to(device).float()/255.0\n            mask = mask.to(device)\n            pred_mask = model(img).cpu()\n        \n            pred_new_mask = (pred_mask.squeeze() > 0.5).float()\n            new_tensor = torch.ones_like(pred_new_mask) - pred_new_mask\n\n            batch_dice = 0.0\n            batch_iou_score = 0.0\n            batch_acc = 0.0\n            batch_prec = 0.0\n            batch_rec = 0.0\n            \n            batch_size = img.size(0)\n            \n            for i in range(batch_size):\n                           \n                batch_dice += dice_coefficient(mask[i].cpu(),  new_tensor[i].squeeze())\n                batch_iou_score += iou(mask[i].cpu(), new_tensor[i].squeeze())\n                batch_acc += accuracy(mask[i].cpu(), new_tensor[i].squeeze())\n                batch_prec += precision(mask[i].cpu(),new_tensor[i].squeeze())\n                batch_rec += recall(mask[i].cpu(), new_tensor[i].squeeze())\n\n            \n            batch_dice = batch_dice/batch_size\n            batch_iou_score = batch_iou_score/batch_size\n            batch_acc = batch_acc/batch_size\n            batch_prec = batch_prec/batch_size\n            batch_rec = batch_rec/batch_size\n            \n\n            total_dice += batch_dice\n            total_iou += batch_iou_score\n            total_acc += batch_acc\n            total_prec += batch_prec\n            total_rec += batch_rec\n            \n            \n            num_batches += 1\n    \n        mean_dice = total_dice / num_batches\n        mean_iou = total_iou / num_batches\n        mean_acc = total_acc / num_batches\n        mean_prec = total_prec / num_batches\n        mean_rec = total_rec / num_batches\n        \n        \n        data = {\n            'Mean Dice' : mean_dice,\n            'IoU' : mean_iou,\n            'Accuracy' : mean_acc,\n            'Precision' : mean_prec,\n            'Recall' : mean_rec,\n        }\n        \n        df = pd.DataFrame.from_dict(data, orient='index').T\n        \n        csv_file_path = '/kaggle/working/metrices-bestvit-b16.csv'\n        df.to_csv(csv_file_path, index = True)\n        \n# Assuming 'model' and 'device' are defined elsewhere\ncalculate_metrices(test_loader, model, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T15:21:39.849999Z","iopub.execute_input":"2025-12-15T15:21:39.850288Z","iopub.status.idle":"2025-12-15T15:21:50.157314Z","shell.execute_reply.started":"2025-12-15T15:21:39.850269Z","shell.execute_reply":"2025-12-15T15:21:50.156369Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\ndef plot_result(validloader, n_sample=2):\n    # 1. 获取一个 Batch 的数据\n    img, mask = next(iter(validloader))\n    \n    # 2. 准备模型输入 (修复之前遇到的 ByteTensor 错误)\n    # img 是 uint8 (0-255), 需要转成 float (0-1) 喂给模型\n    img_input = img.to(device).float() / 255.0\n    \n    # 3. 加载权重并推理\n    # model.load_state_dict(...) # 建议在函数外加载好模型，不要在循环里反复加载文件，太慢了\n    model.eval()\n    \n    with torch.no_grad():\n        output = model(img_input)\n        if isinstance(output, (tuple, list)):\n            output = output[0]\n            \n        # 加上 sigmoid 转成概率 (0~1)，方便可视化\n        pred_mask = output.sigmoid().cpu()\n\n    # 4. 绘图\n    plt.figure(figsize=(15, 10))\n    \n    # 确保 n_sample 不超过 batch size\n    n_sample = min(n_sample, img.shape[0])\n    \n    for i in range(n_sample):\n        # --- (1) 原图 ---\n        plt.subplot(n_sample, 3, 3*i+1)\n        # img 是 (C, H, W)，需要转置为 (H, W, C) 才能画图\n        # 如果 img 是 float，imshow 期望 0-1；如果 uint8，期望 0-255。这里 img 是原始 uint8\n        plt.imshow(np.transpose(img[i].numpy(), (1, 2, 0)))\n        plt.axis('off')\n        plt.title(\"Image\")\n\n        # --- (2) 预测 Mask ---\n        plt.subplot(n_sample, 3, 3*i+2)\n        # pred_mask[i] 是 (1, H, W) -> squeeze -> (H, W)\n        plt.imshow(pred_mask[i].squeeze().numpy(), cmap='gray')\n        plt.axis('off')\n        plt.title(\"Predicted Probability\")\n        \n        # --- (3) 真实 Mask ---\n        plt.subplot(n_sample, 3, 3*i+3)\n        # 🛠️ 关键修复: mask[i] 是 (1, H, W)，必须 squeeze 掉那个 1 变成 (H, W)\n        plt.imshow(mask[i].squeeze().numpy(), cmap='gray')\n        plt.axis('off')\n        plt.title(\"Ground Truth\")\n    \n    plt.tight_layout()\n    plt.savefig('single_batch_result.png')\n    plt.show()\n\nplot_result(test_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T15:24:00.653496Z","iopub.execute_input":"2025-12-15T15:24:00.654232Z","iopub.status.idle":"2025-12-15T15:24:02.672460Z","shell.execute_reply.started":"2025-12-15T15:24:00.654206Z","shell.execute_reply":"2025-12-15T15:24:02.671437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dice_metric = []\niou_metric = []\nacc_metric = []\nprec_metric = []\nrecall_metric = []\n\nimport matplotlib.pyplot as plt\n\ndef plot_result(validloader, model, device, batch_idx, n_sample=4):\n    model.eval()\n    with torch.no_grad():\n        for idx, (img, mask) in enumerate(validloader):\n            if idx == batch_idx:\n                break\n                \n        # img = img.to(device)\n        img = img.to(device).float() / 255.0\n        mask = mask.to(device)\n        pred_mask = model(img).cpu()\n        \n        pred_new_mask = (pred_mask.squeeze() > 0.5).float()\n        new_tensor = torch.ones_like(pred_new_mask) - pred_new_mask\n\n        N = n_sample // 2\n        plt.figure(figsize=(15, 8))\n        for i in range(n_sample):\n            plt.subplot(N, 4, 2*i+1)\n            plt.imshow(mask[i].cpu().squeeze(), cmap = 'gray')\n            plt.axis('off')\n            plt.title('Mask')\n\n            plt.subplot(N, 4, 2*i+2)\n            plt.imshow(new_tensor[i])\n            plt.axis('off')\n            plt.title('Predicted Mask')\n\n            plt.tight_layout()\n            plt.savefig('/kaggle/working/another_batch.png')\n#             print(mask[i].shape)  256 * 256\n#             print(new_tensor[i].squeeze().shape)  #1 * 256 * 256\n            \n            # Compute evaluation metrics\n            dice = dice_coefficient(mask[i].cpu(),  new_tensor[i].squeeze())\n            iou_score = iou(mask[i].cpu(), new_tensor[i].squeeze())\n            acc = accuracy(mask[i].cpu(), new_tensor[i].squeeze())\n            prec = precision(mask[i].cpu(),new_tensor[i].squeeze())\n            rec = recall(mask[i].cpu(), new_tensor[i].squeeze())\n            \n\n            dice_metric.append(dice)\n            iou_metric.append(iou_score)\n            acc_metric.append(acc)\n            prec_metric.append(prec)\n            recall_metric.append(rec)\n            \n\n\n# Assuming 'model' and 'device' are defined elsewhere\nwith torch.no_grad():\n    plot_result(test_loader, model, device, 1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T15:26:44.609484Z","iopub.execute_input":"2025-12-15T15:26:44.610159Z","iopub.status.idle":"2025-12-15T15:26:47.083851Z","shell.execute_reply.started":"2025-12-15T15:26:44.610133Z","shell.execute_reply":"2025-12-15T15:26:47.083180Z"}},"outputs":[],"execution_count":null}]}