ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

DeepLabV3+实战:水面漂浮物语义分割与报警

DeepLabV3+实战:水面漂浮物语义分割与报警 简介这是一份基于开源模型DeepLabV3完成的水体漂浮物像素级分割项目来源于模式识别与机器学习课程小组结课作业同时也应用于极市开发者平台打榜实践。项目以Python为主要开发语言围绕水体与漂浮物分割任务设计并实现了基于面积阈值的自动报警判断流程适合正在学习语义分割、目标检测或参加相关竞赛的本科生、研究生与开发者参考。压缩包共包含258个文件整体约29.77MB其中样本图像有120张PNG和100张JPG覆盖河流、海域等不同水上环境另有26个Python脚本负责数据预处理、模型训练、推理与结果分析7个文本说明用于环境配置与使用指引2个shell脚本可辅助快速运行。资源当前已有223人浏览学习。通过本包可以获得完整的项目文件组织方式既能理解DeepLabV3在真实场景中的标注数据、训练流程与预测逻辑也可基于现有代码修改适配其他分割需求是课程设计与竞赛实战的实用参考。1. 拿课程设计和打榜赛练 DeepLabV3这个项目到底在做什么先抛结论这可能是你见过最典型的“课程设计 平台打榜”二合一项目——用开源模型 DeepLabV3 对水面监控视频做语义分割把水体和水上漂浮物逐像素分出来再按面积阈值判断是否需要报警。数据来自极市开发者平台的漂浮物检测赛道图片命名里的 rivers、sea 已经标明了场景训练集是抽帧后的水面图片标签是像素级的分割掩码。整个任务的技术栈就是 Python PyTorch 图像分割没有复杂的工程框架适合模式识别与机器学习课程的小组结课也适合想在实际分割任务上验证模型能力的入门者。这类项目的价值不在“跑通一个模型”——那只是及格线。真正拉开差距的地方有三个数据怎么切分和清洗、训练参数怎么匹配数据集特性、推理后处理里的报警阈值怎么定才不误报不漏报。这三件事做扎实了答辩有东西讲打榜分数也不会难看。后面的篇幅就按这个顺序先把数据准备讲透再给训练配置和推理报警的完整代码最后把最容易翻车的地方单独列出来。2. 数据准备水面漂浮物数据集的切分与标注处理2.1 先读懂数据集的命名和目录结构拿到的数据集压缩包解压后训练图片基本是这样的命名格式ZDSfloating_objects20230206_V3_train_rivers_27_000044.jpg、ZDSfloating_objects20230206_V3_train_sea_1_011323.jpg。拆开看20230206是数据采集日期版本V3是标注版本rivers和sea是场景类别后面的数字是视频序号和帧号。这个命名规律在写数据加载器时非常有用——按场景字段做分层采样就可以了。初次拿到数据第一件事不是写模型而是把目录结构和标注文件彻底摸清楚。我一般会先执行一段检查代码输出所有图片和对应的标注掩码是否一一对应以及每张图的尺寸是否统一import os import cv2 from collections import defaultdict img_dir ./data/images mask_dir ./data/masks img_files sorted(os.listdir(img_dir)) mask_files sorted(os.listdir(mask_dir)) print(f图片数量: {len(img_files)}, 掩码数量: {len(mask_files)}) # 检查文件名是否成对出现 img_names {f.split(.)[0] for f in img_files} mask_names {f.split(.)[0] for f in mask_files} print(f缺失掩码的文件: {len(img_names - mask_names)}) print(f多余掩码的文件: {len(mask_names - img_names)}) # 统计图片尺寸分布 size_counter defaultdict(int) for f in img_files[:500]: img cv2.imread(os.path.join(img_dir, f)) size_counter[img.shape[:2]] 1 print(f前500张图的尺寸分布: {dict(size_counter)})这段代码的逻辑并不复杂先把图片和掩码文件名都提出来做集合差确认没有“有图没标注”或“有标注没图”的情况然后抽样统计尺寸分布。水面监控视频抽帧出来的图尺寸通常是一致的比如 1920x1080 或 1280x720但如果存在多分辨率混用的情况训练前就必须统一处理不然后面 DataLoader 的 collate 会直接报错。提示如果发现掩码是 PNG 格式而图片是 JPG不要惊讶这是分割任务的常规做法。JPG 有损压缩会污染类别边缘PNG 无损才能保证像素级标注的精度。2.2 标签规则背景、水体、漂浮物三分类的类目设计水面漂浮物分割的类别设计最常用的是三分类0 表示背景天空、岸边、船体等1 表示水体2 表示漂浮物。有些版本会把水体当作背景直接做二分类但那样的话漂浮物边缘和水面反光的区分度会明显变差——DeepLabV3 的 ASPP 模块能利用多尺度上下文信息把水体单独作为一类能帮助模型学到“漂浮物是浮在水面上的物体”这个空间关系。标注掩码的图像模式需要确认一下灰度图还是调色板图这会直接影响读取方式from PIL import Image import numpy as np mask Image.open(./data/masks/ZDSfloating_objects20230206_V3_train_sea_1_011323.png) print(图像模式:, mask.mode) print(尺寸:, mask.size) mask_array np.array(mask) print(像素值范围:, np.unique(mask_array)) print(类别统计:, {cls: (mask_array cls).sum() for cls in np.unique(mask_array)})如果mask.mode是P说明是调色板索引图像素值本身就是类别 ID直接用np.array(mask)得到的就是类别矩阵读取后要检查类别 ID 是否连续比如 0、1、2如果mask.mode是RGB说明标注被存成了彩图需要先做颜色到类别 ID 的映射表类别统计这里特别值得看一眼漂浮物类别的像素占比通常非常低可能不到 1%这就是典型的类别不平衡问题我实际处理时会把统计结果做成一张表格放在项目 README 里答辩的时候直接贴出来说明你做过数据探索分析。这一点在课程评分里往往比模型结构讲解更加分。2.3 按场景切分训练集和验证集而不是随机切这里有一个很容易被忽视的细节切分数据集时不要用random_split全量随机切要按场景rivers、sea分层切。原因很直接——河流和海洋的水面纹理、光照条件、漂浮物类型差异很大如果随机切很可能验证集里某个场景占比过高或过低导致验证指标失真。import os import random import shutil random.seed(42) image_dir ./data/images mask_dir ./data/masks train_img_dir ./data/train/images val_img_dir ./data/val/images train_mask_dir ./data/train/masks val_mask_dir ./data/val/masks os.makedirs(train_img_dir, exist_okTrue) os.makedirs(val_img_dir, exist_okTrue) os.makedirs(train_mask_dir, exist_okTrue) os.makedirs(val_mask_dir, exist_okTrue) images sorted(os.listdir(image_dir)) # 按场景分组 scene_groups {} for img_name in images: scene rivers if rivers in img_name else sea scene_groups.setdefault(scene, []).append(img_name) val_ratio 0.15 for scene, scene_images in scene_groups.items(): random.shuffle(scene_images) val_count int(len(scene_images) * val_ratio) val_set set(scene_images[:val_count]) for img_name in scene_images: src_img os.path.join(image_dir, img_name) src_mask os.path.join(mask_dir, img_name.replace(.jpg, .png)) if img_name in val_set: shutil.copy(src_img, os.path.join(val_img_dir, img_name)) shutil.copy(src_mask, os.path.join(val_mask_dir, img_name.replace(.jpg, .png))) else: shutil.copy(src_img, os.path.join(train_img_dir, img_name)) shutil.copy(src_mask, os.path.join(train_mask_dir, img_name.replace(.jpg, .png))) print(切分完成验证集比例: 15%)这段代码的关键在于先按rivers和sea分组再在组内各自按 15% 的比例抽验证集。这样无论哪个场景验证集都有足够且均衡的样本量。另外一个附带的好处是如果后续你想做场景对比实验比如只看河流场景的 IoU数据已经按场景分好了直接用文件名前缀过滤即可。3. 训练配置DeepLabV3 的骨干网络选择与损失函数设计3.1 backbone 选型ResNet50、ResNet101 与 MobileNetV2 的取舍DeepLabV3 的框架结构可以拆成两部分理解Encoder 部分用带空洞卷积的骨干网络提取特征再通过 ASPPAtrous Spatial Pyramid Pooling模块用不同膨胀率的空洞卷积并行捕捉多尺度上下文Decoder 部分把深层特征与浅层特征融合恢复物体边缘细节。这个“多尺度空洞卷积 编解码结构”的组合特别适合水面漂浮物这种目标尺寸差异极大的任务——小到饮料瓶大到成片的水葫芦到了 512x512 输入尺度下目标尺寸可能差出几十倍。骨干网络的选择直接决定显存占用和训练速度。我整理了一个对比表课程设计的机器配置下怎么选一目了然骨干网络参数量输入尺寸单卡显存占用batch8建议场景ResNet50约 26M512x512约 6-8 GB有 1080Ti 以上显卡兼顾速度与精度ResNet101约 47M512x512约 10-12 GB追求精度上限显存充足MobileNetV2约 3M512x512约 2-3 GB显卡吃紧或需要快速迭代调参这里的经验是第一次跑通流程用 MobileNetV2验证整套代码逻辑没问题再切到 ResNet50 做正式训练。不要一上来就用 ResNet101——训练时间翻倍而课程设计的数据量下ResNet50 和 ResNet101 的最终分数差距往往只有一两个百分点不值得。3.2 损失函数交叉熵为主Dice Loss 辅助解决类别不平衡前面统计过漂浮物类别的像素占比经常不到 1%直接用标准交叉熵损失训练模型会严重偏向背景类和水体类漂浮物类别几乎学不出来。常见做法是用交叉熵和 Dice Loss 的加权组合import torch import torch.nn as nn import torch.nn.functional as F class CombinedLoss(nn.Module): def __init__(self, num_classes3, weightNone, dice_weight0.4): super().__init__() self.ce nn.CrossEntropyLoss(weightweight) self.dice_weight dice_weight self.num_classes num_classes def dice_loss(self, logits, targets): probs F.softmax(logits, dim1) # [B, C, H, W] targets_onehot F.one_hot(targets, num_classesself.num_classes) targets_onehot targets_onehot.permute(0, 3, 1, 2).float() # [B, C, H, W] smooth 1.0 intersection (probs * targets_onehot).sum(dim(0, 2, 3)) union probs.sum(dim(0, 2, 3)) targets_onehot.sum(dim(0, 2, 3)) dice (2.0 * intersection smooth) / (union smooth) # 对背景类权重降低重点提升漂浮物类 class_weights torch.tensor([0.2, 0.3, 1.0], devicelogits.device) dice_loss_value 1.0 - (dice * class_weights).mean() return dice_loss_value def forward(self, logits, targets): ce_loss self.ce(logits, targets) dice_loss_value self.dice_loss(logits, targets) return ce_loss self.dice_weight * dice_loss_value这个复合损失有两个设计意图交叉熵负责稳定的梯度信号保证收敛Dice Loss 对类别不平衡天然不敏感因为它计算的是预测区域和真实区域的像素级重叠率。特别留意class_weights的设定——给漂浮物类别更高的权重让损失函数明确知道“分错漂浮物的代价比分错背景更大”。参数dice_weight0.4是经验值太大超过 0.6会导致训练初期损失震荡太小低于 0.2又起不到平衡作用。3.3 超参数设置与训练脚本骨架训练的超参数设置看起来常规但每个值都值得细说。输入尺寸 512x512 是 DeepLabV3 官方预训练的标准尺寸直接沿用可以加载 ImageNet 预训练权重不用修正主干网络的全连接层尺寸batch size 取 8 是在 ResNet50 11GB 显存下的安全值初始学习率 1e-4 配 Poly 学习率衰减策略是分割任务里最稳的搭配——线性衰减到 0比 StepLR 的阶梯下降更适合语义分割这种需要长时间精细收敛的任务。import torch from torch.utils.data import DataLoader from torchvision import models import segmentation_models_pytorch as smp # 使用 segmentation_models_pytorch 库直接实例化 DeepLabV3 model smp.DeepLabV3Plus( encoder_nameresnet50, # 骨干网络 encoder_weightsimagenet, # 加载 ImageNet 预训练权重 in_channels3, classes3, # 背景、水体、漂浮物 ) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) # Poly 学习率衰减 def poly_lr(base_lr, iter, max_iter, power0.9): return base_lr * (1 - iter / max_iter) ** power optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) criterion CombinedLoss(num_classes3, dice_weight0.4) # 随机翻转 随机裁剪做数据增强 transform torch.nn.Sequential( # 在实际代码中用 albumentations 实现 ) train_loader DataLoader(train_dataset, batch_size8, shuffleTrue, num_workers4, drop_lastTrue) max_epochs 50 iters_per_epoch len(train_loader) global_step 0 for epoch in range(max_epochs): model.train() for images, masks in train_loader: images images.to(device) masks masks.to(device).long() lr poly_lr(1e-4, global_step, max_epochs * iters_per_epoch) for g in optimizer.param_groups: g[lr] lr optimizer.zero_grad() logits model(images) loss criterion(logits, masks) loss.backward() optimizer.step() global_step 1 if (epoch 1) % 10 0: torch.save(model.state_dict(), f./checkpoints/deeplabv3p_epoch{epoch1}.pth) print(fEpoch {epoch1}/{max_epochs}, Loss: {loss.item():.4f}, LR: {lr:.6f})关于预训练权重的加载多说一句encoder_weightsimagenet这个参数如果你用的是 TorchVision 自带的 DeepLabV3Plus 实现需要确认权重版本与 PyTorch 版本匹配。更省心的方式是直接用segmentation_models_pytorch这个第三方库它的预训练权重管理做得更友好ResNet50 的 encoder 能自动下载对应权重不用自己处理状态字典的键名映射问题。提示训练到第 30 轮左右如果损失还在明显下降属于正常现象。分割任务收敛本来就慢建议至少训练 50 轮再下结论。4. 推理与报警从分割掩码到面积阈值判定4.1 推理管线加载模型、输出掩码、连通域分析模型训练完成后推理阶段要做的事情和训练完全不同训练时关注 Loss 下降推理时关注的是能否从分割掩码中准确提取漂浮物区域并给出报警决策。完整的推理流程是——读图、预处理、模型前向推理、argmax取类别、连通域分析、面积计算、阈值判断。import cv2 import numpy as np import torch import segmentation_models_pytorch as smp def postprocess_mask(prob_map, min_area500): prob_map: [3, H, W] 的 softmax 输出 min_area: 最小连通域面积阈值像素数 # 取每个像素概率最大的类别 mask np.argmax(prob_map, axis0) # [H, W] # 只保留漂浮物类别类别 ID 2 floating_mask (mask 2).astype(np.uint8) * 255 # 形态学开运算先腐蚀再膨胀去除孤立噪点 kernel cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3)) floating_mask cv2.morphologyEx(floating_mask, cv2.MORPH_OPEN, kernel) # 连通域分析 num_labels, labels, stats, _ cv2.connectedComponentsWithStats(floating_mask, connectivity8) alert_boxes [] total_floating_pixels 0 for i in range(1, num_labels): # 跳过背景标签 0 area stats[i, cv2.CC_STAT_AREA] if area min_area: x, y, w, h (stats[i, cv2.CC_STAT_LEFT], stats[i, cv2.CC_STAT_TOP], stats[i, cv2.CC_STAT_WIDTH], stats[i, cv2.CC_STAT_HEIGHT]) alert_boxes.append((x, y, w, h, area)) total_floating_pixels area return mask, floating_mask, alert_boxes, total_floating_pixels def inference_single_image(model, image_path, device, min_area500): model.eval() image cv2.imread(image_path) image_rgb cv2.cvtColor(image, cv2.COLOR_BGR2RGB) h, w image_rgb.shape[:2] # 缩放到 512x512 并归一化 input_tensor cv2.resize(image_rgb, (512, 512)) input_tensor input_tensor.astype(np.float32) / 255.0 input_tensor torch.from_numpy(input_tensor).permute(2, 0, 1).unsqueeze(0).to(device) with torch.no_grad(): output model(input_tensor) # [1, 3, 512, 512] prob torch.softmax(output, dim1).squeeze(0).cpu().numpy() # 还原到原始分辨率 prob_resized np.transpose(prob, (1, 2, 0)) # [512, 512, 3] prob_resized cv2.resize(prob_resized, (w, h)) prob_resized np.transpose(prob_resized, (2, 0, 1)) mask, floating_mask, alert_boxes, total_area postprocess_mask(prob_resized, min_area) return mask, floating_mask, alert_boxes, total_area这段代码里最值得关注的是postprocess_mask函数中的两个设计形态学开运算和连通域分析。开运算能有效去掉水面反光造成的零星误检点——这些误检点通常只有几十个像素但如果不处理会被误判成漂浮物报警连通域分析则把分散的漂浮物像素聚合成一个个目标区域方便后续按面积逐一判断。4.2 报警判断逻辑面积阈值的三层设计报警判断是整个项目从“图像分割”走向“实际应用”的关键一环。分割模型输出的是像素级类别但用户真正关心的是“现在需不需要处理”。阈值设低了一片落叶飘过就报警系统成了“狼来了”阈值设高了真正的大面积污染反而漏报。我常用的设计是三层阈值层级参数名默认值判定逻辑单目标面积min_area300 像素单个连通域面积小于该值的忽略总覆盖面积比max_area_ratio0.005漂浮物总像素占全图比例超过 0.5% 触发报警连续帧确认confirm_frames3 帧连续 3 帧都触发才输出报警防止单帧误检def check_alert(alert_boxes, total_floating_pixels, image_area, frame_count): 根据面积阈值输出报警消息 alert_messages [] # 条件1: 单个大面积漂浮物 for (x, y, w, h, area) in alert_boxes: if area 5000: # 单个目标面积超过 5000 像素立刻报警 alert_messages.append(f检测到大面积漂浮物: 位置({x},{y}), 面积{area}像素) # 条件2: 漂浮物总覆盖比例 area_ratio total_floating_pixels / image_area if area_ratio 0.005: alert_messages.append(f水面漂浮物覆盖率异常: {area_ratio:.4%}) # 条件3: 连续帧确认 if len(alert_messages) 0: if frame_count 3: alert_messages.append(连续3帧确认触发报警) return True, alert_messages else: return False, [检测到异常等待连续帧确认...] return False, [正常]frame_count控制的连续帧确认机制需要配合视频流使用这里单独提出来是因为它在课程设计的演示环节特别有用——现场用一段测试视频跑推理如果没有连续帧确认偶尔几帧的误检会让报警信息一闪一闪看起来非常不专业。加了这个机制后报警消息稳定输出答辩演示效果好得多。5. 避坑指南训练与推理中最容易翻车的五个问题5.1 训练损失不降反升Loss 震荡剧烈现象训练到第 10 轮左右CrossEntropy Loss 在 0.7 到 1.2 之间反复横跳完全没有下降趋势。原因最常见的是学习率设置过大或者 Poly 学习率衰减的初始值太高另一个隐患是用了torchvision自带的 COCO 预训练权重而不是 ImageNet 权重导致 encoder 的初始特征分布与水面场景差异过大。解决把初始学习率从1e-3降到1e-4并把 encoder 的 BN 层从train模式改成eval模式用encoder.freeze_bn()方法观察 5 轮内 Loss 是否稳定下降。我一般还会顺手确认一下类别权重 tensor 是否在正确的 device 上cuda还是cpu这个低级错误会导致对部分类别权重不生效位置很隐蔽。5.2 验证集 mIoU 很高但可视化结果却很糟现象验证集 mIoU 能到 0.85 以上但把预测掩码贴回原图一看漂浮物边缘粗糙很多细长的漂浮物直接断裂成好几段。原因mIoU 是全局像素统计指标漂浮物类别像素数少即使它整片预测错误对 mIoU 的影响也有限但人眼对目标的完整性非常敏感边缘断裂一眼就能看出来。解决关注每个类别的 IoU 单独数值特别是漂浮物类别的 IoU 是否超过 0.5另外在推理后处理里加一步稠密条件随机场DenseCRF后处理能显著改善边缘连续性。虽然会多花几百毫秒但对于课程设计演示来说视觉效果带来的答辩加分远远大于这点延迟。5.3 水面反光被大面积误检成漂浮物现象推理结果里阳光照射下的水面波纹区域被分割成漂浮物而且连通域面积很大触发报警。原因水面镜面反射的亮斑在纹理和亮度上确实与某些白色漂浮物相似模型无法仅凭空间特征区分训练数据里如果反光样本少这个问题会更严重。解决三步走——第一训练数据扩增时加入亮度扰动和对比度扰动RandomBrightnessContrast(0.2, 0.2)第二后处理中先用 HSV 色彩空间提取高饱和区域做参考如果判定为漂浮物的区域在 HSV 的 V 通道上亮度超过 250 且饱和度低于 30很可能是反光要降权处理第三把误检样本单独收集重新标注后加入训练集做难例挖掘。这三步做完误检率一般能下降一半以上。5.4 输出掩码保存成图片后全黑现象用cv2.imwrite保存预测掩码打开一看是纯黑图。原因模型的输出是类别索引像素值只有 0、1、2直接以灰度图保存时这些数值映射到 8 位灰度几乎全黑。解决保存前乘一个映射系数或者转成调色板模式# 正确做法: 将类别ID映射到可见灰度 vis_mask (mask * 85).astype(np.uint8) # 0, 85, 170 对应三个类别 cv2.imwrite(output_mask.png, vis_mask) # 或者用伪彩色映射 color_map np.array([[0, 0, 0], [0, 255, 0], [255, 0, 0]], dtypenp.uint8) # 黑/绿/红 color_vis color_map[mask] cv2.imwrite(output_color.png, cv2.cvtColor(color_vis, cv2.COLOR_RGB2BGR))这个坑几乎每个第一次做分割项目的人都会踩但在答辩演示时如果出现预测结果全黑观感会非常差。做模型训练前先跑一次单图推理并保存可视化结果确认颜色映射没问题再大规模训练。5.5 推理速度太慢视频流卡成 PPT现象单张 1920x1080 的图推理耗时 1.5 秒处理视频时画面严重卡顿。原因直接对全分辨率图做推理计算量过大而且torch.no_grad()虽然省了梯度计算但数据在 CPU 与 GPU 之间逐张拷贝的耗时被忽略了。解决先下采样到 512x512 推理然后上采样回原分辨率用torch.cuda.stream做异步传输如果机器支持半精度把模型转成 FP16 推理速度能提升约 40%。另外把多次推理循环外的公共操作比如模型复制到 device提到循环外避免每一帧都重复执行。6. 效果验证与提分技巧把分割结果做成可视化的对比报告课程设计答辩时评委最想看到的并不是模型在验证集上的 mIoU 数字——他们想看到“你这个模型真正解决了什么问题”。一个有效的做法是生成对比可视化报告随机抽 6 张测试图每张图三列排列——原图、真实掩码、预测掩码并配上每张图的漂浮物像素占比、连通域数量、是否触发报警的判断结果。import matplotlib.pyplot as plt def visualize_results(image_path, true_mask_path, pred_mask, save_pathreport.png): fig, axes plt.subplots(2, 3, figsize(15, 10)) image cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB) true_mask cv2.imread(true_mask_path, cv2.IMREAD_GRAYSCALE) # 类别颜色映射: 背景黑, 水体绿, 漂浮物红 color_map np.array([[0, 0, 0], [0, 200, 0], [200, 0, 0]], dtypenp.uint8) true_vis color_map[true_mask] pred_vis color_map[pred_mask] axes[0, 0].imshow(image); axes[0, 0].set_title(Original Image) axes[0, 1].imshow(true_vis); axes[0, 1].set_title(Ground Truth Mask) axes[0, 2].imshow(pred_vis); axes[0, 2].set_title(Prediction Mask) # 计算每类像素占比 total_pixels pred_mask.shape[0] * pred_mask.shape[1] floating_ratio (pred_mask 2).sum() / total_pixels axes[1, 0].text(0.5, 0.5, fFloating Object Area Ratio:\n{floating_ratio:.4%}, hacenter, vacenter, fontsize14) axes[1, 0].axis(off) # 输出报警判断 alert_triggered ALERT: Floating object detected if floating_ratio 0.005 else NO ALERT axes[1, 1].text(0.5, 0.5, alert_triggered, hacenter, vacenter, fontsize12, colorred if ALERT in alert_triggered else green) axes[1, 1].axis(off) plt.tight_layout() plt.savefig(save_path, dpi150, bbox_inchestight) plt.close()做报告时还可以把不同 epoch 的 checkpoint 分别跑一次同样的可视化流程从中挑出效果最好的那个用于演示。这个做法有一个实际意义很多时候最后一个 epoch 的模型由于学习率衰减到接近零可能在验证集上的指标更高但由于破坏了之前学到的某些特征分布可视化效果反而不如倒数第二个 epoch 的 checkpoint。这也是我自己的血泪教训——有一次打榜项目就是盲目用最后一个 checkpoint 提交分数反而掉了。从那以后我每次训练完都会保留所有 checkpoint至少对每 10 个 epoch 的模型各做一次可视化对比再决定用哪个版本做最终推理。这个过程一共不到半小时但能避免“模型指标高、实际效果差”的尴尬局面。希望这篇文章能帮你把这个课程设计做扎实答辩顺利。本文还有配套的精品资源点击获取
返回列表