ResNet图像多分类实战:从数据预处理到混淆矩阵全解析
发布时间:2026/9/15 2:44:27来源:尧图网络
简介这是一个基于ResNet实现的2D图像简单多分类完整工程面向深度学习初学者与图像分类入门者。资源围绕数据准备、残差网络搭建、训练验证与结果可视化展开提供完整的PyTorch工程代码。压缩包共20个文件以12个Python脚本为主涵盖数据增强、onehot标签生成、模型定义、进度条日志工具、评估及可视化模块另有6个pyc编译文件和2张图片结果图包体仅382KB轻量高效。项目结构清晰从train_main.py主入口到model/resnet.py模型定义逐层对应便于对照学习图像分类完整流程。目前已有215人学习下载可帮助理解ResNet残差连接与训练调参思路也可直接作为课设或项目改写的基底。1. ResNet 2D多分类别忽略数据管线里的坑很多人第一次用ResNet做2D图像多分类第一反应是“加载预训练模型替换全连接层开始训练”。这个思路本身没有错但我在实际运行这个项目时发现真正卡住精度和收敛速度的不是模型而是数据怎么送进网络。常见的做法是直接读图片、缩放到224x224、除以255但这样处理过的数据在ResNet的BatchNorm层上会产生明显分布偏移导致前几个epoch损失下降极慢。另一个容易忽视的点是简单多分类任务虽然类别少但样本不均衡和标签噪声带来的影响比模型结构更大。本文基于一个完整的ResNet分类工程从数据预处理、残差块实现、训练闭环到混淆矩阵可视化把每一步的关键参数和踩坑点拆开讲。2. 数据预处理与增强data_process.py到底做了什么2.1 从原始图片到ResNet能吃的张量ResNet的输入规格一般是224x224的RGB图像通道顺序为C,H,W数值范围根据预训练模型要求决定。项目里的data_process.py核心职责就是把磁盘上的散落图片统一成固定尺寸的numpy数组并生成标签。最常见的做法是用OpenCV读取图片后执行resize和归一化但需要注意两点resize的插值方式不能默认用双线性对于物体边缘敏感的任务cv2.INTER_AREA在缩小图片时保留纹理信息的效果更好归一化系数不是简单的除以255而是按均值[0.485, 0.456, 0.406]和方差[0.229, 0.224, 0.225]做标准化这个数值来自ImageNet统计迁移到自己数据集时如果不使用预训练权重也可以改成数据集的真实均值和标准差但需要从头训练。# 项目 data_process.py 中归一化的典型写法 import cv2 import numpy as np def preprocess_image(image_path, target_size(224, 224), is_trainTrue): img cv2.imread(image_path) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) if is_train: # 训练时期随机裁剪模拟多尺度输入 h, w img.shape[:2] new_h, new_w int(h * 0.9), int(w * 0.9) img cv2.resize(img, (new_w, new_h), interpolationcv2.INTER_AREA) img cv2.random_crop(img, target_size) else: img cv2.resize(img, target_size, interpolationcv2.INTER_LINEAR) # 除255后按通道标准化 img img.astype(np.float32) / 255.0 mean np.array([0.485, 0.456, 0.406], dtypenp.float32) std np.array([0.229, 0.224, 0.225], dtypenp.float32) img (img - mean) / std return img.transpose(2, 0, 1) # HWC - CHW这里cv2.random_crop是放大的模拟实现实际工程里常用PIL写一个随机裁剪类。注意末尾的transpose决定了网络输入的通道顺序PyTorch的卷积层要求为[batch, channel, height, width]如果漏掉这一行模型会报形状错误。很多人在自家代码里踩过这个坑明明用OpenCV读图喂给ResNet却报维度不匹配多半是忘了这一步。2.2 数据增强策略不只是翻转摘要里提到的imgaug_dataProcess.py说明项目支持第三方增强库。我一般会在简单多分类任务里控制总增强强度不做过度的仿射变换。原因是类别少时图片语义往往集中在全局形状过度旋转反而破坏物体局部特征。常用策略是随机水平翻转、对比度微调、高斯噪声以及0.2概率的随机灰度化。增强操作关键参数建议值适用场景RandomHorizontalFlipp0.5通用数据集避免方位敏感RandomResizedCropscale(0.7, 1.0)强化尺度不变性ColorJitterbrightness/contrast0.2/0.2光照变化明显的场景GaussNoisestd0.01传感器噪声SmallCauseRotationdegrees15倾斜不影响语义的任务增强写进数据加载器而不是写进预处理脚本是因为训练和验证需要不同的增强路径。验证集只做resize和归一化不添加任何随机扰动这能保证混淆矩阵指标稳定复现。2.3 标签生成防止类别错位generateOnehotLabel_txt.py看起来是生成onehot标签的脚本但实际工程里我建议直接使用索引标签即类别ID等计算Loss时再由训练框架内部转onehot避免在高维numpy数组里存储稀疏编码。关键点是标签排序要和训练列表文件保持一致。我的做法是先从文件夹名构建类别字典再遍历所有图片生成image_path label_index的txt文件。这个顺序写死之后无论增强还是分批次采样都以这个txt为准防止不同进程里的随机种子不同导致数据错位。3. 残差块与模型封装在resnet.py里改出适合自己的结构3.1 残差块的原始意图与实现差异ResNet由Kaiming He在2015年提出核心是残差块内的恒等映射。对于50层以下的ResNet使用的是BasicBlock每个块包含两个3x3卷积激活函数后置50层以上使用Bottleneck用1x1卷积降维再升维以减少计算量。项目的model/resnet.py里大概率实现了这两个结构。要注意的是BasicBlock的跳跃连接在特征图尺寸减半时不能直接相加必须通过1x1卷积调整通道数。很多复现的代码在第一个残差块之后直接把输入池化这不正确。stride2的卷积要在主路径的第一个卷积上执行同时跳跃连接的1x1卷积也要设置stride2这样才能使尺寸匹配。import torch.nn as nn import torch.nn.functional as F class BasicBlock(nn.Module): expansion 1 def __init__(self, in_channels, out_channels, stride1): super(BasicBlock, self).__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, stride1, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_channels) self.shortcut nn.Sequential() if stride ! 1 or in_channels ! out_channels: self.shortcut nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size1, stridestride, biasFalse), nn.BatchNorm2d(out_channels) ) def forward(self, x): identity x out F.relu(self.bn1(self.conv1(x))) out self.bn2(self.conv2(out)) out self.shortcut(identity) out F.relu(out) return out这段代码里biasFalse是因为卷积后面接BatchNorm偏置会被归一化层抵消不加反而省参数。shortcut在输入输出维度不一致时做映射注意不是每层都加否则模型训练不稳定。我在实际训练中发现跳跃连接后的ReLU必须放在加法完成之后如果先对主路径分支做ReLU再相加梯度回传时会出现两条路径值域打架收敛变慢。3.2 网络组装与全连接层的修正项目里的ResNet封装会按照[3, 4, 6, 3]等层数配置堆叠残差块。自定义分类时经典做法是保留预训练模型前四层提取特征替换最后的平均池化和全连接层为适合自己类别数的结构。但简单多分类任务里类别间差异往往只依赖局部纹理这时可以考虑在全局平均池化后接一个Dropout(0.2)再加线性层能抑制过拟合。如果使用对类别数不敏感的调查我会先固定resnet18做基线把全连接层改成nn.Sequential(nn.Dropout(0.2), nn.Linear(512, num_classes))验证集准确率有明显提升因为项目数据量并不支持深层的50/101层结构发挥优势。3.3 预训练模型的选择与加载热搜里包含resnet预训练模型这确实值得单独说。加载torchvision.models.resnet18(pretrainedTrue)时默认会下载在ImageNet上训练好的权重。但注意预训练权重里的全连接层输出是1000维不能直接用于自己的分类任务。加载时应该利用state_dict的键名过滤只加载卷积和BN层的权重。很多初学项目会在这一步报错因为num_classes改了之后fc权重尺寸不匹配。常见做法是忽略不匹配的key或者显式剔除最后一层等待新训练的fc层覆盖。import torchvision.models as models model models.resnet18(pretrainedTrue) in_features model.fc.in_features model.fc nn.Sequential(nn.Dropout(0.2), nn.Linear(in_features, 2)) # 只加载除了fc以外的预训练权重 pretrained_dict models.resnet18(pretrainedTrue).state_dict() model_dict model.state_dict() pretrained_dict {k: v for k, v in pretrained_dict.items() if k in model_dict and fc not in k} model_dict.update(pretrained_dict) model.load_state_dict(model_dict)这段代码里先保存了原始fc的输入维度再替换成自己的分类头。过滤fc not in k时还要注意如果自己修改了残差块内部的层数某些键会消失需要打印model_dict里所有键检查是否匹配。我在一个猫狗分类项目里就是这样加载的最后的准确率比随机初始化高15个点左右。4. 训练闭环与超参train_main.py的运行机制与调参4.1 优化器与损失函数的选择先设计训练过程再对每个细节逐项讨论。优化器我优先推荐SGD配合Nesterov动量而不是无脑选Adam。因为分类任务的损失曲面比较平滑Adam短时间收敛快但容易停留在泛化性较差的极值点。损失函数选交叉熵它在PyTorch中的实现会结合LogSoftmax和NLLLoss所以传入网络输出的logits不要手动归一化。这样处理数值稳定性更好同时允许低概率类别有足够梯度。一个容易被忽略的参数是label_smoothing。多分类任务如果训练样本存在标注错误onehot标签会让模型对正确类别过度自信损失容易震荡。我在这个简单分类项目里设0.1的平滑系数验证集集准确率波动明显降低。4.2 训练循环的动态操作train_main.py中一个epoch的完整流程包括遍历训练数据清零梯度前向计算计算loss反向传播梯度裁剪更新参数。关键细节是在每个epoch结束后立即做梯度裁剪max_norm5.0能防止异常样本导致梯度爆炸这在用SGD训练时尤其重要。import torch from utils.logger import Logger def train_epoch(model, dataloader, criterion, optimizer, epoch, writer): model.train() running_loss 0.0 correct 0 total 0 for batch_idx, (images, labels) in enumerate(dataloader): images, labels images.cuda(), labels.cuda() optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() running_loss loss.item() * images.size(0) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() if batch_idx % 20 0: Logger().info(fepoch:{epoch} batch:{batch_idx} loss:{loss.item():.4f}) train_loss running_loss / total train_acc correct / total writer.add_scalar(train_loss, train_loss, epoch) writer.add_scalar(train_acc, train_acc, epoch) return train_loss, train_acc这里running_loss通过乘images.size(0)累加样本总量的loss避免dataloader最后一批不够时计算偏差。predicted和labels比较时注意predicted是Tensor类型与labels维度要一致。PyTorch里常见错误是predicted labels时如果两者都带grad_fn会报错因此必须用.data或.detach()分离。此脚本的writer来自TensorBoard工具记录标量图。4.3 学习率调度与超参表学习率策略我采用ReduceLROnPlateau当验证损失连续三个epoch不降时按0.1倍回调初始学习率。不要使用每隔固定step数衰减的StepLR因为简单任务的收敛节奏不稳定自适应回调更安全。初始学习率对SGD设为0.01如果使用Adam可以改成0.001。下面给出我在该工程中实验后的推荐超参数组合参数推荐值备注batch_size32显存不足时降到16epochs50结合early stoppinginit_lr0.01 (SGD)或0.001 (Adam)momentum0.9NesterovTrueweight_decay1e-4大模型可调1e-5label_smoothing0.1标签嘈杂时有效这些参数在train_main.py里以argparse方式声明。验证时要注意模型切换验证阶段必须model.eval()并且torch.no_grad()包裹计算过程。否则BatchNorm和Dropout会在验证时继续更新状态导致验证指标失真。我见过有人在这里踩坑结果验证准确率周期性震荡其实是模型还在训练模式。4.4 验证集的评估与混淆矩阵热搜词中有python多分类混淆矩阵代码训练过程中的验证阶段正好用上。多分类不能只看准确率混淆矩阵能揭示哪些类别互相干扰。项目里提供的eval.py和misc.py负责这个功能。import numpy as np from sklearn.metrics import confusion_matrix def compute_confusion(model, loader, num_classes): model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in loader: images, labels images.cuda(), labels.cuda() outputs model(images) _, preds torch.max(outputs, dim1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) return confusion_matrix(all_labels, all_preds, labelslist(range(num_classes)))confusion_matrix的labels参数指定了行的顺序如果不加这个参数sklearn会从输入里自动推导类别集合如果某类样本在验证集中一个都没出现矩阵维度就会短缺。在多分类中我每次都会显式传num_classes防止预测类别和真实类别集合不一致。得到的矩阵可以配合display_labels绘制可视化热力图直观看出误分类的成对类别。5. 混淆矩阵与可视化验证模型不是只看准确率5.1 用混淆矩阵定位系统性错误训练完成后我会先打印出归一化后的混淆矩阵。归一化有两种方式按行归一化表示每个真实类别的召回率按列归一化表示每个预测类别的查准率。对多分类任务来说更值得关注行归一化因为能看出哪个类别的样本被频繁误判。例如某个类被另一个类大面积吸收通常是特征相似或标注边界不清。看这个矩阵时我习惯把阈值设定在0.1以上也就是说真实类别的样本有超过10%被错分就值得返回训练集检查数据质量。5.2 可视化与导出技巧visualize.py里提供了类似的方法但我在生产环境中更倾向把矩阵输出成CSV文件方便和同事协同分析。代码可以在上面的compute_confusion返回后利用pandas写入磁盘。import pandas as pd import seaborn as sns import matplotlib.pyplot as plt # cm 是由 compute_confusion 得到的结果矩阵 cm_df pd.DataFrame(cm, index[f真实{i} for i in range(cm.shape[0])], columns[f预测{i} for i in range(cm.shape[1])]) cm_df.to_csv(confusion_matrix.csv, encodingutf-8-sig) plt.figure(figsize(8, 6)) sns.heatmap(cm_df, annotTrue, fmtd, cmapBlues) plt.xlabel(预测类别) plt.ylabel(真实类别) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi300)代码里encodingutf-8-sig很重要否则用Excel打开中文文件名会乱码。fmtd让热力图显示整数索引。如果类别数量比较大超过20类我建议把annot关闭只保留颜色深浅否则热力图全是数字看不清结构。这个技巧在简单多分类任务也适用。5.3 把结果回灌到训练流程拿到混淆矩阵后一个实用的做法是根据对角线外的误报数量动态调整类别权重。在train_main.py里如果某个真实类别经常被误判到另一个类别可以提高另一个类别的权重或在采样器里对该类样本多采样一个epoch。我没有采用复杂的手动调权重方法而是用一个回调函数读取每个epoch验证集混淆矩阵如果发现某类准确率低于0.7就在下个epoch把该类的交叉熵权重乘以1.2。项目中的utils/bar_utils.py提供了进度条正好把这一指标显示出来。最后把验证结果写入日志文件方便比较多个版本模型的混淆矩阵差异这才是多分类任务最值得长期维护的产物。本文还有配套的精品资源点击获取
网站建设高端定制企业官网