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

资讯详情

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

VGG模型实战:珊瑚识别迁移学习与PyTorch实现

VGG模型实战:珊瑚识别迁移学习与PyTorch实现 简介面向需要入门卷积神经网络图像分类的PyTorch学习者这份资源围绕珊瑚种类识别场景提供了从数据集划分、模型训练到PyQt界面展示的完整CNN流程。代码基于PyTorch搭建VGG风格网络三个Python文件分别负责生成训练列表、执行训练与推理、以及可视化交互界面每一行均配有中文注释并附带Word说明文档和环境依赖清单便于小白对照理解与本地复现。资源包共8个文件以Python脚本、JPG示意图片、环境依赖txt及Word说明文档为主压缩包仅213KB轻量易下载。需注意压缩包不含真实数据集图片使用者需按文件夹提示自行搜集脑珊瑚、软珊瑚等类别的图片放入对应目录。目前已有70人学习适合希望结合具体项目快速上手PyTorch图像分类的初学者参考实践。1. 珊瑚识别最大的坑不在模型而在图片怎么进模型珊瑚种类识别经常被当成一个普通的图像分类任务来处理但这恰恰是错误率高发的起点。同一颗珊瑚在水下不同深度、不同光照下拍出来的颜色差异可能比两个种类之间的差异还要大模型如果只学到色彩统计换一个拍摄环境就失效。VGG模型在这个场景里反而有优势它通过多层小卷积核堆叠出对纹理和结构的敏感度而不是依赖大面积的颜色均值。加上这个压缩包明确不含数据集图片核心工作并不是翻论文找结构而是把散落的图片整理成模型能读的清单再让预训练VGG在自有数据上完成迁移学习。适合读这篇内容的人是手头有水下影像或保护区监测图、但没有现成训练集想快速搭起实验基线的开发者。2. VGG模型网络结构拆解选16还是19取决于你的数据集2.1 小卷积核堆叠的归纳偏置3×3如何扩大感受野理解VGG模型先看它最核心的设计没有大卷积核全部使用3×3卷积加2×2最大池化。两个3×3卷积堆叠感受野等效于一个5×5卷积三个堆叠等效于7×7卷积。这样做的两个好处是参数更少、非线性更强两个3×3总共18个权重替代一个5×5需要的25个权重中间还多了一次ReLU激活让特征组合更复杂。对珊瑚识别来说真正有效的判别线索通常不是单点颜色而是小孔、枝杈、表面脊线之间的排列关系。3×3小卷积核适合提取这类局部纹理池化层又负责保留主要结构整体过程正好完成从边缘、纹理、局部形状到整体形态的逐级抽象。水下影像往往有轻微模糊和色偏但这种结构化的纹理特征比颜色特征稳定得多。还有一个常被忽略的细节VGG没有引入BatchNorm它在迁移学习时少了一类统计量跟踪问题。微调时不需要担心running mean被小batch size带偏这在样本量小的场景里省了不少事。2.2 VGG16与VGG19的差异小样本场景别只追层数标题里有vgg模型落地时先要确认用哪个变体。VGG16和VGG19的差别只在卷积层数量前者13层后者16层多出的三层都堆在features块末端。下表是我在珊瑚识别项目里常用的判断依据对比维度VGG16VGG19珊瑚场景判断卷积层数量1316差异集中在最后几层特征提取主体一致参数总量约1.38亿约1.44亿VGG19略大但差距不构成主要压力底层特征与VGG19几乎相同与VGG16共享两者都适合加载ImageNet预训练权重深层语义足够处理纹理分类略强一些珊瑚的类间差异小提升不明显结论默认选择数据量大时再试几百到几千张的样本量VGG16更稳VGG19多出的卷积层对全局语义识别有帮助但珊瑚识别更多依赖局部纹理。而且多出来的参数在小样本下更容易放大多数类频率的影响验证集上的波动会比VGG16明显。我在实际项目里固定用VGG16做基线除非样本量过万且类别数大于30才会拿VGG19做对照。2.3 不含数据集图片的工程CSV清单要承担更多逻辑这个压缩包明确不带图片意味着数据加载逻辑必须能接受外部图片目录。常见做法是维护一个CSV清单每行指向一张图片和它的标签。除了image_path和label建议再写两列split表示划分集合source_id表示同一颗珊瑚的编号防止随机切分造成数据泄漏。image_path,label,split,source_id images/acr_001.jpg,Acropora,train,A01 images/acr_002.jpg,Acropora,train,A01 images/poc_003.jpg,Pocillopora,val,P07字段要求并不复杂但坑都在细节里。image_path建议用相对路径这样换机器时只要移动整个项目目录不必改脚本里的绝对路径。label的拼写和大小写必须统一否则同一个类会被拆成两个类类别数直接翻倍。split列写进清单而不是由脚本随机生成是为了保证每次复现实验时训练和验证的划分一致。提示source_id用于5.1节的分组切分。如果清单里没有这列同一颗珊瑚的不同照片很可能被同时分进训练集和验证集导致准确率虚高。3. 用PyTorch搭珊瑚图片加载管线VGG模型怎么在你的数据集上跑起来3.1 从CSV清单到Dataset不能靠ImageFolder硬凑既然没有现成目录结构torchvision自带的ImageFolder用起来很别扭。它要求图片按类别放在子目录里而我们的图片可能散落在多个采集批次中。更直接的方式是自己写一个Dataset从CSV清单按split过滤样本。import csv import torch from PIL import Image from torch.utils.data import Dataset class CoralDataset(Dataset): def __init__(self, csv_path, splittrain, transformNone): self.samples [] self.class_to_idx {} with open(csv_path, r, encodingutf-8) as f: reader csv.DictReader(f) for row in reader: if row[split] ! split: continue self.samples.append((row[image_path], row[label])) # 类别索引按 CSV 中出现顺序动态生成换数据集不用改代码 for _, label in self.samples: if label not in self.class_to_idx: self.class_to_idx[label] len(self.class_to_idx) self.transform transform def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] # 统一转 RGBVGG 的 Conv3d 期望三通道输入 # 混入灰度图会在 batch 拼接时报错 image Image.open(path).convert(RGB) if self.transform is not None: image self.transform(image) return image, self.class_to_idx[label]这段代码的核心是class_to_idx动态生成换自己的数据集时只需要改CSV路径不需要碰模型代码这也是标题里“不含数据集图片”的真实含义。convert(RGB)是必要防御水下相机可能输出灰度或带Alpha通道的图不做转换模型输入维度会不一致。split过滤放在__init__里而不是__getitem__里是为了保证每个epoch遍历时样本顺序稳定也避免每次迭代都扫描一遍整个CSV。3.2 用预训练VGG权重替换分类头冻结策略怎么定加载预训练VGG模型并替换最后一层是迁移学习的标准动作。这里要注意两点权重参数用官方接口加载避免自己写路径分类头的输出维度改成当前数据集的类别数。import torch.nn as nn from torchvision import models def build_vgg(num_classes, variantvgg16, freeze_featuresTrue): if variant vgg16: model models.vgg16(weightsmodels.VGG16_Weights.IMAGENET1K_V1) else: model models.vgg19(weightsmodels.VGG19_Weights.IMAGENET1K_V1) if freeze_features: for param in model.features.parameters(): param.requires_grad False # VGG16 分类头原结构是 4096 - 4096 - 1000# 只替换最后输出层 model.classifier[6] nn.Linear(4096, num_classes) return model需要注意model.classifier不是单一线性层而是一个Sequential容器前两个全连接层仍然是4096维度所以新输出层可以保持4096 - num_classes不变。常见做法里还有一种更激进的替换策略把整个分类头改成更小的结构# 如果不想背 1.38 亿参数的包袱可以换成更小的分类头 model.classifier nn.Sequential( nn.Linear(25088, 256), # 25088 来自 7 * 7 * 512 nn.ReLU(inplaceTrue), nn.Dropout(0.3), nn.Linear(256, num_classes) )这里25088是features层输出特征图展平后的长度。小分类头在样本量少时更不容易过拟合代价是如果后续想解冻features层做联合微调容量可能不够。关于features层解冻多少我在项目里常用下面这个策略表冻结策略适用情况显存占用收敛速度冻结features只训练分类头图片量少想快速验证管线低快解冻features最后几个卷积块水下光照分布与ImageNet差异大中中全量微调样本量大且计算资源充足高慢解冻features时不需要手动数层数用切片方式即可for param in model.features[20:].parameters(): param.requires_grad True。起始下标建议打印model.features结构后确认不同网络变体布局不同。3.3 训练主循环与最优模型保存只看准确率不够训练循环本身不复杂但有两个点必须处理好优化器只接收requires_grad为True的参数验证阶段必须关闭dropout和数据增强。import torch import torch.optim as optim from torch.utils.data import DataLoader def run_training(train_ds, val_ds, num_classes, epochs30): model build_vgg(num_classes) device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) # 分类头是重新初始化的用稍大的学习率 train_params [p for p in model.classifier.parameters() if p.requires_grad] optimizer optim.Adam(train_params, lr1e-4) loss_fn nn.CrossEntropyLoss() train_loader DataLoader(train_ds, batch_size16, shuffleTrue, num_workers2) val_loader DataLoader(val_ds, batch_size16, shuffleFalse, num_workers2) best_acc 0.0 for epoch in range(epochs): model.train() for images, labels in train_loader: images, labels images.to(device), labels.to(device) logits model(images) loss loss_fn(logits, labels) optimizer.zero_grad() loss.backward() optimizer.step() model.eval() correct total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) preds model(images).argmax(dim1) correct (preds labels).sum().item() total labels.size(0) acc correct / total print(fepoch{epoch} val_acc{acc:.4f}) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_coral.pt)保存state_dict而不是整个model是为了后续加载时结构想改就能改只要先用build_vgg重建模型再load_state_dict就不会被序列化版本问题卡住。优化器只传给train_params是因为冻结层不需要计算梯度。如果解冻了features层要用optimizer.add_param_group({params: unfrozen_params, lr: 1e-5})追加一组更小学习率的参数而不是把所有参数混在一起。提示Windows环境下DataLoader的num_workers建议设为0否则直接在交互环境里运行可能报多进程错误。Linux服务器上可以调到4以上。4. 珊瑚类别训练参数调整学习率、损失函数与数据增强的边界4.1 双段学习率与自动衰减分类头和特征层不能同速率珊瑚识别里最常被问的参数就是学习率。直接照搬ImageNet分类任务里的1e-3会出问题分类头是随机初始化的它需要较快的收敛速度但features层是预训练好的学习率一大几步就会把学到的通用纹理特征冲掉。常见做法是把两组参数分开设置用验证集准确率做监控触发平台期后自动降学习率。from torch.optim import lr_scheduler optimizer optim.Adam([ {params: model.classifier.parameters(), lr: 1e-4}, {params: model.features.parameters(), lr: 1e-5}, # 要求解冻过 ]) scheduler lr_scheduler.ReduceLROnPlateau( optimizer, modemax, factor0.5, patience5, verboseTrue ) # 每个 epoch 结束后 scheduler.step(val_acc)factor0.5表示每次降低一半patience5表示连续5个epoch验证集准确率不提升才触发。这两个值的组合比较保守适合样本量小、指标容易波动的珊瑚场景。如果batch_size只有8分类头学习率降到5e-5更稳妥batch_size到32以上可以回到2e-4。如果loss一直不下降先别急着调学习率。优先检查类别数是否和CSV里的label去重数一致label拼写不一致会让类别数翻倍损失函数永远学不完。其次看图片路径是否真的读到了图很多管线在读取失败时用上一张图兜底训练看起来正常实际输入完全错位。4.2 类别不均衡的损失函数CrossEntropy权重与FocalLoss水下调查数据里优势种和稀有种数量差距往往很大某一类占比可能超过80%。这种数据下验证准确率没有意义因为模型只需预测多数类就能拿到高分。需要先按类别频率算权重把少数类的损失信号放大。import torch import torch.nn.functional as F def make_class_weight(label_list, num_classes): cnt torch.zeros(num_classes) for idx in label_list: cnt[idx] 1 weight 1.0 / (cnt 1e-6) # 数量越少权重越大 weight weight / weight.mean() # 归一化到均值 1保持学习率语义 return weight把weight传给nn.CrossEntropyLoss(weightweight.to(device))即可。归一化这步很关键如果不缩放少数类权重可能达到10以上等效于把学习率放大十倍优化过程会剧烈震荡。如果类别严重不均衡且多数类样本太容易被分对可以用FocalLoss把已分对样本的损失继续压低。简化实现如下class FocalLoss(nn.Module): def __init__(self, alphaNone, gamma2.0): super().__init__() self.alpha alpha self.gamma gamma def forward(self, logits, targets): ce F.cross_entropy(logits, targets, weightself.alpha, reductionnone) p torch.exp(-ce) return (ce * (1 - p) ** self.gamma).mean()gamma2.0是常见起步值它压低的只是那些confidence很高的正确预测对还分不对的样本影响不大。alpha可以沿用上一节的category weight。要提醒的是类别权重解决的是损失被多数类主导的问题解决不了样本量个位数的类别本身无法学习的问题。少于10张的类别我一般会直接剔除或合并到相近的超类硬保留只会让验证集指标来回跳动。4.3 图像增强CNN算法的合理边界不是越强越准数据增强的直觉是越多越好但珊瑚识别有其特殊性颜色和结构本身就是判别依据增强幅度过大会直接扭曲类别。比如hue旋转过大会把硬珊瑚的褐色改成接近软珊瑚的颜色。下面这组参数是我在增强CNN算法场景里常用的起点。增强操作参数范围作用风险边界RandomResizedCropscale(0.7, 1.0)模拟不同拍摄距离scale小于0.5会丢失整体结构RandomRotationdegrees20模拟相机角度变动角度过大违背珊瑚朝上生长的先验RandomHorizontalFlipp0.5镜面对称水下常见无方向语义时可一直开RandomVerticalFlipp0.1少数俯仰拍摄场景垂直翻转改变光照方向p不宜高ColorJitterbrightness0.2, contrast0.2, saturation0.2, hue0.05模拟水质与光照变化hue超过0.1会让颜色失真成另一类RandomErasingp0.25模拟鱼群或泥沙遮挡纹理细密的类别慎用对应的transform代码是from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomRotation(20), transforms.RandomHorizontalFlip(p0.5), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2, hue0.05), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])Normalize沿用ImageNet的均值和方差即可。迁移学习中不建议自己统计数据集的RGB分布一是样本量小统计不稳定二是预训练权重已经适配这套归一化参数。验证集transform只做Resize、ToTensor和Normalize任何随机增强都不能出现在验证阶段否则指标不可比。5. 验证方法与说明文档把调参经验留给下一个看代码的人5.1 按个体分组切分别让同源照片泄漏珊瑚识别最容易犯的数据集切分错误是把同一颗珊瑚的多张照片随机分进训练和验证。水下视频抽帧得到的图片中相邻帧几乎相同随机切分会把近重复样本放进两个集合验证准确率虚高到90%以上真实场景表现却只有60%。正确做法是按source_id做分组切分。from sklearn.model_selection import GroupShuffleSplit import pandas as pd df pd.read_csv(coral_manifest.csv) gss GroupShuffleSplit(n_splits1, test_size0.2, random_state42) train_idx, val_idx next(gss.split(df, groupsdf[source_id])) df.iloc[train_idx].to_csv(train.csv, indexFalse) df.iloc[val_idx].to_csv(val.csv, indexFalse)groups参数决定了同一颗珊瑚的图片只能落在一个集合里。验证准确率要反映的是模型见过足够多珊瑚后对新珊瑚照片的判断能力而不是对同一颗珊瑚不同帧的识别能力。5.2 用混淆矩阵定位珊瑚的易混类别只看准确率看不到模型具体卡在哪两类之间。scikit-learn的classification_report可以直接给出逐类的precision、recall和F1。from sklearn.metrics import classification_report, confusion_matrix cm confusion_matrix(y_true, y_pred, labelslist(dataset.class_to_idx.values())) report classification_report(y_true, y_pred, target_nameslist(dataset.class_to_idx.keys()), digits3) print(report)如果Acropora和Pocillopora频繁互相误判说明两者的枝状结构在有限分辨率下确实难以区分。这时候最有效的动作不是改模型而是把预测置信度接近0.5的样本截图找出来肉眼核对标签是否标错再把典型易混样本整理成对照表放进说明文档。这类样本往往比新增几十张普通图片更能提升验证指标。5.3 逐行注释的粒度写“为什么”不写“是什么”逐行注释最容易犯的错是把代码翻译成中文比如img Image.open(path)旁边写“打开图片”这对理解项目没有增量。真正有价值的注释是解释约束比如为什么必须convert(RGB)为什么class_to_idx要按出现顺序生成。这些信息决定了别人换自己的数据集时能不能避开同样的坑。# 统一转 RGBVGG 的 Conv3d 期望三通道输入 # 混入单通道灰度图会在 batch 拼接时维度不一致 image Image.open(path).convert(RGB)说明文档则要交代四件事数据清单的字段含义、环境版本、一条可复现的训练命令、一次baseline实验记录。baseline表里写清模型变体、学习率、增强参数、验证集F1和训练耗时后面任何人调参都有对照。最后一点经验算完各类别数量后先把最少的几类合并进相近的超类再开始训练这比加任何注意力模块都更容易换来验证集上的稳定提升。本文还有配套的精品资源点击获取
返回列表