尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch UNet实战指南:从数据准备到可调试分割闭环

PyTorch UNet实战指南:从数据准备到可调试分割闭环 简介本资源是一份面向深度学习初学者与图像分割实践者的PyTorch U-Net实战项目聚焦医学及通用图像分割场景提供从网络搭建、数据准备到端到端训练的完整实现路径。压缩包共27个文件含7个核心Python脚本如net.py定义U-Net架构、train.py封装训练流程、data.py处理数据加载、5张示例图像与结果图png、5份Markdown文档含中英文README与参数说明、5个XML标注文件支持自定义数据集转换辅以Git配置、IDE配置及LICENSE等辅助文件整体仅602KB轻量易部署。已有760人学习下载资源结构清晰根目录下分utils、data、result等模块含mask生成、评估脚本及训练日志管理特别适配个人小规模数据集训练需求开箱即可运行并快速验证模型效果。1. PyTorch版UNet不是“抄个net.py就能跑通”的黑匣子它是一套可调试、可插拔、能落地到你手头那几张CT片或工业缺陷图的完整训练闭环你手上有27张标注好的PCB板缺陷图或者刚从医院导出的13例肝脏CT切片每张带手动勾画的mask想快速验证一个分割模型是否能初步圈出病灶/焊点虚焊——这时候搜“pytorch unet 训练自己的数据集”90%的教程会给你扔一个net.py、一个train.py、再加一句“把图片放data目录下运行就行”。结果呢RuntimeError: expected stride to be a multiple of 32卡在dataloaderlossnan飘在终端test.py跑出来全是灰块。这不是代码有问题是整套流程缺了数据-模型-训练-评估四层对齐的锚点。这个pytorch-UNet.zip包之所以值得拆是因为它把UNet从论文公式落地成可调试的工程实体utils.py里封装了带边界裁剪的mask生成逻辑make_mask_data.py能自动把VOC格式转成UNet吃的双通道输入get_evaluation.py不只算IoU还输出逐类召回率和混淆矩阵热力图。它不承诺“一键训练”但保证你改三行路径、调两个参数就能看到第一张预测图上的红色轮廓线——这才是工程师要的“最小可行分割”。2. 从零构建UNet网络为什么必须重写net.py而不是直接import torchvision.modelsUNet不是ResNet那种即插即用的骨干网它的跳跃连接、上采样方式、通道数缩放策略每一处都和你的数据特性强耦合。这个包里的net.py不是玩具代码而是按医学/工业场景真实约束设计的编码器用3×3卷积BNReLU组合替代简单堆叠解码器强制使用nn.ConvTranspose2d而非nn.UpsampleConv2d避免棋盘伪影跳跃连接前做channel对齐nn.Conv2d(in_c, out_c, 1)。下面拆解核心模块实现逻辑2.1 编码器4级下采样中的通道膨胀与信息保全class UNetEncoder(nn.Module): def __init__(self, in_channels3, init_features32): super(UNetEncoder, self).__init__() features init_features # Level 0: 输入层不做下采样仅特征提取 self.encoder1 UNetBlock(in_channels, features) # out: [B,32,H,W] # Level 1: 第一次下采样H/2, W/2, channel×2 self.pool1 nn.MaxPool2d(kernel_size2, stride2) # 注意stride2而非1避免重叠池化 self.encoder2 UNetBlock(features, features * 2) # out: [B,64,H/2,W/2] # Level 2: 第二次下采样H/4, W/4, channel×4 self.pool2 nn.MaxPool2d(kernel_size2, stride2) self.encoder3 UNetBlock(features * 2, features * 4) # out: [B,128,H/4,W/4] # Level 3: 第三次下采样H/8, W/8, channel×8 self.pool3 nn.MaxPool2d(kernel_size2, stride2) self.encoder4 UNetBlock(features * 4, features * 8) # out: [B,256,H/8,W/8] # Level 4: 最深层瓶颈无池化纯卷积增强语义 self.bottleneck UNetBlock(features * 8, features * 16) # out: [B,512,H/8,W/8] def forward(self, x): enc1 self.encoder1(x) # [B,32,H,W] enc2 self.encoder2(self.pool1(enc1)) # [B,64,H/2,W/2] enc3 self.encoder3(self.pool2(enc2)) # [B,128,H/4,W/4] enc4 self.encoder4(self.pool3(enc3)) # [B,256,H/8,W/8] bottleneck self.bottleneck(enc4) # [B,512,H/8,W/8] return enc1, enc2, enc3, enc4, bottleneck参数说明init_features32是关键调节旋钮。当你的图像分辨率高如1024×1024 CT图且GPU显存≥12GB时可设为64以提升小目标分割能力若只有27张训练图且显存≤8GB必须压到16否则bottleneck层会OOM。UNetBlock内部是Conv-BN-ReLU-Conv-BN-ReLU双卷积结构比单卷积更能抑制梯度消失——这是作者在README.md里没明说但utils.py中所有预处理函数都默认适配的底层假设。2.2 解码器转置卷积的步长陷阱与跳跃连接的通道对齐class UNetDecoder(nn.Module): def __init__(self, init_features32): super(UNetDecoder, self).__init__() features init_features # 上采样层必须用ConvTranspose2d且kernel_size2,stride2,padding0 # 这是避免棋盘效应的硬性要求见arXiv:1806.02679 self.upconv4 nn.ConvTranspose2d( features * 16, features * 8, kernel_size2, stride2 ) # 将[B,512,H/8,W/8] → [B,256,H/4,W/4] self.decoder4 UNetBlock((features * 8) * 2, features * 8) # 拼接enc4后通道翻倍 self.upconv3 nn.ConvTranspose2d( features * 8, features * 4, kernel_size2, stride2 ) # [B,256,H/4,W/4] → [B,128,H/2,W/2] self.decoder3 UNetBlock((features * 4) * 2, features * 4) self.upconv2 nn.ConvTranspose2d( features * 4, features * 2, kernel_size2, stride2 ) # [B,128,H/2,W/2] → [B,64,H,W] self.decoder2 UNetBlock((features * 2) * 2, features * 2) self.upconv1 nn.ConvTranspose2d( features * 2, features, kernel_size2, stride2 ) # [B,64,H,W] → [B,32,2H,2W] self.decoder1 UNetBlock(features * 2, features) def forward(self, enc1, enc2, enc3, enc4, bottleneck): # 注意upconv后必须crop enc_i以匹配空间尺寸 dec4 torch.cat([self.upconv4(bottleneck), self._center_crop(enc4, bottleneck)], dim1) dec4 self.decoder4(dec4) dec3 torch.cat([self.upconv3(dec4), self._center_crop(enc3, dec4)], dim1) dec3 self.decoder3(dec3) dec2 torch.cat([self.upconv2(dec3), self._center_crop(enc2, dec3)], dim1) dec2 self.decoder2(dec2) dec1 torch.cat([self.upconv1(dec2), self._center_crop(enc1, dec2)], dim1) dec1 self.decoder1(dec1) return dec1 def _center_crop(self, layer, target_layer): 将layer中心裁剪至target_layer尺寸解决转置卷积导致的尺寸偏移 _, _, h, w target_layer.size() diff_y (layer.size()[2] - h) // 2 diff_x (layer.size()[3] - w) // 2 return layer[:, :, diff_y:(diff_y h), diff_x:(diff_x w)]关键逻辑_center_crop不是可选项是UNet训练收敛的生死线。nn.ConvTranspose2d(kernel_size2,stride2)理论上应完美上采样但实际因padding机制会导致输出尺寸比理论值大1像素如H/8→H/4时本该512→1024却变成1025。若不裁剪直接拼接torch.cat会报size mismatch。这个函数在utils.py中被所有训练脚本调用但99%的GitHub UNet复现者会忽略它——直到loss突然爆炸。2.3 输出头二分类与多分类的激活函数选择class UNet(nn.Module): def __init__(self, in_channels3, out_channels1, init_features32): super(UNet, self).__init__() self.encoder UNetEncoder(in_channels, init_features) self.decoder UNetDecoder(init_features) # 最终卷积层out_channels决定任务类型 self.final_conv nn.Conv2d(init_features, out_channels, kernel_size1) # 激活函数根据任务动态选择 if out_channels 1: self.activation nn.Sigmoid() # 二分类前景/背景 else: self.activation nn.Softmax(dim1) # 多分类如肿瘤/正常组织/坏死 def forward(self, x): enc1, enc2, enc3, enc4, bottleneck self.encoder(x) dec1 self.decoder(enc1, enc2, enc3, enc4, bottleneck) logits self.final_conv(dec1) # [B,out_c,H,W] return self.activation(logits)参数说明out_channels必须与你的mask通道数严格一致。若用make_mask_data.py生成单通道灰度mask0为背景255为前景则设为1若用SegmentationClass目录下的彩色PNG每个类别用不同RGB值需先用utils.py中的rgb_to_class_id()转换为单通道整数ID图此时out_channels等于类别数1含背景。这里埋着一个巨坑很多教程直接让final_conv输出3通道去拟合RGB mask结果sigmoid后三个通道互相干扰边缘严重模糊——正确做法永远是先转ID再训练。3. 数据准备为什么make_mask_data.py比LabelImg更适配UNet训练流UNet对输入数据有隐式假设图像与mask必须空间对齐、尺寸可被16整除、mask值域为[0,1]或[0,C-1]。make_mask_data.py不是简单的格式转换器它是专为UNet定制的数据管道预处理器。它解决三个核心问题1VOC格式的SegmentationObject.png常含抗锯齿边缘直接二值化会丢失细节2工业缺陷图常有极细裂纹resize到512×512后可能完全消失3医学CT窗宽窗位差异导致同一组织在不同序列中灰度值漂移。下面看它如何破局3.1 VOC转UNet抗锯齿mask的锐化重建# make_mask_data.py 核心逻辑节选 def voc_to_unet_mask(voc_mask_path: str, output_dir: str, class_names: List[str]): voc_mask_path: VOCdevkit/VOC2012/SegmentationObject/xxx.png class_names: [background, tumor, cyst] # 必须按ID顺序排列 mask cv2.imread(voc_mask_path, cv2.IMREAD_UNCHANGED) # 读取4通道PNG # 步骤1分离alpha通道并二值化VOC的SegmentationObject用alpha表示实例 if mask.shape[2] 4: alpha mask[:, :, 3] binary_mask (alpha 128).astype(np.uint8) * 255 else: # 若无alpha用HSV阈值提取前景针对彩色SegmentationObject hsv cv2.cvtColor(mask, cv2.COLOR_RGB2HSV) lower np.array([0, 0, 50]) upper np.array([180, 255, 255]) binary_mask cv2.inRange(hsv, lower, upper) # 步骤2形态学闭运算填充细小空洞修复抗锯齿导致的孔洞 kernel np.ones((3,3), np.uint8) binary_mask cv2.morphologyEx(binary_mask, cv2.MORPH_CLOSE, kernel) # 步骤3连通域分析保留最大连通域去除噪点 num_labels, labels, stats, centroids cv2.connectedComponentsWithStats(binary_mask) if num_labels 1: largest_idx np.argmax(stats[1:, cv2.CC_STAT_AREA]) 1 clean_mask np.where(labels largest_idx, 255, 0).astype(np.uint8) else: clean_mask binary_mask # 步骤4保存为UNet标准格式单通道0-255 filename os.path.basename(voc_mask_path) cv2.imwrite(os.path.join(output_dir, filename), clean_mask)逻辑说明这段代码直击VOC数据集痛点。VOC的SegmentationObject.png本质是实例分割图其边缘经抗锯齿处理后呈半透明过渡alpha值128~255直接cv2.threshold会切掉所有过渡区导致mask收缩。make_mask_data.py用alpha通道二值化形态学闭运算连通域筛选三步重建出紧贴物体边界的clean mask。这比LabelImg手动描边效率高10倍且保证所有训练图mask质量一致——你在data/SegmentationClass里看到的每一张PNG都是这样炼出来的。3.2 尺寸归一化16整除约束下的智能resize策略# utils.py 中的 resize_with_pad 函数 def resize_with_pad(image: np.ndarray, target_size: int 512) - np.ndarray: 将图像resize至target_size但保持长宽比并用0填充至正方形 确保最终尺寸可被16整除UNet最小下采样因子2^416 h, w image.shape[:2] scale target_size / max(h, w) new_h, new_w int(h * scale), int(w * scale) # 双线性插值resize resized cv2.resize(image, (new_w, new_h), interpolationcv2.INTER_LINEAR) # 计算padding使最终尺寸为target_size且可被16整除 pad_h target_size - new_h pad_w target_size - new_w # 向上取整到16的倍数 final_h ((target_size 15) // 16) * 16 final_w ((target_size 15) // 16) * 16 # 实际padding量 pad_h_final final_h - new_h pad_w_final final_w - new_w # 填充上下左右均分 padded cv2.copyMakeBorder( resized, toppad_h_final//2, bottompad_h_final//2 pad_h_final%2, leftpad_w_final//2, rightpad_w_final//2 pad_w_final%2, borderTypecv2.BORDER_CONSTANT, value0 ) return padded参数说明target_size512是默认值但你可以根据GPU显存调整。若显存紧张设为384384÷1624仍为整数若处理高分辨率病理图可设为768。关键在final_h/final_w的计算——它确保无论原始图比例如何最终输入UNet的tensor尺寸必为[B,3,512,512]或[B,3,384,384]等16的倍数。这是train.py中dataloader不报错的前提也是net.py里ConvTranspose2d能稳定上采样的物理基础。3.3 数据增强为什么旋转弹性形变比单纯Flip更有效# data.py 中的 train_transform 定义 train_transform A.Compose([ A.HorizontalFlip(p0.5), A.VerticalFlip(p0.5), A.RandomRotate90(p0.5), # 随机旋转90/180/270度 A.ElasticTransform( p0.7, alpha120, # 弹性变换强度 sigma120 * 0.05, # 高斯核标准差 alpha_affine120 * 0.03 # 仿射变换强度 ), A.RandomBrightnessContrast( brightness_limit0.2, contrast_limit0.2, p0.3 ), A.Normalize( mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225), max_pixel_value255.0 ), ToTensorV2() ])避坑逻辑A.ElasticTransform是此包的灵魂增强。医学图像中器官形变、工业图像中金属热胀冷缩都属于非刚性形变。单纯HorizontalFlip只能模拟镜像而弹性形变通过控制点网格扭曲能生成更真实的病理组织拉伸/压缩效果。参数alpha120对应中等强度实测值alpha80形变太弱150会破坏器官拓扑结构。注意p0.7——70%概率应用避免所有样本都被扭曲导致过拟合。这个配置在README.md的“训练技巧”章节有实测对比用弹性形变后小肿瘤检出率提升12.3%而仅用Flip提升仅2.1%。4. 训练与调试train.py里的四个隐藏开关与loss曲线诊断法train.py表面是标准PyTorch训练循环实则暗藏四个影响收敛的关键开关。它们不在命令行参数里而在代码深处且相互制约。不理解它们你看到的loss下降只是假象。4.1 学习率预热Warmup为什么前10轮loss必然震荡# train.py 片段 def get_scheduler(optimizer, warmup_epochs10, total_epochs100): 学习率预热余弦退火组合调度器 def lr_lambda(epoch): if epoch warmup_epochs: # 线性预热epoch0时lr0epochwarmup_epochs时lrbase_lr return float(epoch) / float(max(1, warmup_epochs)) else: # 余弦退火从base_lr平滑降至0 return 0.5 * (1.0 math.cos(math.pi * (epoch - warmup_epochs) / (total_epochs - warmup_epochs))) return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda) # 在main()中调用 scheduler get_scheduler(optimizer, warmup_epochs10, total_epochsparams.epochs)现象与原理如果你在第1-10轮看到loss从1.2跳到0.8再跳回1.0别慌——这是预热期的正常震荡。原因UNet初始权重随机前几轮梯度方向混乱若直接用base_lr1e-3权重更新幅度过大导致loss突变。预热机制强制lr从0线性增至1e-3让网络先用小步长“试探”梯度方向。warmup_epochs10是经验值少于5轮预热不足多于15轮收敛变慢。这个参数必须和batch_size联动调整——若batch_size从8提到16warmup_epochs应减半因梯度更稳定。4.2 损失函数选择Dice Loss BCE Loss的黄金配比# train.py 中的 loss_fn 定义 class DiceBCELoss(nn.Module): def __init__(self, dice_weight0.5): super(DiceBCELoss, self).__init__() self.dice_weight dice_weight self.bce_loss nn.BCELoss() def forward(self, inputs, targets): # BCE Loss bce self.bce_loss(inputs, targets) # Dice Loss smooth 1e-5 inputs_flat inputs.view(-1) targets_flat targets.view(-1) intersection (inputs_flat * targets_flat).sum() dice (2. * intersection smooth) / (inputs_flat.sum() targets_flat.sum() smooth) dice_loss 1 - dice # 加权求和 return self.dice_weight * dice_loss (1 - self.dice_weight) * bce # 初始化时设置 criterion DiceBCELoss(dice_weight0.7) # Dice占70%BCE占30%参数说明dice_weight0.7是此包针对小目标分割的调优值。Dice Loss对前景区域敏感分子是交集BCE Loss对全局像素分布敏感。当你的数据集中目标占比5%如CT中的微小结节纯BCE会因背景像素过多而主导梯度导致前景召回率低纯Dice又易受噪声干扰。0.7:0.3的配比在result/目录下的loss_curve.png中得到验证当dice_weight0.5时val_loss在50轮后平台期IoU停滞在0.62调至0.7后IoU升至0.68且曲线更平滑。这个值不能硬套需根据你的数据集前景占比微调占比1%用0.810%用0.5。4.3 梯度裁剪Gradient Clipping防止UNet深层梯度爆炸的保险丝# train.py 训练循环内 for batch_idx, (data, target) in enumerate(train_loader): optimizer.zero_grad() output model(data) loss criterion(output, target) loss.backward() # 关键梯度裁剪阈值设为1.0 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step()避坑逻辑UNet的跳跃连接虽缓解梯度消失但深层编码器尤其是bottleneck仍易梯度爆炸。max_norm1.0是经过utils.py中gradient_check()函数验证的安全阈值。若不裁剪lossnan通常在第3-5轮出现若设为5.0虽不报错但val_loss波动剧烈。这个值与init_features强相关当init_features64时需降至0.7init_features16时可放宽至1.5。train.py中注释明确写着“修改init_features后务必同步调整clip_grad_norm_的max_norm”。4.4 验证频率与早停Early Stopping如何避免过拟合而不浪费GPU时间# train.py 中的验证逻辑 best_iou 0.0 patience_counter 0 patience 15 # 连续15轮val_iou未提升则停止 for epoch in range(params.epochs): # 训练... if epoch % 5 0: # 每5轮验证一次平衡精度与速度 val_iou validate(model, val_loader, device) if val_iou best_iou: best_iou val_iou torch.save(model.state_dict(), best_model.pth) patience_counter 0 else: patience_counter 1 if patience_counter patience: print(fEarly stopping at epoch {epoch}) break参数说明epoch % 5 0是折中方案。每轮验证虽精准但validate()耗时是train_step的3倍需全图推理IoU计算每10轮验证又可能错过最佳保存点。5轮是实测最优在100轮训练中best_model.pth平均出现在第62轮而val_iou峰值在第65轮误差3轮可接受。patience15对应15×575轮无提升足够覆盖UNet常见的“平台期-回升”波动。若你的数据集极小50图建议改为patience8。5. 避坑指南UNet训练中五个血泪换来的“现象-原因-解决”清单注意以下问题均来自真实复现过程非理论推演。每个问题都对应pytorch-UNet.zip中某行代码的修改或某参数的调整。5.1 现象train.py运行到第2轮报错RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor) should be the same原因train.py中model.to(device)调用位置错误。原代码在DataLoader初始化后才执行model.to(device)但optimizer在之前已创建其内部状态仍绑定CPU权重。解决将model.to(device)移至optimizer定义之前并确保optimizer用model.parameters()初始化model UNet(in_channels3, out_channels1, init_features32) model.to(device) # 必须在此处 optimizer torch.optim.Adam(model.parameters(), lrparams.lr)5.2 现象test.py输出的result.png全是黑色或只有零星白点原因test.py中torch.no_grad()后未调用.cpu()和.numpy()导致output仍是GPU tensorcv2.imwrite无法处理。解决在保存前强制转CPU和numpyoutput model(image_tensor.unsqueeze(0)) # [1,1,H,W] pred_mask output.squeeze().cpu().numpy() # 关键.cpu().numpy() pred_mask (pred_mask 0.5).astype(np.uint8) * 255 cv2.imwrite(result.png, pred_mask)5.3 现象get_evaluation.py计算的IoU为0.0但肉眼可见预测图有重叠原因get_evaluation.py中mask读取方式错误。原代码用cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)但若mask是16位PNG常见于医学图像cv2.imread会截断高位导致所有像素值为0。解决统一用PIL.Image.open读取自动适配位深from PIL import Image mask np.array(Image.open(mask_path)) # 自动处理8/16位 pred np.array(Image.open(pred_path)) # 确保值域为0-255 mask (mask / mask.max() * 255).astype(np.uint8) if mask.max() 255 else mask5.4 现象训练loss持续下降但验证IoU停滞在0.4远低于预期原因data.py中train_transform和val_transform的归一化参数不一致。原包val_transform用ImageNet均值而train_transform用自定义均值导致训练/验证分布偏移。解决统一归一化参数在data.py顶部定义# 全局归一化参数根据你的数据集计算 MEAN [0.342, 0.287, 0.261] # 用utils.py中的calc_mean_std.py计算 STD [0.215, 0.198, 0.189] # train_transform和val_transform均使用此参数5.5 现象make_mask_data.py生成的mask边缘有毛刺与原图不严丝合缝原因cv2.morphologyEx的MORPH_CLOSE操作过度填充。原代码用kernelnp.ones((3,3))对细长裂纹会桥接不应连接的区域。解决改用椭圆核并降低迭代次数kernel cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3,3)) binary_mask cv2.morphologyEx(binary_mask, cv2.MORPH_CLOSE, kernel, iterations1)6. 进阶技巧用get_evaluation.py的混淆矩阵热力图定位模型弱点get_evaluation.py不只是算IoU它生成的confusion_matrix.png是调试UNet的终极地图。当你发现整体IoU达0.75但医生反馈“小结节总漏检”传统方法只能盲调loss权重而混淆矩阵能精确定位问题层——是编码器早期特征提取不足还是解码器上采样丢失细节下面教你如何用它做根因分析6.1 生成混淆矩阵从预测图到可解释热力图# 运行评估脚本需先运行test.py生成result/目录 python get_evaluation.py \ --pred_dir result/ \ --gt_dir data/SegmentationClass/ \ --num_classes 2 \ --save_dir evaluation/脚本会输出iou_per_class.txt: 各类IoU背景、前景confusion_matrix.npy: 2×2混淆矩阵数组confusion_matrix.png: 可视化热力图6.2 解读热力图三类典型模式与对应优化策略热力图模式数值表现根本原因优化动作主对角线亮右上角FP暗左下角FN亮FN0.35, FP0.05模型过于保守不敢预测前景降低final_conv后sigmoid阈值test.py中0.5改为0.3或增加Dice Loss权重至0.8主对角线亮右上角FP亮左下角FN暗FP0.28, FN0.03模型过度敏感把背景误判为前景增加train_transform中RandomBrightnessContrast强度或在net.py解码器最后加nn.Dropout2d(p0.1)主对角线暗非对角线亮对角线均0.5FP/FN均0.2编码器特征提取失败无法区分前景/背景检查net.py中UNetBlock的卷积核数是否与init_features匹配或更换编码器为ResNet34需重写UNetEncoder实战案例我曾用此包处理肺部CT结节数据confusion_matrix.png显示FN高达0.41。放大热力图发现所有FN样本的结节直径5mm。于是我在data.py中添加了A.RandomScale(scale_limit0.3, p0.5)随机放大至1.3倍强制模型学习小目标特征。再训练后FN降至0.19且confusion_matrix.png中FN区块明显收缩——这比调10次learning rate更高效。6.3 定制化评估为你的场景添加新指标get_evaluation.py预留了扩展接口。比如工业检测需关注“边缘精度”可添加Hausdorff距离计算# 在get_evaluation.py末尾添加 def hausdorff_distance(mask_true, mask_pred): 计算两掩膜的Hausdorff距离像素 from scipy.spatial.distance import directed_hausdorff coords_true np.argwhere(mask_true) coords_pred np.argwhere(mask_pred) if len(coords_true) 0 or len(coords_pred) 0: return float(inf) d1 directed_hausdorff(coords_true, coords_pred)[0] d2 directed_hausdorff(coords_pred, coords_true)[0] return max(d1, d2) # 在evaluate()函数中调用 hd hausdorff_distance(gt_mask, pred_mask) print(fHausdorff Distance: {hd:.2f} pixels)我的习惯每次新项目启动我必先跑一遍get_evaluation.py盯着confusion_matrix.png看5分钟。如果热力图不对称如FN远大于FP立刻停掉训练先检查数据质量——90%的“模型不行”问题根源在make_mask_data.py生成的mask有偏差。从那以后我每次生成mask都强制本文还有配套的精品资源点击获取
返回列表