Python实现U-Net图像语义分割:从数据加载到部署落地的完整工程实践
简介这是一份面向Python初学者与计算机视觉进阶学习者的图像语义分割实战资源聚焦Unet模型原理与工程实现适用于课程设计、毕业设计、工程实训及AI项目快速启动。资源包含24个文件涵盖4个核心Python训练/预测脚本如unet_train.py、unet_predict.py、14张标注图像png格式及对应XML标签文件、1个已训练好的Keras模型.h5、1份PPTX项目说明文档以及数据生成与结果可视化相关辅助文件整体压缩包达478.98MB结构清晰模块分离明确——含data原始与标签图像、src源码、test测试样本等目录便于理解数据流与模型调用逻辑。已有204人下载学习读者可直接复现端到端流程从数据集生成、模型训练、权重保存到图像预测与可视化同时获得可调试的完整代码框架、典型排错提示及模型部署基础支持。1. 为什么用 Python U-Net 做图像语义分割不是“跑个 demo”就完事你手头有一批工业缺陷图、医疗 CT 切片、或者遥感卫星影像想让模型自动标出“裂纹在哪”“肿瘤边界在哪”“水体/建筑/农田各占多少像素”——这不是目标检测框个 bounding box 就能解决的。它要的是像素级分类每个像素都得打上类别标签。U-Net 正是为此而生它在小样本、高精度、边缘敏感的场景下比 ResNetFCN 或 DeepLab 系列更稳、更准、更易收敛。而 Python 是整个 CV 生态的“事实标准语言”PyTorch/TensorFlow 生态成熟OpenCV 和 scikit-image 处理预处理一步到位tqdm 和 matplotlib 调参可视化不费劲。但现实很骨感很多人 pip install torch torchvision 后直接 copy-paste GitHub 上的 U-Net 代码一跑就 OOM一训就 Dice 系数卡在 0.6 不动验证集 loss 突然爆炸——问题不在模型本身而在数据加载管道没对齐、损失函数没加权重、验证逻辑漏了 mask 归一化。这篇笔记不讲论文复现只讲我在产线部署 3 类 U-Net医学、遥感、工业时从数据准备到部署上线踩过的所有坑以及每一步必须调的 5 个核心参数。适合已经写过 PyTorch DataLoader、能看懂 nn.Conv2d 参数、但还没把语义分割真正跑通落地的工程师。2. 从零搭起可复现的 U-Net 训练框架不是 clone 仓库而是亲手焊死每一根管线U-Net 的结构看似简单编码器-解码器跳跃连接但真正决定效果的是它和数据、训练策略、评估方式咬合的紧密度。我从不直接用torchvision.models.segmentation里的现成模型——它的输入输出接口、预处理逻辑、loss 设计和实际业务数据往往错位。下面这套结构是我在线上项目中稳定迭代 18 个月的最小可行骨架所有模块都可插拔、可 debug、可 profile。2.1 数据加载器别让 DataLoader 成为性能黑洞和标签错位元凶U-Net 对输入尺寸极其敏感原始图像若未做合理裁剪或 padding会导致 batch 内 shape 不一致触发 dynamic paddingGPU 显存暴涨若 mask 标签未与 image 同步做几何变换如旋转、缩放训练时 label 就会“漂移”Dice 永远上不去。我们不用transforms.Compose简单堆叠而是用albumentations实现像素级同步增强它比 torchvision.transforms 更可靠尤其对 maskimport albumentations as A from albumentations.pytorch import ToTensorV2 # 注意所有几何变换必须同时作用于 image 和 mask train_transform A.Compose([ A.Resize(256, 256, interpolationcv2.INTER_NEAREST), # mask 必须用最近邻插值避免灰度值污染 A.HorizontalFlip(p0.5), A.VerticalFlip(p0.5), A.RandomRotate90(p0.5), 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)), # ImageNet 标准化仅对 image ToTensorV2(), # 自动转 tensor 并 permute (H,W,C) - (C,H,W) ]) # 验证/测试阶段不做随机增强但必须 resize normalize to tensor val_transform A.Compose([ A.Resize(256, 256, interpolationcv2.INTER_NEAREST), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ToTensorV2(), ])关键说明interpolationcv2.INTER_NEAREST是给 mask 用的硬性要求双线性插值会让 mask 边界出现 0.3、0.7 这类非整数值后续torch.argmax或F.cross_entropy会报错或学习失效ToTensorV2()会自动将 uint8 图像除以 255.0并转换为 float32无需手动归一化A.Normalize默认只作用于第一个通道即 imagemask 会被跳过——这是 albumentations 的默认行为安全。接着实现自定义 Dataset重点在于__getitem__中严格保证 image/mask 同步加载与变换import cv2 import numpy as np from torch.utils.data import Dataset class SegmentationDataset(Dataset): def __init__(self, image_paths, mask_paths, transformNone): self.image_paths image_paths self.mask_paths mask_paths self.transform transform def __len__(self): return len(self.image_paths) def __getitem__(self, idx): # 读取原图BGR → RGB image cv2.imread(self.image_paths[idx]) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # 读取 mask务必用 cv2.IMREAD_GRAYSCALE否则三通道 mask 会导致类别混淆 mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) # 若 mask 是彩色如 RGB 编码的多类需转换为单通道索引图 # 示例假设 red1, green2, blue3则 # mask np.zeros((h,w), dtypenp.long) # mask[np.where((mask_rgb[:,:,0]255) (mask_rgb[:,:,1]0) (mask_rgb[:,:,2]0))] 1 if self.transform: # 注意传入字典albumentations 会自动对 image 和 mask 同步变换 transformed self.transform(imageimage, maskmask) image transformed[image] mask transformed[mask] # 此时 mask 已是 torch.Tensorshape(H,W)dtypetorch.long若为分类任务 return image, mask参数说明mask必须是单通道uint8或int64且像素值为0, 1, 2, ..., C-1C 为类别数。U-Net 输出 logits 后接nn.CrossEntropyLoss要求 target 为 long 类型若你的 mask 是 0/255 二值图前景/背景需在__getitem__中做mask mask // 255cv2.IMREAD_GRAYSCALE是底线绝不能用PIL.Image.open().convert(L)后者在某些 PNG 透明通道处理上会引入 0–1 之外的灰度值。2.2 U-Net 主干自己写不套现成才能掌控每一层的生死网上大量 U-Net 实现存在致命隐患skip connection 的 channel 数不匹配比如 encoder 第 2 层输出 128 通道decoder 对应层却拼接了 64 通道、上采样方式用nn.Upsample导致 checkerboard artifacts、最后输出未做nn.Softmax或nn.LogSoftmax导致 loss 计算错误。下面是最简但最稳的 PyTorch 实现兼容二类/多类支持 deep supervisionimport torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): Conv → BN → ReLU ×2 def __init__(self, in_ch, out_ch, mid_chNone): super().__init__() if mid_ch is None: mid_ch out_ch self.double_conv nn.Sequential( nn.Conv2d(in_ch, mid_ch, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(mid_ch), nn.ReLU(inplaceTrue), nn.Conv2d(mid_ch, out_ch, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.double_conv(x) class Down(nn.Module): MaxPool → DoubleConv def __init__(self, in_ch, out_ch): super().__init__() self.maxpool_conv nn.Sequential( nn.MaxPool2d(2), DoubleConv(in_ch, out_ch) ) def forward(self, x): return self.maxpool_conv(x) class Up(nn.Module): 上采样 → 拼接 → DoubleConv注意 channel 对齐 def __init__(self, in_ch, out_ch, bilinearTrue): super().__init__() if bilinear: self.up nn.Upsample(scale_factor2, modebilinear, align_cornersTrue) self.conv DoubleConv(in_ch, out_ch, in_ch // 2) # 拼接后 channel in_ch//2 in_ch//2 in_ch else: self.up nn.ConvTranspose2d(in_ch, in_ch // 2, kernel_size2, stride2) self.conv DoubleConv(in_ch, out_ch) def forward(self, x1, x2): x1 self.up(x1) # input is CHW diffY x2.size()[2] - x1.size()[2] diffX x2.size()[3] - x1.size()[3] x1 F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x torch.cat([x2, x1], dim1) # 拼接在 channel 维度 return self.conv(x) class OutConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Conv2d(in_ch, out_ch, kernel_size1) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, n_channels3, n_classes2, bilinearTrue): super(UNet, self).__init__() self.n_channels n_channels self.n_classes n_classes self.bilinear bilinear self.inc DoubleConv(n_channels, 64) self.down1 Down(64, 128) self.down2 Down(128, 256) self.down3 Down(256, 512) factor 2 if bilinear else 1 self.down4 Down(512, 1024 // factor) self.up1 Up(1024, 512 // factor, bilinear) self.up2 Up(512, 256 // factor, bilinear) self.up3 Up(256, 128 // factor, bilinear) self.up4 Up(128, 64, bilinear) self.outc OutConv(64, n_classes) def forward(self, x): x1 self.inc(x) x2 self.down1(x1) x3 self.down2(x2) x4 self.down3(x3) x5 self.down4(x4) x self.up1(x5, x4) x self.up2(x, x3) x self.up3(x, x2) x self.up4(x, x1) logits self.outc(x) return logits关键设计点说明Up模块中F.pad是必须的因Upsample可能导致尺寸非整数倍增长x1和x2H/W 不一致时直接cat会 crashOutConv用kernel_size1是标准做法避免引入额外空间偏差若n_classes1二分类输出 logits 后需接torch.sigmoid若n_classes1则用nn.CrossEntropyLoss内部已含 softmax切勿再手动加 softmax否则梯度爆炸bilinearTrue是默认推荐ConvTranspose2d在实践中容易产生棋盘伪影checkerboard artifacts尤其在遥感/医学边缘区域明显。2.3 训练循环loss、optimizer、scheduler 的黄金配比U-Net 训练极易陷入局部最优尤其当正负样本极度不均衡如肿瘤区域只占图像 0.5%。单纯用nn.CrossEntropyLoss会导向“全预测背景”的懒惰解。必须组合 loss并用 warmup cosine decay 提升稳定性import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR # Loss 函数Dice CrossEntropy 加权实测在遥感和医学场景下 Dice 权重 0.7 最稳 class DiceLoss(nn.Module): def __init__(self, smooth1.): super(DiceLoss, self).__init__() self.smooth smooth def forward(self, logits, target): # logits: (B, C, H, W), target: (B, H, W) —— 注意维度 probs torch.softmax(logits, dim1) # 多类必须 softmax target_onehot F.one_hot(target, num_classeslogits.shape[1]).permute(0,3,1,2).float() intersection (probs * target_onehot).sum(dim(2,3)) union probs.sum(dim(2,3)) target_onehot.sum(dim(2,3)) dice (2. * intersection self.smooth) / (union self.smooth) return 1 - dice.mean() # 返回 scalar loss # 主 lossDice CE权重可调 dice_loss DiceLoss() ce_loss nn.CrossEntropyLoss(ignore_index255) # ignore_index 防止 mask 中无效像素干扰 def combined_loss(logits, targets, alpha0.7): ce ce_loss(logits, targets) dice dice_loss(logits, targets) return alpha * dice (1 - alpha) * ce # OptimizerAdamW 替代 Adamweight_decay 更干净 optimizer optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-3) # Scheduler先 warmup 10 epoch再 cosine decay 到 1e-6 lr_scheduler torch.optim.lr_scheduler.SequentialLR( optimizer, schedulers[ LinearLR(optimizer, start_factor0.01, total_iters10), CosineAnnealingLR(optimizer, T_maxepochs-10, eta_min1e-6) ], milestones[10] )参数说明alpha0.7是血泪经验Dice 关注重叠率CE 关注分类置信度过高 Dice 权重会导致边缘模糊ignore_index255是关键若 mask 中有未标注区域如遥感图中的云遮挡区将其 pixel value 设为 255并在 CE loss 中忽略否则会拖垮梯度AdamW比Adam更适合视觉任务weight_decay1e-3能有效抑制过拟合尤其在小数据集上LinearLR CosineAnnealingLR组合比单一 StepLR 收敛更快、最终精度更高已在 7 个不同数据集上验证。3. 验证与推理别让“训练 loss 下降”骗了你真功夫在 mask 后处理里训练脚本跑出 0.92 的 train Dice验证集却只有 0.78这太常见了。问题往往不出在模型而出在验证逻辑本身没做 thresholding、没做 CRF 后处理、没按原始分辨率还原、甚至没把 logits 正确转成 label map。下面给出一套端到端验证 pipeline确保你看到的指标就是线上真实效果。3.1 验证时的 mask 生成logits → prob → argmax → resize → save验证阶段必须关闭 dropout/batch norm 的 training mode并用torch.no_grad()节省显存model.eval() val_metrics {dice: [], iou: []} with torch.no_grad(): for images, masks in val_loader: images images.to(device) masks masks.to(device) # shape: (B, H, W) logits model(images) # shape: (B, C, H, W) # 多类softmax argmax二类sigmoid 0.5 if n_classes 1: probs torch.sigmoid(logits) # (B, 1, H, W) preds (probs 0.5).long().squeeze(1) # (B, H, W) else: probs torch.softmax(logits, dim1) # (B, C, H, W) preds torch.argmax(probs, dim1) # (B, H, W) # 计算 batch-level metric注意preds 和 masks 必须同 dtype device for i in range(len(images)): pred_mask preds[i].cpu().numpy() true_mask masks[i].cpu().numpy() val_metrics[dice].append(dice_coeff(pred_mask, true_mask)) val_metrics[iou].append(iou_coeff(pred_mask, true_mask)) # 自定义 metric 函数避免 torchmetrics 依赖 def dice_coeff(pred, target, eps1e-6): intersection (pred target).sum() union pred.sum() target.sum() return (2. * intersection eps) / (union eps) def iou_coeff(pred, target, eps1e-6): intersection (pred target).sum() union pred.sum() target.sum() - intersection return (intersection eps) / (union eps)关键说明preds torch.argmax(probs, dim1)是多类语义分割的唯一正确路径probs.max(dim1)返回的是 (value, index)index 才是 labeldice_coeff和iou_coeff必须用 numpy 逐样本计算不能用 batch-level mean否则小目标被大背景淹没eps1e-6是防除零不是玄学——当某张图无目标时分母为 0不加 eps 会返回 nan破坏整个 epoch 指标。3.2 推理时的工程化输出不只是 .png而是可交付的 mask overlay stats线上部署不要只输出一张白底黑字的 mask.png。客户需要的是带原始比例的彩色分割图、叠加原图的可视化图、每个类别的像素统计、以及可导入 GIS 或 CAD 的 GeoJSON遥感或 DICOM ROI医疗。以下是一个生产级 inference 函数def predict_and_save(model, image_path, output_dir, device, n_classes2, paletteNone): 输入单张图像路径 输出mask.png, overlay.jpg, stats.json model.eval() image cv2.imread(image_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # 预处理resize 到模型输入尺寸如 256x256normalize transform A.Compose([ A.Resize(256, 256, interpolationcv2.INTER_NEAREST), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ToTensorV2() ]) transformed transform(imageimage) tensor_img transformed[image].unsqueeze(0).to(device) # (1,3,H,W) with torch.no_grad(): logits model(tensor_img) if n_classes 1: pred (torch.sigmoid(logits) 0.5).long().squeeze(0).squeeze(0).cpu().numpy() else: pred torch.argmax(torch.softmax(logits, dim1), dim1).squeeze(0).cpu().numpy() # 还原到原始尺寸双线性插值mask 用 nearest orig_h, orig_w image.shape[:2] pred_resized cv2.resize(pred, (orig_w, orig_h), interpolationcv2.INTER_NEAREST) # 保存 maskuint80~C-1 mask_path os.path.join(output_dir, f{os.path.basename(image_path).split(.)[0]}_mask.png) cv2.imwrite(mask_path, pred_resized.astype(np.uint8)) # 生成 overlay用 palette 上色并叠加 if palette is None: palette [[0,0,0], [255,0,0], [0,255,0], [0,0,255]] # 背景、类1、类2、类3... overlay image.copy() for cls_id in range(1, n_classes): # 跳过背景 mask_binary (pred_resized cls_id) overlay[mask_binary] (overlay[mask_binary] * 0.5 np.array(palette[cls_id]) * 0.5).astype(np.uint8) overlay_path os.path.join(output_dir, f{os.path.basename(image_path).split(.)[0]}_overlay.jpg) cv2.imwrite(overlay_path, cv2.cvtColor(overlay, cv2.COLOR_RGB2BGR)) # 统计像素占比 stats {} for cls_id in range(n_classes): count (pred_resized cls_id).sum() stats[fclass_{cls_id}] { pixel_count: int(count), percentage: round(count / (orig_h * orig_w) * 100, 2) } stats_path os.path.join(output_dir, f{os.path.basename(image_path).split(.)[0]}_stats.json) with open(stats_path, w) as f: json.dump(stats, f, indent2) print(f✅ Saved: {mask_path}, {overlay_path}, {stats_path}) return pred_resized, stats参数说明palette是可配置的医疗常用 [0,0,0]背景、[255,0,0]肿瘤、[0,255,0]正常组织遥感常用 [0,0,0]背景、[0,0,255]水体、[255,255,0]建筑cv2.INTER_NEAREST用于 mask resizecv2.INTER_LINEAR用于 overlay二者不可互换stats.json是交付刚需客户不关心 Dice只关心“这片农田面积是多少亩”像素数 × 地面分辨率即可换算。4. 避坑U-Net 训练中最常翻车的 5 个现场附定位命令和修复命令这些坑我都在凌晨三点的服务器上亲手 debug 过不是理论推演是血泪经验。每一条都对应一个真实报错、一句print()定位命令、一行修复代码。4.1 现象RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation原因nn.ReLU(inplaceTrue)在某些 PyTorch 版本1.12与torch.compile或DDP冲突或F.pad时用了inplaceTrue但F.pad无此参数实为误读最常见是torch.sigmoid_()或x y这类 inplace 操作。定位在 loss.backward() 前加torch.autograd.set_detect_anomaly(True)运行后报错行会指向具体操作。解决全局搜索_结尾的函数如sigmoid_,relu_,copy_全部替换为非 inplace 版本x y改为x x ynn.ReLU(inplaceTrue)改为nn.ReLU(inplaceFalse)。4.2 现象训练 loss 为 nan或验证 Dice 突然跳到 0.0原因mask 中存在非法像素值如 -1、256、nan导致CrossEntropyLoss计算 log(0)或softmax输入过大logits 中有 infsoftmax输出 nan。定位在 dataloader 返回前加断点print(mask unique:, torch.unique(masks)) print(mask min/max:, masks.min().item(), masks.max().item()) print(mask nan count:, torch.isnan(masks).sum().item())解决在__getitem__中强制 clipmask np.clip(mask, 0, n_classes-1) # 确保所有值在 [0, C-1] mask torch.from_numpy(mask).long()4.3 现象验证 Dice 持续 0.0但 train loss 下降原因验证时用了model.train()模式BN 层用 batch statistics 而非 running mean/var导致输出不稳定或torch.no_grad()没加显存溢出触发 OOM 后部分 batch 被 skipmetric 计算失真。定位验证 loop 开头加print(model.training)应为False用nvidia-smi观察显存是否周期性 spike。解决严格保证model.eval()torch.no_grad()成对出现验证 batch size 设为 train 的 1/2留足显存余量。4.4 现象预测 mask 全是 0纯黑或全是 1纯白原因nn.CrossEntropyLoss的 target 必须是torch.long若误传torch.float32loss 会静默失败不报错但梯度为 0或sigmoid后 threshold 设为 0.99太严所有像素被判为背景。定位检查 loss 值是否恒为常数如始终 0.6931即 -log(0.5)打印masks.dtype和masks.max()。解决masks masks.long()强制转换threshold 从 0.5 开始调用plt.hist(probs.flatten())看概率分布再定。4.5 现象训练速度极慢1 iter/secGPU 利用率 10%原因DataLoader的num_workers0时Windows 下cv2.imread与多进程冲突OpenCV 默认使用 forkWindows 用 spawn或albumentations的p参数设为 0.999每次都要做 heavy transform。定位num_workers0运行若速度恢复即为多进程问题用htop观察 CPU 是否满载。解决Windows 用户必须设num_workers0Linux 用户设num_workersmin(8, os.cpu_count())并在__getitem__开头加cv2.setNumThreads(0)关闭 OpenCV 多线程。5. 进阶技巧如何让 U-Net 在小样本100 张下依然达到 0.85 Dice小样本是工业检测、罕见病诊断、新遥感地物识别的常态。U-Net 天然适合小样本但需三招组合拳强数据增强 伪标签自训练 损失函数重加权。下面给出可直接粘贴的完整 pipeline。5.1 强增强用imgaug补充albumentations未覆盖的域迁移操作albumentations擅长几何光度但对模拟传感器噪声、运动模糊、镜头畸变支持弱。imgaug更底层可精准控制import imgaug.augmenters as iaa # 针对工业缺陷图模拟产线相机抖动、LED 光源闪烁、CMOS 热噪声 iaa_seq iaa.Sequential([ iaa.MotionBlur(k3, angle[-45, 45]), # 模拟轻微抖动 iaa.AdditiveGaussianNoise(loc0, scale(0.0, 0.05*255), per_channelTrue), # CMOS 噪声 iaa.JpegCompression(compression(70, 95)), # 模拟传输压缩 iaa.Sometimes(0.3, iaa.CoarseDropout(0.02, size_percent0.3)) # 模拟 sensor dead pixel ], random_orderTrue) # 在 Dataset.__getitem__ 中调用注意iaa 输入是 uint8 numpy输出也是 image_aug iaa_seq(imageimage) # image 是 uint8 RGB numpy array为什么有效在轴承裂纹数据集87 张上加iaa_seq后验证 Dice 从 0.72 → 0.81因为模型学会了忽略噪声模式专注裂纹几何特征。5.2 伪标签用初始模型给无标注图打标迭代提升伪标签不是“猜”而是用置信度阈值筛选高质量预测。关键在只选probs.max(dim1).values 0.95的像素且该像素所在连通域面积 100pxdef generate_pseudo_labels(model, unlabeled_images, threshold0.95, min_area100): model.eval() pseudo_masks [] for img_path in unlabeled_images: image cv2.imread(img_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # ... 预处理同 train ... with torch.no_grad(): logits model(tensor_img) probs torch.softmax(logits, dim1) conf, preds probs.max(dim1) # conf: (1,H,W), preds: (1,H,W) # 置信度过滤 连通域过滤 mask_np preds.squeeze(0).cpu().numpy() conf_np conf.squeeze(0).cpu().numpy() # 二值掩膜高置信 非背景类 high_conf (conf_np threshold) (mask_np 0) labeled_mask np.zeros_like(mask_np) # 对每个类别分别做连通域分析 for cls_id in np.unique(mask_np[high_conf]): if cls_id 0: continue cls_mask (mask_np cls_id) high_conf num_labels, labels cv2.connectedComponents(cls_mask.astype(np.uint8)) for i in range(1, num_labels): area (labels i).sum() if area min_area: labeled_mask[labels i] cls_id pseudo_masks.append(labeled_mask) return pseudo_masks # 使用将 pseudo_masks 加入训练集权重设为 0.3降低噪声影响参数说明threshold0.95是底线低于此值伪标签噪声太大min_area100防止单个噪点被误标伪标签样本 loss 权重0.3是经验值太高会污染梯度。5.3 损失重加权用有效像素占比动态调整 class weightnn.CrossEntropyLoss(weight...)的 weight 是静态的但每张图的前景占比不同。我们用 batch 内动态 weightdef dynamic_weighted_ce(logits, targets, eps1e-6): # logits: (B,C,H,W), targets: (B,H,W) B, C, H, W logits.shape targets_onehot F.one_hot(targets, num_classesC).permute(0,3,1,2).float() # (B,C,H,W) # 计算每个类在 batch 内的有效像素数忽略 ignore_index valid_mask (targets ! 255).float() # (B,H,W) class_counts (targets_onehot * valid_mask.unsqueeze(1)).sum(dim(2,3)) # (B,C) total_valid valid_mask.sum(dim(1,2)) # (B,) # 动态 weight 1 / (class_count eps) 归一化到 sum1 weights 1.0 / (class_counts eps) # (B,C) weights weights / weights.sum(dim1, keepdimTrue) # (B,C) # 按 batch 元素加权 CE log_probs F.log_softmax(logits, dim1) # (B,C,H,W) ce -(targets_onehot * log_probs).sum(dim1) # (B,H,W) weighted_ce (ce * valid_mask).sum(dim(1,2)) # (B,) weighted_ce (weighted_ce * weights.sum(dim1)).mean() # 加权平均 return weighted_ce为什么比 static weight 强在 PCB 缺陷数据集缺陷像素仅占 0.3%上dynamic weight 使缺陷类 Dice 提升 12.7%因为每张图都根据其真实分布校准了梯度强度。最后说句实在话U-Net 不是银弹但它是最可靠的“基线模型”。我见过太多团队花三个月调 Transformer 分割模型本文还有配套的精品资源点击获取

相关新闻

知识图谱+图神经网络:电影推荐系统完整实现与避坑指南

知识图谱+图神经网络:电影推荐系统完整实现与避坑指南

简介:一份基于Python的知识图谱与图神经网络(KGCN)电影推荐系统毕业设计项目,面向计算机相关专业学生及推荐系统入门开发者,适用于课程设计、毕业设计或项目实战。资源共31个文件,包含21个Python脚本、5个d…

2026/10/11 23:49:54 阅读更多 →
Understand-Anything大型Monorepo基准测试体系详解:如何获得可复现的性能数据

Understand-Anything大型Monorepo基准测试体系详解:如何获得可复现的性能数据

Understand-Anything大型Monorepo基准测试体系详解:如何获得可复现的性能数据 【免费下载链接】Understand-Anything Graphs that teach > graphs that impress. Turn any code into an interactive knowledge graph you can explore, search, and ask questions…

2026/10/11 23:49:54 阅读更多 →
ComfyUI+Flux三视角OOTD实战:从单图到多视角协同建模

ComfyUI+Flux三视角OOTD实战:从单图到多视角协同建模

简介:本资源是面向ComfyUI图像生成工作流开发者的轻量级OOTD(Outfit Our Diffusion)三视角自定义模特配置方案,适用于服装试穿、虚拟穿搭等AIGC应用场景,适合具备基础ComfyUI节点逻辑理解能力的中级开发者快速复用。压…

2026/10/11 23:48:53 阅读更多 →

最新新闻

VMware NAT模式虚拟机端口转发到宿主机:两种配置方法与排查指南

VMware NAT模式虚拟机端口转发到宿主机:两种配置方法与排查指南

1. 先把这张网络拓扑图看明白:这个需求到底在解决什么问题先说结论:你要做的事情,就是把一台运行在VMware NAT网络里的虚拟机,它的9980端口“搬”到宿主机上,让局域网里其他电脑通过访问宿主机IP就能用上这个端口背后的…

2026/10/12 0:34:17 阅读更多 →
技术社区高效求助:正确提问与自我排查,让大佬愿意回复你的帖子

技术社区高效求助:正确提问与自我排查,让大佬愿意回复你的帖子

“求助各位大佬”这五个字,几乎是所有技术社区、交流群里最常出现的标题,也是最容易沉底、最容易让人划过的一条。我刚入行那几年,没少发过这样的帖子,也没少干过把问题描述得云里雾里然后干等三天无人问津的事。后来自己技术慢慢…

2026/10/12 0:34:17 阅读更多 →
烟火检测数据集与YOLO训练实战:从标注格式到模型部署

烟火检测数据集与YOLO训练实战:从标注格式到模型部署

简介:烟火检测数据集面向目标检测与YOLO系列模型训练,包含一千张真实场景图像及对应XML标注,适用于烟火识别、安全监控等视觉任务,也适合作为目标检测入门学习的练习数据。压缩包整体约89.87MB,共两千个文件&#xff0…

2026/10/12 0:34:16 阅读更多 →
非小米电脑安装小米电脑管家教程:跨设备互联与踩坑指南

非小米电脑安装小米电脑管家教程:跨设备互联与踩坑指南

最近一段时间的装机圈子里,“小米电脑管家”这五个字的讨论热度一直没降下来过。原因很简单:小米官方生态里那套跨设备互联体验,确实做得够顺滑,但官方说明一直写着仅限小米自家笔记本使用。可这几个月网上的玩法已经变了&#xf…

2026/10/12 0:34:16 阅读更多 →
把 Claude Code 的默认模型切成 GLM-5.3 Flash:CC Switch 换模型全流程,30 分钟上手

把 Claude Code 的默认模型切成 GLM-5.3 Flash:CC Switch 换模型全流程,30 分钟上手

把 Claude Code 的默认模型切成 GLM-5.3 Flash:CC Switch 换模型全流程,30 分钟上手 【免费下载链接】GLM-5.3 GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。 …

2026/10/12 0:34:16 阅读更多 →
智慧能源管理平台如何让光伏电站从监控走向高效管控

智慧能源管理平台如何让光伏电站从监控走向高效管控

我最早在一线跑光伏电站的时候,对“智慧能源管理平台”这六个字是有怀疑的。当时装了远程监控、能看到实时功率和发电量,我就觉得电站已经管起来了。后来巡检次数多了才发现,监控大屏上的曲线往往一片祥和,可实际发电量却在悄悄缩…

2026/10/12 0:33:16 阅读更多 →

日新闻

复古胶片颗粒感噪点合成器:Canvas ImageData 像素高斯杂色注入算法

复古胶片颗粒感噪点合成器:Canvas ImageData 像素高斯杂色注入算法

在数码相机、高清显示屏与现代矢量图形技术高度发达的今天,画面可以做到绝对的锐利、平滑与无瑕。然而,当一张秋日手账插画或拍立得照片过于“平整无瑕”时,往往会散发出一种冰冷生硬的“数码塑料感(Digital Plasticity&#xff0…

2026/10/12 0:00:59 阅读更多 →
活字印刷古籍线装排版:Canvas 竖排文字与栏线自适应算法

活字印刷古籍线装排版:Canvas 竖排文字与栏线自适应算法

在现代网页与移动端设计中,横排(Horizontal Layout)早已经成为了绝对的主流。然而,当我们翻开泛黄的线装古籍、宋版木刻诗集,或是欣赏一张茶道雅集的手写便签时,那种**自上而下纵向书写、自右向左逐列铺展&…

2026/10/12 0:00:59 阅读更多 →
周日晚间的“精神松绑减震器”:无压力情绪倾倒箱与温和轻声陪伴

周日晚间的“精神松绑减震器”:无压力情绪倾倒箱与温和轻声陪伴

每到周日的晚上八点到十点,很多人心里都会悄悄亮起一盏警示灯。 在心理学上,这种现象有一个专门的称谓——“周日夜晚焦虑症(Sunday Scaries)”。明天又是周一,闹钟又要重新在七点响彻卧房;脑海里仿佛有一个…

2026/10/12 0:00:59 阅读更多 →

周新闻

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

流感时间序列预测实战:ARIMA/LSTM全流程拆解与避坑指南

简介:基于 ARIMA、LSTM、Transformer 等模型的流感时间序列预测 Python 源码,面向计算机相关专业课程设计与期末大作业学生,以及项目实战学习者。内容覆盖预处理、平稳性检验、定阶、残差分析、多模型对比预测的完整时序建模流程,…

2026/10/12 0:16:30 阅读更多 →
影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别

影刀RPA新手教程:键盘模拟输入实战——输入文本与模拟按键的区别 做影刀RPA自动化,十个新手有八个栽在"往输入框里填东西"这件事上:要么填不进去,要么填了一半,要么直接把原来内容追加在后面。这背后的根因&…

2026/10/12 0:16:38 阅读更多 →
影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容

影刀RPA新手教程:阅文起点小说数据采集实战——书籍信息与章节内容 1. 认识影刀:什么场景该用RPA采小说数据 起点中文网的页面结构相对稳定——分类榜单、书籍详情、章节内容三块独立页面,跳转链路清晰。这种场景非常适合影刀自动化&#x…

2026/10/12 0:16:43 阅读更多 →

月新闻

我发现了一个新思路:用 Remotion + Claude Code 像写代码一样自动化生成短视频

我发现了一个新思路:用 Remotion + Claude Code 像写代码一样自动化生成短视频

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/11 10:45:37 阅读更多 →
Windows下 Codex 中 Chrome 和 Computer Use 插件不可用问题排查及解决参考方式:TaoToken 统一 Key 配置与验证

Windows下 Codex 中 Chrome 和 Computer Use 插件不可用问题排查及解决参考方式:TaoToken 统一 Key 配置与验证

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/11 14:36:53 阅读更多 →
黑夜航拍船只数据集训练YOLOV5模型全流程解析

黑夜航拍船只数据集训练YOLOV5模型全流程解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/10/11 14:36:54 阅读更多 →