简介:本资源是面向医学影像AI研究者与深度学习初学者的息肉肿瘤三分类分割数据集,聚焦结肠镜图像中背景、非肿瘤性息肉与肿瘤性息肉的精准像素级区分,直接支撑模型训练、验证与可视化评估。压缩包共2000个文件,主体为1998张JPEG格式医学影像及对应mask(训练集800对、测试集200对),辅以1个说明文档和1个Python可视化脚本——该脚本能自动加载任意样本,同步展示原始图像、人工标注GT及GT叠加蒙板效果,便于直观比对分割质量。资源大小263.69MB,结构清晰:images与masks目录严格配对,开箱即用。已有841人学习下载,适合开展U-Net、TransUNet等分割模型的 baseline 实验、消融分析或课程设计,配套脚本显著降低结果可视化门槛,提升算法调试效率。
1. 息肉肿瘤分割数据集不是“拿来即用”的图片包,而是医学AI落地的临床标尺
很多刚接触医学图像分割的工程师,第一反应是去GitHub搜“polyp dataset”下载zip包,解压后发现几十GB的DICOM或PNG文件夹里混着mask、json、csv,却不知道哪张图对应哪个病灶分期、哪些切片被排除在训练集外、标签里0和1之外的像素值是否代表腺瘤性息肉还是增生性息肉——这恰恰暴露了对“息肉肿瘤分割数据集”本质的误读。它不是通用图像数据集(如COCO),而是一套带临床语义约束的标注协议+严格划分的训练/测试边界+可复现的预处理链路。真正能支撑结肠镜辅助诊断系统上线的数据集,必须同时满足三重验证:放射科医生标注一致性≥0.85(Dice系数)、测试集覆盖不同内镜型号与光照条件、标签格式支持nnUNet或MONAI等主流框架直接加载。本文聚焦于如何从原始数据集出发,完成从“有标签的图片”到“可训练的torch.utils.data.Dataset实例”的可信转化,覆盖Kvasir-SEG、CVC-ClinicDB、ETIS-LaribPolypDB三大主流息肉公开数据集的共性处理逻辑,并给出验证分割掩码拓扑正确性的最小代码集。
2. 为什么必须重写数据加载器:原始息肉数据集的三大结构性缺陷
2.1 标签掩码的像素值语义不统一:从二值到多类的隐式编码
主流息肉分割数据集表面看都是“图像+mask”,但mask中像素值的临床含义差异极大。Kvasir-SEG将所有息肉统一编码为255(uint8),而ETIS-LaribPolypDB使用1表示息肉、0为背景;更复杂的是CVC-ClinicDB,其原始标注包含3类:0(背景)、1(息肉主体)、2(息肉边缘增强区)。若直接用cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)读取并除以255,会导致CVC-ClinicDB的边缘区域被错误归为前景(1→1.0,2→2.0→截断为1),而ETIS的1值在float32下变成极小浮点数(1/255≈0.0039),破坏sigmoid输出的数值稳定性。
提示:不要依赖
np.unique()盲目归一化。必须查阅各数据集官方文档确认label schema——Kvasir-SEG在 论文附录A 明确说明“mask values are 0 for background and 255 for polyp”,而ETIS-LaribPolypDB的 README.md 指出“binary mask with 0 and 1”。
2.1.1 统一标签映射的Python实现
import numpy as np import cv2 def load_polyp_mask(mask_path: str, dataset_name: str) -> np.ndarray: """ 加载并标准化息肉分割掩码,返回0-1二值数组 :param mask_path: 掩码文件路径 :param dataset_name: 数据集名称('kvasir', 'etis', 'cvc') :return: shape=(H,W), dtype=float32, 值域[0.0, 1.0] """ mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) if mask is None: raise FileNotFoundError(f"Mask not found: {mask_path}") if dataset_name == "kvasir": # Kvasir-SEG: 0=background, 255=polyp → 0/1 mask = (mask == 255).astype(np.float32) elif dataset_name == "etis": # ETIS: 0=background, 1=polyp → 直接转float mask = mask.astype(np.float32) # 0→0.0, 1→1.0 elif dataset_name == "cvc": # CVC-ClinicDB: 0=bg, 1=polyp, 2=edge → 只取主体,忽略边缘 mask = (mask == 1).astype(np.float32) else: raise ValueError(f"Unknown dataset: {dataset_name}") return mask # 验证:检查Kvasir-SEG某张mask的唯一值 test_mask = load_polyp_mask("Kvasir-SEG/masks/1.png", "kvasir") print(f"Kvasir mask unique values: {np.unique(test_mask)}") # 应输出 [0. 1.]该函数强制将所有数据集映射到标准二值空间,避免模型因输入分布偏移导致梯度爆炸。关键参数dataset_name不可省略——这是数据集元信息(metadata)的最小必要字段,后续划分训练/测试集时需按此分组。
2.2 训练集与测试集的物理隔离:避免数据泄露的硬性约束
多数公开息肉数据集(如Kvasir-SEG)仅提供“images/”和“masks/”两个文件夹,未显式区分train/test。但医学AI部署要求测试集样本在模型训练全程不可见,包括数据增强、超参搜索、早停判断。常见错误是随机划分80/20,却未保证同一患者的多帧图像不跨集合分布——结肠镜视频序列中相邻帧高度相关,若第10帧在训练集、第11帧在测试集,模型会学到帧间时序伪影而非解剖特征。
2.2.1 基于患者ID的严格划分策略
import os import pandas as pd from pathlib import Path from sklearn.model_selection import train_test_split def split_polyp_dataset( image_dir: str, mask_dir: str, dataset_name: str, test_size: float = 0.2, random_state: int = 42 ) -> dict: """ 按患者ID划分训练/测试集,返回文件路径字典 :param image_dir: 图像根目录 :param mask_dir: 掩码根目录 :param dataset_name: 数据集名(影响ID解析逻辑) :param test_size: 测试集比例 :param random_state: 随机种子 :return: {'train': [(img_path, mask_path), ...], 'test': [...]} """ image_paths = sorted(list(Path(image_dir).glob("*.png")) + list(Path(image_dir).glob("*.jpg"))) mask_paths = sorted(list(Path(mask_dir).glob("*.png")) + list(Path(mask_dir).glob("*.jpg"))) # 构建(图像路径,掩码路径,患者ID)三元组 samples = [] for img_path in image_paths: # 不同数据集患者ID提取规则 if dataset_name == "kvasir": # Kvasir-SEG文件名形如 "1.png", "2.png" → ID=文件名数字部分 patient_id = int(img_path.stem) elif dataset_name == "etis": # ETIS文件名形如 "1_1.png", "1_2.png" → ID=下划线前数字 patient_id = int(img_path.stem.split("_")[0]) elif dataset_name == "cvc": # CVC-ClinicDB文件名形如 "1.png", "2.png" → 同Kvasir patient_id = int(img_path.stem) else: raise ValueError(f"Unsupported dataset: {dataset_name}") # 匹配掩码文件(同名) mask_path = Path(mask_dir) / f"{img_path.stem}{img_path.suffix}" if not mask_path.exists(): continue samples.append((str(img_path), str(mask_path), patient_id)) # 转DataFrame便于按patient_id分组 df = pd.DataFrame(samples, columns=["image", "mask", "patient_id"]) # 获取唯一患者ID列表 patient_ids = df["patient_id"].unique() # 按患者ID划分,确保同一患者所有样本在同一集合 train_pids, test_pids = train_test_split( patient_ids, test_size=test_size, random_state=random_state, stratify=None # 息肉数据集通常无分层需求 ) # 分配样本 train_df = df[df["patient_id"].isin(train_pids)] test_df = df[df["patient_id"].isin(test_pids)] return { "train": list(zip(train_df["image"], train_df["mask"])), "test": list(zip(test_df["image"], test_df["mask"])) } # 示例:为Kvasir-SEG生成划分 split_dict = split_polyp_dataset( image_dir="Kvasir-SEG/images", mask_dir="Kvasir-SEG/masks", dataset_name="kvasir", test_size=0.2 ) print(f"Train samples: {len(split_dict['train'])}, Test samples: {len(split_dict['test'])}")此函数输出的split_dict是后续构建PyTorch Dataset的直接输入。注意stratify=None——医学数据集常存在类别不平衡(如小息肉占比高),但按患者划分时无法保证每组患者息肉大小分布一致,强行分层可能破坏临床真实性,故采用简单随机划分患者ID。
2.3 图像-掩码配对校验:防止文件名错位导致的训练崩溃
当数据集由多人协作标注或经历多次格式转换时,极易出现“图像文件存在但掩码缺失”或“掩码文件名与图像不匹配”问题。若在DataLoader中才报错(如FileNotFoundError),将导致训练中断且难以定位。必须在数据集初始化阶段完成全量校验。
2.3.1 配对完整性检查表
| 检查项 | 方法 | 失败示例 | 修复动作 |
|---|---|---|---|
| 文件名完全匹配 | os.path.basename(img_path) == os.path.basename(mask_path) | img/001.jpgvsmask/001.png | 统一后缀或重命名 |
| 尺寸严格一致 | cv2.imread(img).shape[:2] == cv2.imread(mask).shape[:2] | 图像(1024,768) vs 掩码(1024,767) | 裁剪/填充至相同尺寸 |
| 掩码值域合规 | np.unique(mask) in ([0,1], [0,255]) | 掩码含[0,1,2,3] | 按2.1节规则映射 |
def validate_polyp_pairs(sample_list: list) -> list: """ 批量校验图像-掩码配对质量,返回错误报告 :param sample_list: [(img_path, mask_path), ...] :return: 错误条目列表,每项为(dict)含'error_type','img_path','mask_path','details' """ errors = [] for img_path, mask_path in sample_list: try: # 1. 文件存在性 if not os.path.exists(img_path): errors.append({ "error_type": "image_missing", "img_path": img_path, "mask_path": mask_path, "details": "Image file not found" }) continue if not os.path.exists(mask_path): errors.append({ "error_type": "mask_missing", "img_path": img_path, "mask_path": mask_path, "details": "Mask file not found" }) continue # 2. 尺寸匹配 img = cv2.imread(img_path) mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) if img.shape[:2] != mask.shape[:2]: errors.append({ "error_type": "size_mismatch", "img_path": img_path, "mask_path": mask_path, "details": f"Image {img.shape[:2]} != Mask {mask.shape[:2]}" }) continue # 3. 掩码值域检查(粗筛) unique_vals = np.unique(mask) if not set(unique_vals).issubset({0, 1, 255}): errors.append({ "error_type": "invalid_mask_values", "img_path": img_path, "mask_path": mask_path, "details": f"Mask contains invalid values: {unique_vals}" }) except Exception as e: errors.append({ "error_type": "io_error", "img_path": img_path, "mask_path": mask_path, "details": str(e) }) return errors # 运行校验 errors = validate_polyp_pairs(split_dict["train"][:100]) # 先校验前100个 if errors: print(f"Found {len(errors)} errors:") for err in errors[:5]: # 打印前5个 print(f"- {err['error_type']}: {err['details']} ({err['img_path']})") else: print("All pairs validated successfully.")该校验函数在数据集加载前执行,将错误集中暴露,避免训练中因单个坏样本导致整个epoch失败。生产环境应将errors写入日志文件供标注团队追溯。
3. 构建可复现的PyTorch Dataset:从路径到张量的确定性流水线
3.1 定义PolypDataset类:封装预处理与增强逻辑
import torch from torch.utils.data import Dataset from torchvision import transforms import albumentations as A from albumentations.pytorch import ToTensorV2 class PolypDataset(Dataset): def __init__( self, sample_list: list, dataset_name: str, image_size: tuple = (384, 384), transform: A.Compose = None, is_train: bool = True ): """ 息肉分割数据集 :param sample_list: [(img_path, mask_path), ...] 来自split_polyp_dataset :param dataset_name: 'kvasir', 'etis', 'cvc' :param image_size: 输出图像尺寸 (H, W) :param transform: albumentations增强流水线 :param is_train: 是否训练模式(影响transform选择) """ self.sample_list = sample_list self.dataset_name = dataset_name self.image_size = image_size self.transform = transform self.is_train = is_train # 定义基础变换(始终应用) self.base_transform = A.Compose([ A.Resize(height=image_size[0], width=image_size[1], interpolation=cv2.INTER_LINEAR), A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), # ImageNet统计量 ToTensorV2() ]) def __len__(self): return len(self.sample_list) def __getitem__(self, idx: int) -> dict: img_path, mask_path = self.sample_list[idx] # 加载图像(BGR→RGB) image = cv2.imread(img_path) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # 加载并标准化掩码 mask = load_polyp_mask(mask_path, self.dataset_name) # 应用增强(仅训练时) if self.is_train and self.transform is not None: augmented = self.transform(image=image, mask=mask) image, mask = augmented["image"], augmented["mask"] # 应用基础变换(Resize+Normalize+ToTensor) base_augmented = self.base_transform(image=image, mask=mask) image, mask = base_augmented["image"], base_augmented["mask"] return { "image": image, # torch.Tensor, shape=(3,H,W), range=[0,1] "mask": mask, # torch.Tensor, shape=(1,H,W), range=[0,1] "path": img_path # 用于debug或可视化 } # 定义训练增强(仅用于训练集) train_transform = A.Compose([ A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.5), A.GaussNoise(var_limit=(10.0, 50.0), p=0.3), A.MotionBlur(blur_limit=3, p=0.2), ]) # 创建数据集实例 train_dataset = PolypDataset( sample_list=split_dict["train"], dataset_name="kvasir", image_size=(384, 384), transform=train_transform, is_train=True ) test_dataset = PolypDataset( sample_list=split_dict["test"], dataset_name="kvasir", image_size=(384, 384), transform=None, # 测试时不增强 is_train=False ) # 验证数据集输出 sample = train_dataset[0] print(f"Image shape: {sample['image'].shape}, dtype: {sample['image'].dtype}") print(f"Mask shape: {sample['mask'].shape}, unique values: {torch.unique(sample['mask'])}")此PolypDataset类的关键设计点:
load_polyp_mask内联调用:确保标签映射逻辑与数据集实例绑定,避免外部函数污染;base_transform强制应用:无论是否训练,都执行Resize+Normalize+ToTensor,保证输入张量维度和数值范围绝对一致;is_train开关控制增强:生产环境中可扩展为eval_transform支持TTA(Test Time Augmentation)。
3.2 DataLoader配置:医学图像的批处理特殊性
医学图像分割对batch size敏感——过小(如bs=2)导致BN层统计量不准,过大(bs=32)易OOM。需根据GPU显存动态调整,并启用pin_memory加速CPU→GPU传输。
from torch.utils.data import DataLoader def create_polyp_dataloaders( train_dataset: PolypDataset, test_dataset: PolypDataset, batch_size: int = 8, num_workers: int = 4, pin_memory: bool = True ) -> dict: """ 创建训练/测试DataLoader :param batch_size: 批大小(建议Kvasir-SEG用8,ETIS用12) :param num_workers: 数据加载进程数(设为CPU核心数-1) :param pin_memory: 启用内存锁定,加速GPU传输 :return: {'train': DataLoader, 'test': DataLoader} """ train_loader = DataLoader( train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, pin_memory=pin_memory, drop_last=True, # 防止最后batch size不足 persistent_workers=True # PyTorch 1.7+,避免worker重启开销 ) test_loader = DataLoader( test_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers, pin_memory=pin_memory, drop_last=False ) return {"train": train_loader, "test": test_loader} # 创建DataLoader dataloaders = create_polyp_dataloaders( train_dataset=train_dataset, test_dataset=test_dataset, batch_size=8, num_workers=4 ) # 验证loader输出 for batch in dataloaders["train"]: print(f"Batch image shape: {batch['image'].shape}") # torch.Size([8, 3, 384, 384]) print(f"Batch mask shape: {batch['mask'].shape}") # torch.Size([8, 1, 384, 384]) breakdrop_last=True对训练至关重要——若最后一个batch只有5个样本(bs=8),BN层计算的均值/方差将严重失真。而测试集保留drop_last=False,确保所有测试样本参与评估。
4. 验证分割掩码的临床合理性:超越像素级准确率的三层检查
4.1 形态学验证:检测掩码中的孔洞与断裂
临床医生关注息肉的连续性和边界完整性。一个合格的分割掩码不应出现:
- 内部孔洞(hole):息肉组织被错误预测为空洞,可能漏诊癌变区域;
- 边界断裂(disconnection):息肉边缘被切成多个孤立小块,影响尺寸测量。
import numpy as np import cv2 def validate_mask_morphology(mask: np.ndarray, min_hole_area: int = 50) -> dict: """ 形态学验证息肉掩码 :param mask: 二值掩码 (H,W), dtype=float32, [0,1] :param min_hole_area: 最小孔洞面积阈值(像素) :return: 验证结果字典 """ mask_uint8 = (mask * 255).astype(np.uint8) # 1. 检测孔洞:对掩码取反,找连通域 inverted = cv2.bitwise_not(mask_uint8) num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(inverted, connectivity=8) holes = [] for i in range(1, num_labels): # 跳过背景标签0 area = stats[i, cv2.CC_STAT_AREA] if area >= min_hole_area: holes.append({ "area": int(area), "bbox": stats[i, cv2.CC_STAT_LEFT:cv2.CC_STAT_LEFT+4].tolist() # [x,y,w,h] }) # 2. 检测断裂:计算主连通域数量 num_labels_fg, _, _, _ = cv2.connectedComponents(mask_uint8, connectivity=8) # 减去背景(标签0) foreground_components = num_labels_fg - 1 return { "hole_count": len(holes), "holes": holes, "component_count": foreground_components, "is_connected": foreground_components <= 3 # 临床接受≤3个主要组件 } # 对测试集前10个掩码验证 for i in range(10): sample = test_dataset[i] mask_np = sample["mask"].numpy().squeeze() # (H,W) result = validate_mask_morphology(mask_np) print(f"Sample {i}: holes={result['hole_count']}, components={result['component_count']}, connected={result['is_connected']}")该函数返回的is_connected布尔值可作为数据质量过滤器——若大量测试样本is_connected=False,说明模型过度分割,需调整损失函数(如加入轮廓损失)或后处理(形态学闭运算)。
4.2 解剖结构一致性检查:息肉位置是否符合结肠解剖学
息肉在结肠镜图像中具有典型位置偏好:直肠乙状结肠交界处、升结肠近肝曲。若模型在图像顶部(内镜视野上方)频繁预测息肉,大概率是标注噪声或模型学习了设备伪影(如镜头反光)。需建立位置先验知识库。
def check_anatomical_plausibility(mask: np.ndarray, image_height: int) -> str: """ 基于息肉在图像中的垂直位置判断解剖合理性 :param mask: 二值掩码 (H,W) :param image_height: 图像高度(像素) :return: 'plausible', 'suspicious_top', 'suspicious_bottom' """ # 计算息肉质心纵坐标 y_coords, x_coords = np.where(mask > 0.5) if len(y_coords) == 0: return "no_polyp" centroid_y = np.mean(y_coords) # 结肠镜图像中,有效视野通常在中间2/3区域 # 顶部1/6(0~H/6)和底部1/6(5H/6~H)为器械/气泡区域 top_boundary = image_height / 6 bottom_boundary = 5 * image_height / 6 if centroid_y < top_boundary: return "suspicious_top" elif centroid_y > bottom_boundary: return "suspicious_bottom" else: return "plausible" # 统计测试集位置合理性 plausibility_stats = {"plausible": 0, "suspicious_top": 0, "suspicious_bottom": 0, "no_polyp": 0} for i in range(len(test_dataset)): sample = test_dataset[i] mask_np = sample["mask"].numpy().squeeze() result = check_anatomical_plausibility(mask_np, image_height=384) plausibility_stats[result] += 1 print("Anatomical plausibility distribution:") for k, v in plausibility_stats.items(): print(f" {k}: {v}/{len(test_dataset)} ({v/len(test_dataset)*100:.1f}%)")若suspicious_top占比超过15%,需检查数据集是否混入未裁剪的原始内镜视频帧(含操作手柄),或模型是否过拟合设备品牌特征。
4.3 临床指标计算:Dice与ASD的PyTorch原生实现
最终评估必须回归临床指标。Dice系数衡量重叠率,ASD(Average Surface Distance)衡量边界距离,二者缺一不可。
import torch import torch.nn.functional as F def compute_dice_coefficient(pred: torch.Tensor, target: torch.Tensor, smooth: float = 1e-6) -> float: """ 计算Dice系数(PyTorch原生,支持GPU) :param pred: 预测掩码 (B,1,H,W), sigmoid输出 :param target: 真实掩码 (B,1,H,W), 0/1 :param smooth: 平滑项 :return: Dice系数(标量) """ pred_flat = pred.view(pred.size(0), -1) target_flat = target.view(target.size(0), -1) intersection = (pred_flat * target_flat).sum(dim=1) dice = (2. * intersection + smooth) / (pred_flat.sum(dim=1) + target_flat.sum(dim=1) + smooth) return dice.mean().item() def compute_asd(pred: torch.Tensor, target: torch.Tensor) -> float: """ 计算平均表面距离(简化版,基于欧氏距离变换) :param pred: 预测掩码 (B,1,H,W) :param target: 真实掩码 (B,1,H,W) :return: ASD(毫米,假设1px=0.1mm) """ # 转换为numpy进行距离变换(torchvision暂无原生实现) pred_np = pred.cpu().numpy().squeeze(1) # (B,H,W) target_np = target.cpu().numpy().squeeze(1) asd_list = [] for i in range(len(pred_np)): # 计算预测边界到真实边界的平均距离 pred_dist = cv2.distanceTransform((pred_np[i] > 0.5).astype(np.uint8), cv2.DIST_L2, 3) target_dist = cv2.distanceTransform((target_np[i] > 0.5).astype(np.uint8), cv2.DIST_L2, 3) # 表面距离 = pred边界上点到target边界的最短距离 pred_boundary = cv2.morphologyEx((pred_np[i] > 0.5).astype(np.uint8), cv2.MORPH_GRADIENT, np.ones((3,3))) if pred_boundary.sum() > 0: surface_dist = (pred_dist * pred_boundary).sum() / pred_boundary.sum() else: surface_dist = 0 asd_list.append(surface_dist * 0.1) # 转换为毫米 return np.mean(asd_list) if asd_list else 0.0 # 在测试循环中使用 model.eval() dice_scores = [] asd_scores = [] with torch.no_grad(): for batch in dataloaders["test"]: images = batch["image"].cuda() targets = batch["mask"].cuda() preds = model(images) # 假设model输出logits,需sigmoid preds_sigmoid = torch.sigmoid(preds) dice = compute_dice_coefficient(preds_sigmoid, targets) asd = compute_asd(preds_sigmoid, targets) dice_scores.append(dice) asd_scores.append(asd) print(f"Mean Dice: {np.mean(dice_scores):.4f} ± {np.std(dice_scores):.4f}") print(f"Mean ASD: {np.mean(asd_scores):.4f} mm ± {np.std(asd_scores):.4f}")此实现避免依赖medpy等第三方库,全部使用OpenCV+PyTorch原语,确保在无网络环境的医疗设备上可部署。ASD计算中0.1的像素-毫米换算因子需根据具体内镜型号校准,此处为示例值。
5. 息肉分割数据集的终极验证技巧:用Grad-CAM定位模型注意力焦点
即使Dice达到0.92,模型也可能在“作弊”——例如通过识别内镜器械反光区域而非息肉纹理进行预测。必须验证模型关注区域是否与临床专家标注区域一致。Grad-CAM是最轻量级的可解释性工具。
import torch import torch.nn.functional as F from torchvision import models def generate_gradcam(model: torch.nn.Module, image: torch.Tensor, target_layer: torch.nn.Module) -> np.ndarray: """ 生成Grad-CAM热力图 :param model: 训练好的分割模型(需支持forward返回features) :param image: 输入图像 (1,3,H,W) :param target_layer: 最后一个卷积层(如model.encoder.layer4[-1]) :return: 热力图 (H,W) """ model.eval() image.requires_grad_(True) # 前向传播获取特征图 features = model.encoder(image) # 假设模型有encoder属性 output = model.segmentation_head(features) # 假设分割头 # 获取预测最大值对应的类别(息肉类) pred_class = output.argmax(dim=1, keepdim=True) # (1,1,H,W) # 反向传播计算梯度 model.zero_grad() output.backward(torch.ones_like(output), retain_graph=True) # 获取目标层梯度 gradients = target_layer.weight.grad pooled_gradients = torch.mean(gradients, dim=[0, 2, 3]) # 加权组合特征图 features = features[0] for i in range(features.shape[0]): features[i, :, :] *= pooled_gradients[i] heatmap = torch.mean(features, dim=0).detach().cpu().numpy() heatmap = np.maximum(heatmap, 0) # ReLU heatmap /= np.max(heatmap) # 归一化 return heatmap # 使用示例(需适配具体模型结构) # cam = generate_gradcam(model, sample["image"].unsqueeze(0).cuda(), model.encoder.layer4[-1]) # overlay = cv2.applyColorMap(np.uint8(255*cam), cv2.COLORMAP_JET) # cv2.imshow("Grad-CAM", overlay)将Grad-CAM热力图与原始图像叠加,若热点集中在息肉区域(而非器械、文字水印、气泡),则证明模型学习到了临床相关特征。这是数据集质量与模型鲁棒性的双重验证——当热力图与医生标注掩码的IoU>0.6时,该数据集才真正具备临床价值。
本文还有配套的精品资源,点击获取