news 2026/10/1 12:45:48

基于U-Net的COVID肺部感染分割实战:从数据集到训练全流程

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
基于U-Net的COVID肺部感染分割实战:从数据集到训练全流程

简介:这份资源面向医学图像分割方向的算法学习者与研究者,提供约2500张256×256分辨率的肺部感染(COVID)图像分割数据,前景标注为感染区域,mask采用前景255的二值图像,便于直观观察与训练。数据在DRIVE数据集基础上做了扩充,并按训练集1864张、验证集466张、测试集583张划分,每个子集均包含images图片目录与masks标签目录,可直接用于模型训练与评估。压缩包共约2000个文件,以1998个png图像与标签为主,另附1个txt说明和1个py脚本,整体约80.6MB,体积轻便易于下载与部署。其中可视化脚本可随机抽取一张图片,展示原始图像、GT图像及GT在原图上的蒙板效果并保存至当前目录,方便快速核验标注质量。目前已有402人学习关注,适合需要开展肺部感染分割实验、复现分割流程或扩充医学影像数据的中高级读者参考使用。

1. 肺部感染 COVID 分割数据集:2500 张 256×256 样本到底能跑出什么结果

手里这份医学图像分割数据集,核心就一件事:把 256×256 的肺部影像里感染区域抠出来。训练集 1864 对、验证集 466 对、测试集 583 对,合计约 2900 张图和对应 mask,标签是前景 255 的二值图,肉眼一看就知道哪里是感染灶。它适合谁?想入门医学图像分割但拿不到临床数据的算法同学、要做 COVID 病灶定量分析原型的工程师、以及需要一套干净二值标签来验证 U-Net 类模型的人。我第一眼看到「对 DRIVE 数据集做了扩充」这句时愣了一下——DRIVE 是眼底血管数据集,这里应该是借用了它的目录组织习惯,images + masks 双目录结构,训练/验证/测试三分。别被这句话带偏,实际内容就是肺部感染分割,不是眼底。这份资源能解决的核心痛点是:你不需要自己从 DICOM 开始标注,直接拿到就能喂进分割网络,省掉最耗时的数据准备环节。

2. 目录结构与标签格式:先搞清楚 images 和 masks 怎么对齐

2.1 三分目录的实际组织方式

拿到压缩包解压后,常见做法是看到这样的层级:

dataset/ ├── train/ │ ├── images/ │ │ ├── covid_0001.png │ │ └── ... │ └── masks/ │ ├── covid_0001.png │ └── ... ├── val/ │ ├── images/ │ └── masks/ └── test/ ├── images/ └── masks/

训练集 1864 对、验证集 466 对、测试集 583 对,注意验证集和测试集数量不一样,别想当然按 8:1:1 去反推。images 里是原始肺部影像,masks 里是二值标签,前景 255、背景 0。文件名一一对应,这是最省心的设计——你不需要额外维护一个映射表,直接按文件名配对就行。

我一般会先跑一段校验脚本,确认三件事:图片和 mask 数量是否相等、文件名是否完全交集、mask 的像素值是否只有 0 和 255。这三个检查能挡掉后面 80% 的玄学报错。

import os import numpy as np from PIL import Image def check_dataset(root): for split in ['train', 'val', 'test']: img_dir = os.path.join(root, split, 'images') mask_dir = os.path.join(root, split, 'masks') imgs = set(os.listdir(img_dir)) masks = set(os.listdir(mask_dir)) # 数量与文件名交集检查 print(f"{split}: images={len(imgs)}, masks={len(masks)}, " f"only_in_images={len(imgs - masks)}, only_in_masks={len(masks - imgs)}") # 抽查 mask 像素值分布 sample = list(masks)[:5] for name in sample: m = np.array(Image.open(os.path.join(mask_dir, name))) uniq = np.unique(m) print(f" {name}: shape={m.shape}, dtype={m.dtype}, unique={uniq}") check_dataset('./dataset')

这段脚本的逻辑很直白:用集合运算找出只出现在一边的文件,再抽查 mask 的唯一值。参数上,root指向解压后的数据集根目录。如果only_in_images或only_in_masks不为 0,说明配对有问题,后面训练必然对不上。mask 的unique应该输出[0 255],如果出现中间值,说明标签不是严格二值,需要先做阈值化。

2.2 二值 mask 的读取陷阱

PNG 格式存二值图有个经典坑:PIL 读进来可能是mode='P'或mode='L',前者是调色板模式,直接转 numpy 会得到索引值而不是 0/255。我踩过一次,mask 读出来全是 0 和 1,模型训练 loss 看着在降,但预测出来全是背景。后来强制加了一步转换:

def load_mask(path): m = Image.open(path) # 调色板模式或灰度模式统一转成 L 再二值化 if m.mode != 'L': m = m.convert('L') arr = np.array(m) # 保险起见做一次阈值化,防止有压缩噪声 arr = (arr > 127).astype(np.uint8) * 255 return arr

convert('L')把调色板或 RGB 统一成 8 位灰度,> 127做阈值化,再乘 255 还原成标准二值。这一步多花几毫秒,但能避免「训练正常、推理全黑」这种查半天的黑匣子问题。参数 127 是经验值,因为原始标签就是 0/255,取中间值最稳。

2.3 为什么 256×256 是个务实的选择

256×256 分辨率在医学分割里属于「够用但不奢侈」。原始 CT 或 X 光往往 512×512 甚至更大,直接训练显存吃不消。降到 256 后,单张图占显存约 256×256×3×4 字节 ≈ 786KB(float32),batch size 开到 16 也就 12MB 左右的特征图开销,一张 8GB 卡跑 U-Net 绰绰有余。代价是小结节或边缘模糊的感染区域可能丢细节,如果你的任务对微小病灶敏感,常见做法是先在这份数据上预训练,再拿原始高分辨率数据微调。这份数据集的价值在于让你快速把 pipeline 跑通,而不是直接拿去发顶会。

3. 从零搭一个 U-Net 训练流程:数据加载、模型、损失函数怎么配

3.1 Dataset 和 DataLoader 的写法

PyTorch 的 Dataset 是标准入口,关键是保证 image 和 mask 同步做增强。我一般用 albumentations,因为它对 image 和 mask 的几何变换能自动对齐。

import torch from torch.utils.data import Dataset, DataLoader import albumentations as A from albumentations.pytorch import ToTensorV2 import cv2 import os class CovidSegDataset(Dataset): def __init__(self, root, split='train', size=256): self.img_dir = os.path.join(root, split, 'images') self.mask_dir = os.path.join(root, split, 'masks') self.names = sorted(os.listdir(self.img_dir)) # 训练集加增强,验证测试只做 resize 和归一化 if split == 'train': self.tf = A.Compose([ A.Resize(size, size), A.HorizontalFlip(p=0.5), A.RandomBrightnessContrast(p=0.3), A.Normalize(mean=(0.5,), std=(0.5,)), ToTensorV2(), ]) else: self.tf = A.Compose([ A.Resize(size, size), A.Normalize(mean=(0.5,), std=(0.5,)), ToTensorV2(), ]) def __len__(self): return len(self.names) def __getitem__(self, idx): name = self.names[idx] img = cv2.imread(os.path.join(self.img_dir, name), cv2.IMREAD_GRAYSCALE) mask = cv2.imread(os.path.join(self.mask_dir, name), cv2.IMREAD_GRAYSCALE) # mask 二值化,确保只有 0/1 mask = (mask > 127).astype('float32') out = self.tf(image=img, mask=mask) image = out['image'] # [1, H, W] mask = out['mask'].unsqueeze(0) # [1, H, W] return image, mask

逻辑说明:cv2.IMREAD_GRAYSCALE保证读进来就是单通道,省去后面转通道的麻烦。Normalize用 mean=0.5、std=0.5 把像素拉到 [-1,1],这是医学图像里常用的简化归一化,因为灰度图没有 ImageNet 那套 RGB 均值。增强只加水平翻转和亮度对比度,医学图像别乱用旋转和裁剪,解剖结构的方向和完整性有临床意义。参数size=256和数据集原生分辨率一致,不需要额外缩放。

DataLoader 那边:

train_ds = CovidSegDataset('./dataset', 'train') val_ds = CovidSegDataset('./dataset', 'val') train_loader = DataLoader(train_ds, batch_size=16, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=8, shuffle=False, num_workers=4)

batch_size=16在 8GB 显存上跑 U-Net 比较稳,num_workers=4看 CPU 核数调整,pin_memory=True加速 GPU 拷贝。

3.2 U-Net 模型的最小实现

不引入额外库,手写一个标准 U-Net,方便你改结构。

import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.net = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.net(x) class UNet(nn.Module): def __init__(self, in_ch=1, out_ch=1): super().__init__() self.d1 = DoubleConv(in_ch, 64) self.d2 = DoubleConv(64, 128) self.d3 = DoubleConv(128, 256) self.d4 = DoubleConv(256, 512) self.pool = nn.MaxPool2d(2) self.bottleneck = DoubleConv(512, 1024) self.up4 = nn.ConvTranspose2d(1024, 512, 2, stride=2) self.u4 = DoubleConv(1024, 512) self.up3 = nn.ConvTranspose2d(512, 256, 2, stride=2) self.u3 = DoubleConv(512, 256) self.up2 = nn.ConvTranspose2d(256, 128, 2, stride=2) self.u2 = DoubleConv(256, 128) self.up1 = nn.ConvTranspose2d(128, 64, 2, stride=2) self.u1 = DoubleConv(128, 64) self.out = nn.Conv2d(64, out_ch, 1) def forward(self, x): c1 = self.d1(x) c2 = self.d2(self.pool(c1)) c3 = self.d3(self.pool(c2)) c4 = self.d4(self.pool(c3)) bn = self.bottleneck(self.pool(c4)) x = self.u4(torch.cat([self.up4(bn), c4], dim=1)) x = self.u3(torch.cat([self.up3(x), c3], dim=1)) x = self.u2(torch.cat([self.up2(x), c2], dim=1)) x = self.u1(torch.cat([self.up1(x), c1], dim=1)) return self.out(x)

in_ch=1对应灰度输入,out_ch=1输出单通道 logits。跳跃连接用torch.cat拼接,这是 U-Net 恢复空间细节的关键。参数量约 31M,256×256 输入下显存占用可控。

3.3 损失函数与评估指标

二值分割最常用 Dice Loss + BCE 组合。Dice 直接优化重叠度,BCE 稳定梯度。

class DiceBCELoss(nn.Module): def __init__(self, weight=0.5): super().__init__() self.weight = weight self.bce = nn.BCEWithLogitsLoss() def forward(self, logits, targets): bce = self.bce(logits, targets) probs = torch.sigmoid(logits) # 展平后计算 Dice probs = probs.view(-1) targets = targets.view(-1) inter = (probs * targets).sum() dice = 1 - (2 * inter + 1e-6) / (probs.sum() + targets.sum() + 1e-6) return self.weight * bce + (1 - self.weight) * dice

weight=0.5是起点,如果验证集 Dice 波动大,可以调到 0.3 让 Dice 主导。1e-6防止除零。评估时用阈值 0.5 把概率转成二值,再算 Dice 和 IoU。

@torch.no_grad() def evaluate(model, loader, device): model.eval() dice_sum, iou_sum, n = 0, 0, 0 for img, mask in loader: img, mask = img.to(device), mask.to(device) logits = model(img) pred = (torch.sigmoid(logits) > 0.5).float() inter = (pred * mask).sum().item() union = pred.sum().item() + mask.sum().item() - inter dice_sum += (2 * inter + 1e-6) / (pred.sum().item() + mask.sum().item() + 1e-6) iou_sum += (inter + 1e-6) / (union + 1e-6) n += 1 return dice_sum / n, iou_sum / n

这套评估在验证集 466 张上跑一遍大概十几秒,能快速判断模型有没有学崩。

4. 可视化脚本怎么用:一张图看清原图、GT 和叠加蒙板

4.1 可视化脚本的核心逻辑

数据集自带一个可视化脚本,随机抽一张图,展示原图、GT 二值图、GT 叠加在原图上的蒙板,并保存到当前目录。我一般会把它改造成可指定索引的版本,方便反复看同一张。

import os import random import numpy as np import cv2 import matplotlib.pyplot as plt def visualize(root, split='val', idx=None, save_path='vis_result.png'): img_dir = os.path.join(root, split, 'images') mask_dir = os.path.join(root, split, 'masks') names = sorted(os.listdir(img_dir)) if idx is None: idx = random.randint(0, len(names) - 1) name = names[idx] img = cv2.imread(os.path.join(img_dir, name), cv2.IMREAD_GRAYSCALE) mask = cv2.imread(os.path.join(mask_dir, name), cv2.IMREAD_GRAYSCALE) mask_bin = (mask > 127).astype(np.uint8) # 构造红色叠加蒙板 overlay = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR) overlay[mask_bin == 1] = [0, 0, 255] blended = cv2.addWeighted(cv2.cvtColor(img, cv2.COLOR_GRAY2BGR), 0.6, overlay, 0.4, 0) fig, axes = plt.subplots(1, 3, figsize=(12, 4)) axes[0].imshow(img, cmap='gray'); axes[0].set_title('Original') axes[1].imshow(mask_bin, cmap='gray'); axes[1].set_title('GT Mask') axes[2].imshow(cv2.cvtColor(blended, cv2.COLOR_BGR2RGB)); axes[2].set_title('Overlay') for ax in axes: ax.axis('off') plt.tight_layout() plt.savefig(save_path, dpi=150) plt.close() print(f"Saved to {save_path}, sample: {name}") visualize('./dataset', split='val', idx=0)

逻辑说明:mask_bin是 0/1 矩阵,用它做布尔索引把原图对应位置染红,再和原图做加权融合,0.6/0.4的权重让底图仍可见。idx=None时随机抽,指定索引可复现。保存用dpi=150,够看清边缘。参数split可切 train/val/test,建议先看 val,因为 val 没参与训练,能反映真实泛化。

4.2 用可视化做快速质检

这个脚本不只是好看,它是你发现数据问题的第一道防线。我习惯在训练前随机抽 20 张跑一遍,重点看三件事:mask 是否和原图解剖位置对齐、有没有全黑或全白的异常标签、感染区域边界是否合理。有一次发现某几张 mask 整体偏移了几个像素,追查发现是某批数据在 resize 时用了不同的插值方式。这种问题不看图根本发现不了,loss 曲线照样好看,但模型学到的就是错位映射。

提示:可视化脚本保存的图片默认在当前工作目录,批量跑的时候记得改save_path,否则会互相覆盖。

5. 避坑与排查:训练不收敛、Dice 虚高、显存爆掉的真实原因

5.1 现象:loss 降到 0.1 以下但预测全黑

原因:mask 读取时没做二值化,或者用了mode='P'的调色板图,导致标签值域不是 0/255 而是 0/1 甚至索引值。模型学到的是「全预测背景」也能拿到很低的 BCE,因为背景占绝大多数像素。

解决:在 Dataset 里强制(mask > 127).astype('float32'),并在训练前用第 2 章的校验脚本确认 mask 唯一值只有 0 和 255。如果已经是 0/1,检查归一化有没有把 mask 也一起 Normalize 了——albumentations 的 Normalize 默认只作用 image,但如果你手动把 mask 也塞进去就会出事。

5.2 现象:验证集 Dice 0.9 以上,换测试集掉到 0.6

原因:验证集和测试集分布不一致。这份数据验证集 466 张、测试集 583 张,如果验证集里简单样本偏多,Dice 就会虚高。另一个可能是你在验证集上调了太多轮超参,相当于变相过拟合。

解决:训练中只拿验证集做早停和模型选择,测试集在最终确定模型前不要碰。如果测试集 Dice 明显低,先可视化几张测试集预测结果,看是漏检还是误检。漏检多就降阈值到 0.3,误检多就升到 0.6。别小看这 0.2 的调整,医学分割里阈值对 Dice 的影响能有 5 个点。

5.3 现象:训练到一半 CUDA out of memory

原因:256×256 输入下 U-Net 的中间特征图不小,batch size 16 在 8GB 卡上如果开了num_workers且用了 pin_memory,再加上验证时没加torch.no_grad(),显存会累积。

解决:验证和推理一律包@torch.no_grad();训练时如果爆显存,先把 batch size 降到 8,再考虑用torch.cuda.amp混合精度。混合精度能把显存占用砍掉近一半,代价是偶尔需要调 loss scale,但对 U-Net 这种结构通常很稳。

scaler = torch.cuda.amp.GradScaler() for img, mask in train_loader: img, mask = img.to(device), mask.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(): logits = model(img) loss = criterion(logits, mask) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()

5.4 现象:Dice 在 0.7 附近震荡上不去

原因:学习率太大或太小。太大导致在最优解附近跳,太小导致卡在局部。另一个常见原因是 Dice Loss 和 BCE 的权重没调好,BCE 主导时模型偏向像素级准确率,而医学分割更看重区域重叠。

解决:先用lr=1e-3跑 10 个 epoch 看趋势,如果 loss 震荡就降到3e-4。损失权重从weight=0.5开始,如果 Dice 上不去就降到 0.3。还可以加ReduceLROnPlateau,验证 Dice 3 轮不升就砍半学习率。

5.5 现象:可视化叠加图里蒙板颜色溢出到肺外

原因:mask 和原图尺寸不一致,或者 resize 时用了不同的插值。原图 resize 常用双线性,mask 必须用最近邻,否则边缘会产生中间值,二值化后边界偏移。

解决:在 albumentations 里对 mask 单独指定interpolation=cv2.INTER_NEAREST。如果你用的是自定义 transform,确保 mask 的 resize 和几何变换都走最近邻。

A.Compose([ A.Resize(256, 256, interpolation=cv2.INTER_LINEAR), # 对 image ], additional_targets={'mask': 'mask'}) # mask 的 resize 需单独处理或使用支持 mask 插值配置的版本

6. 进阶技巧:用测试集做一次严格的泛化验证与阈值搜索

训练流程跑通后,别急着收工。这份数据集给了独立的测试集 583 张,这是你验证泛化能力的唯一依据。我一般会做两件事:一是固定模型权重,在测试集上跑一遍完整评估;二是做阈值搜索,找到 Dice 最优的判定阈值,而不是默认 0.5。

@torch.no_grad() def threshold_search(model, loader, device, thresholds=np.arange(0.3, 0.71, 0.05)): model.eval() all_probs, all_masks = [], [] for img, mask in loader: img = img.to(device) logits = model(img) probs = torch.sigmoid(logits).cpu().numpy() all_probs.append(probs) all_masks.append(mask.numpy()) all_probs = np.concatenate(all_probs, axis=0) all_masks = np.concatenate(all_masks, axis=0) results = [] for t in thresholds: pred = (all_probs > t).astype(np.float32) inter = (pred * all_masks).sum() dice = (2 * inter + 1e-6) / (pred.sum() + all_masks.sum() + 1e-6) results.append((t, dice)) print(f"threshold={t:.2f}, dice={dice:.4f}") best = max(results, key=lambda x: x[1]) print(f"Best threshold: {best[0]:.2f}, Dice: {best[1]:.4f}") return best

这段代码把测试集所有样本的预测概率和 GT 缓存下来,然后遍历 0.3 到 0.7 的阈值,每个阈值算一次全局 Dice。注意这里是全局 Dice,把所有像素展平后一起算,而不是每张图算完再平均。两种算法结果会有差异,全局 Dice 对大病灶更敏感,逐图平均对小病灶更公平。医学分割里我倾向逐图平均,因为临床更关心每个病例的分割质量,而不是整体像素占比。

参数上,thresholds从 0.3 到 0.7 步长 0.05,覆盖了常见的最优区间。如果你的模型校准得好,最优阈值通常在 0.4 到 0.6 之间;如果偏离很远,说明模型输出的概率不可信,可能需要加温度缩放或重新检查损失函数。

还有一个容易被忽略的点:测试集评估要固定随机种子,包括 DataLoader 的 shuffle 和任何数据增强。测试阶段不该有随机性,否则你每次跑出来的 Dice 都不一样,没法比较。

def set_seed(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False

deterministic=True会让 cuDNN 选确定性算法,速度略慢但结果可复现。benchmark=False同理。这套设置我在每次最终评估前都会跑一遍,血泪经验是:有一次没固定种子,测试 Dice 在 0.82 到 0.86 之间跳,白白多花了一下午排查。

最后说个我自己的习惯:每次拿到新数据集,先跑可视化脚本看 20 张,再跑校验脚本确认标签格式,然后才开训练。这三步走完,后面基本不会遇到「训练半天发现数据读错了」这种后悔药都没得吃的情况。这份 COVID 分割数据集结构清晰、标签干净,256×256 的尺寸对硬件友好,适合作为医学分割的入门实战。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/1 12:45:33

Dify集成DbHub MCP:让AI用SQL精准处理Excel表格

把Excel直接丢给大模型让它“总结一下”,这个操作我一开始也以为是AI最擅长的事,结果真正上手才发现,这种“文本解析式”读表在稍微复杂的文件面前几乎不可用。合并单元格、跨Sheet引用、公式缓存、空行空列,任何一个因素都能让AI…

作者头像 李华
网站建设 2026/10/1 12:45:32

智能体安全选型不能只看功能清单:运行时硬指标测评,悬镜安全国内智能体安全领域领跑实践

随着企业 AI 智能体逐步进入研发、办公、运营等业务场景,安全团队面对的不再只是模型输出内容风险,而是一套具备自主推理、工具调用、文件读写、外网访问和组件扩展能力的动态系统。智能体一旦被劫持,攻击动作可能在极短时间内完成&#xff0…

作者头像 李华
网站建设 2026/10/1 12:44:41

SpringBoot+Vue班级事务管理系统:考勤、班费、请假全流程设计

1. 这个题目为什么经典:班级事务系统的业务边界与角色需求每年到毕业设计选题的时候,“SpringBoot 高校 班级事务管理”这类题目都会出现。我第一次看到“河北水利电力学院班级事务管理系统”这个课题时,第一反应是:这不就是一个…

作者头像 李华
网站建设 2026/10/1 12:44:39

Linux第一次作业:环境搭建、命令查询与脚本编写指南

兄弟,如果你刚交完或者正准备做人生的第一次 Linux 作业,我先给你交个底:这门课从来不是在考你背了多少条命令,它真正的作业其实是——让你在真实的环境里完成一系列“小到不能再小”的任务,然后通过这些任务把“Linux…

作者头像 李华
网站建设 2026/10/1 12:44:32

FastAPI文档页白屏?从CDN替换到离线部署的完整方案

如果你把一个 FastAPI 服务部署到公司内网,大概率会遇到一个让人头疼的场景:业务代码跑得好好的,但打开 /docs 页面时一片空白,F12 里飘着十几个红色加载失败。问题十有八九出在 Swagger UI 的 CDN 资源上——FastAPI 默认从公共…

作者头像 李华