ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

ConvNeXt V2图像分类实战:从环境搭建到模型部署全流程指南

ConvNeXt V2图像分类实战:从环境搭建到模型部署全流程指南 简介卷积神经网络CNN作为计算机视觉的经典架构通过局部连接和权值共享高效提取图像特征。其核心原理在于利用卷积核在输入数据上滑动进行特征映射具有平移不变性和参数效率高的优势。在Transformer架构席卷CV领域的背景下ConvNeXt V2通过引入全卷积掩码自编码器FCMAE预训练策略显著提升了纯卷积模型的性能上限证明了CNN架构在现代化训练方法下仍具强大竞争力。该技术在工程实践中展现出优异的硬件友好性和部署便利性特别适用于资源受限场景下的图像分类任务如森林图像分类、病虫害识别等实际应用。本文将以ConvNeXt V2为例系统讲解从环境配置、数据准备到模型微调、训练优化的完整实现路径。1. 从“老树新花”到“森林图像分类”为什么现在还要关注ConvNeXt V2如果你最近在关注图像分类领域尤其是那些听起来很“卷”的竞赛或者实际项目比如“森林图像分类”、“病虫害识别”你可能会被各种层出不穷的模型名字搞得眼花缭乱。从Transformer席卷CV开始好像大家都在讨论ViT、Swin Transformer仿佛传统的卷积神经网络CNN已经成了“古典”技术。但事实真的如此吗ConvNeXt V2的出现恰恰给了我们一个重新审视CNN潜力的绝佳机会。它不是什么颠覆性的新架构而是在经典的ConvNeXt基础上通过引入一个名为“全卷积掩码自编码器”FCMAE的预训练策略让这棵“老树”开出了惊艳的“新花”。简单来说ConvNeXt V2证明了在正确的“训练方法”加持下纯卷积架构的性能天花板远比我们想象的要高。那么对于我们这些需要解决实际问题的开发者来说ConvNeXt V2意味着什么首先它提供了一个在速度和精度之间取得优异平衡的选择。相比于一些计算密集的Transformer模型ConvNeXt V2的卷积操作在通用硬件尤其是没有特殊优化过的GPU上往往能跑得更快内存占用也更友好。其次它的架构清晰没有那么多复杂的注意力机制需要理解对于从经典CNN如ResNet过渡过来的开发者非常友好。最后也是最重要的一点它在包括ImageNet在内的多个标准数据集上达到了与顶尖Transformer模型媲美的性能。这意味着当你下一个“森林图像分类”项目需要在有限的计算资源下追求尽可能高的准确率时ConvNeXt V2是一个非常值得放入候选清单的模型。本系列文章就将手把手地带你完成使用ConvNeXt V2实现图像分类任务的全过程从环境搭建、模型解读到数据准备、训练调优最后到模型部署我们会深入每一个环节并分享那些官方文档里不会写的实操细节和踩坑经验。2. 环境搭建与模型初探不仅仅是pip install在开始写第一行训练代码之前一个稳定、可复现的环境是成功的基石。这里我推荐使用Conda来管理Python环境它能很好地处理不同项目间复杂的依赖冲突。2.1 创建并激活专用环境首先我们创建一个名为convnextv2的Python 3.9环境3.8-3.10通常都是兼容性较好的选择conda create -n convnextv2 python3.9 -y conda activate convnextv22.2 核心依赖安装PyTorch与TorchvisionPyTorch是我们的基础框架。访问PyTorch官网获取最适合你CUDA版本的安装命令。假设你的环境是CUDA 11.7安装命令如下pip install torch torchvision --index-url https://download.pytorch.org/whl/cu117关键细节务必确保PyTorch版本与CUDA版本匹配。你可以通过nvidia-smi查看CUDA版本并通过python -c “import torch; print(torch.__version__)”和python -c “import torch; print(torch.version.cuda)”验证PyTorch是否正确识别了CUDA。版本不匹配是后续很多诡异错误的根源。接下来安装一些必要的工具库pip install opencv-python pillow matplotlib scikit-learn pandas tqdm tensorboardopencv-python和PIL用于图像处理matplotlib用于可视化scikit-learn用于评估指标tqdm提供美观的训练进度条tensorboard用于监控训练过程。2.3 获取ConvNeXt V2官方实现ConvNeXt V2的官方代码库在GitHub上。我们将其克隆到本地git clone https://github.com/facebookresearch/ConvNeXt-V2.git cd ConvNeXt-V2进入目录后安装项目自身的要求依赖pip install -r requirements.txt实操心得官方requirements.txt有时会包含一些版本号非常严格的依赖可能会与你本地已安装的包冲突。如果遇到冲突一个比较稳妥的做法是先注释掉requirements.txt中已有核心包如numpy, torchvision的版本限制使用你当前稳定环境中的版本优先保证PyTorch环境的稳定。2.4 模型架构速览ConvNeXt V2的核心模块在动手训练前花几分钟理解模型的核心构成能让你在调试时更有方向。ConvNeXt V2的主体架构沿用了ConvNeXt可以看作是一个“现代化”的ResNet。其主要模块包括Patchify Stem 替代了传统CNN中堆叠小卷积核的“头部”。它使用一个较大的卷积核如4x4和较大的步长如4直接将输入图像分割成不重叠的图块Patch并进行嵌入这借鉴了ViT的思想能更高效地在下采样初期提取特征。ConvNeXt Block 这是核心构建块。每个Block主要由“深度可分离卷积Depthwise Conv”、“LayerNorm”和“倒瓶颈结构Inverted Bottleneck”的前馈网络FFN组成。特别注意它使用了“大核深度卷积”如7x7这是其获得强大感受野的关键。下采样层Downsampling Layers 在Stage之间使用一个步长为2的2x2卷积进行空间下采样同时增加通道数。全局平均池化与分类头 在提取所有特征后进行全局平均池化将每个通道的特征图压缩为一个标量最后接一个全连接层作为分类器。而ConvNeXt V2的“灵魂”在于其FCMAE预训练。它通过在输入图像上随机掩码掉一部分图块然后让模型去重建这些被掩码的像素。这个过程迫使模型学习到更强大、更通用的视觉表征。对于我们进行下游分类任务通常有两种方式一是直接使用官方发布的、经过FCMAE预训练的模型权重进行微调Fine-tuning这是最常用且高效的方法二是在自己的数据集上从头进行FCMAE预训练这需要海量数据和时间一般只在特定领域且数据充足时考虑。3. 数据准备构建高效的数据管道模型和代码都准备好了接下来就是“喂”给模型的数据。一个鲁棒的数据加载和预处理管道对训练稳定性至关重要。我们以经典的“猫狗分类”或你自己准备的“森林树种分类”数据集为例。3.1 数据集目录结构我强烈推荐使用以下目录结构它与torchvision.datasets.ImageFolder完美兼容能省去大量自己写数据加载逻辑的麻烦。your_dataset/ ├── train/ │ ├── class_a/ # 例如: oak │ │ ├── image1.jpg │ │ └── image2.jpg │ └── class_b/ # 例如: pine │ ├── image3.jpg │ └── image4.jpg └── val/ # 或 test/ ├── class_a/ │ └── image5.jpg └── class_b/ └── image6.jpgtrain和val目录下的子文件夹名就是类别标签。这种结构清晰且易于扩展。3.2 使用Torchvision进行数据加载与增强我们使用torchvision来构建数据管道。首先定义训练和验证时的数据增强Data Augmentation策略。增强能有效提升模型泛化能力防止过拟合。import torch from torchvision import datasets, transforms # 定义训练集的数据增强和预处理 train_transform transforms.Compose([ transforms.RandomResizedCrop(224), # 随机裁剪并缩放到224x224 transforms.RandomHorizontalFlip(p0.5), # 随机水平翻转 transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), # 颜色抖动 transforms.ToTensor(), # 转换为Tensor并归一化到[0,1] transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) # ImageNet标准归一化 ]) # 定义验证集的数据预处理通常不进行随机增强 val_transform transforms.Compose([ transforms.Resize(256), # 将短边缩放到256 transforms.CenterCrop(224), # 中心裁剪到224x224 transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])为什么用ImageNet的均值和标准差ConvNeXt V2的官方预训练权重是在ImageNet上训练的其输入数据经过了这样的归一化。使用相同的统计量可以确保输入分布与预训练时一致这是迁移学习成功的关键。即使你用自己的数据在微调初期也建议先使用这个统计量。接下来使用ImageFolder加载数据# 路径替换成你自己的数据集路径 train_dataset datasets.ImageFolder(rootpath/to/your_dataset/train, transformtrain_transform) val_dataset datasets.ImageFolder(rootpath/to/your_dataset/val, transformval_transform) # 创建数据加载器DataLoader train_loader torch.utils.data.DataLoader( train_dataset, batch_size64, # 根据你的GPU内存调整 shuffleTrue, # 训练集需要打乱 num_workers4, # 并行加载数据的进程数可加速数据读取 pin_memoryTrue # 将数据锁页内存加速GPU传输 ) val_loader torch.utils.data.DataLoader( val_dataset, batch_size64, shuffleFalse, # 验证集不需要打乱 num_workers4, pin_memoryTrue )踩坑提醒num_workers设置并非越大越好。通常设置为CPU核心数的2-4倍。设置过大可能导致内存占用过高甚至死锁。如果在Windows上遇到多进程问题可以尝试将num_workers设为0。pin_memoryTrue在GPU训练时能显著提升数据从CPU到GPU的传输速度务必开启。4. 模型加载与微调策略站在巨人的肩膀上现在我们进入核心环节加载预训练的ConvNeXt V2模型并为其适配我们自己的分类任务。4.1 加载预训练模型ConvNeXt V2官方提供了多种规格的预训练模型如convnextv2_tiny,convnextv2_base等。我们以convnextv2_tiny为例。你需要从官方仓库或Model Zoo下载对应的权重文件.pth或.npz格式。假设我们已将权重文件convnextv2_tiny_1k_224_fcmae.pt放在当前目录。加载模型并替换分类头的代码如下import torch import torch.nn as nn from models.convnextv2 import convnextv2_tiny # 根据官方代码结构导入模型定义 # 1. 初始化模型不加载预训练权重 model convnextv2_tiny(num_classes1000) # 先按原始1000类初始化 # 2. 加载预训练权重 checkpoint torch.load(convnextv2_tiny_1k_224_fcmae.pt, map_locationcpu) # 注意权重文件的key可能与模型state_dict的key不完全匹配可能需要处理 model.load_state_dict(checkpoint[model], strictFalse) # strictFalse允许不匹配的key # 3. 修改分类头适配我们自己的类别数 num_ftrs model.head.in_features # 获取原分类头输入特征维度 model.head nn.Linear(num_ftrs, len(train_dataset.classes)) # train_dataset.classes是类别列表 # 将模型移动到GPU device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device)关键细节解析strictFalse参数非常有用。因为预训练模型的分类头是针对ImageNet的1000类而我们新建的分类头参数是随机初始化的key对不上。strictFalse会加载所有能匹配的键即主干特征提取部分的权重而忽略不匹配的键即分类头这正是我们需要的。4.2 微调策略哪些层需要学习一个常见的误区是微调就是将所有参数都放开训练。实际上更精细的策略能带来更好的效果和更快的收敛。通常我们采用分层学习率和选择性冻结的策略。策略一仅训练分类头快速基准在数据量较少时可以先将模型主干特征提取器的所有参数冻结只训练新换上的分类头。这是最快的方案用于快速验证数据管道和任务可行性。# 冻结所有主干参数 for param in model.parameters(): param.requires_grad False # 仅解冻分类头的参数 for param in model.head.parameters(): param.requires_grad True策略二全模型微调标准做法当数据量相对充足时解冻所有参数进行训练。但为了稳定我们通常为**主干backbone和分类头head**设置不同的学习率。分类头是全新的需要更大的学习率快速学习而主干部分已有较好的特征需要用较小的学习率进行精细调整防止破坏已有的好特征。# 后续在定义优化器时为不同参数组设置不同学习率 optimizer torch.optim.AdamW([ {params: model.head.parameters(), lr: 1e-3}, # 分类头较大学习率 {params: model.parameters(), lr: 1e-4, weight_decay: 0.05} # 主干较小学习率 ])经验之谈对于ConvNeXt V2这类大型模型我强烈推荐使用AdamW优化器而非传统的SGD它通常收敛更快且对超参数尤其是学习率不那么敏感。weight_decay权重衰减是防止过拟合的重要正则化手段对于微调同样重要。4.3 损失函数与评估指标对于多分类任务交叉熵损失CrossEntropyLoss是标准选择。criterion nn.CrossEntropyLoss()评估指标我们主要看准确率Accuracy但为了更细致地分析模型在各类别上的表现可以同时计算混淆矩阵Confusion Matrix这在类别不平衡的“森林图像分类”等任务中尤其有用。from sklearn.metrics import accuracy_score, confusion_matrix, classification_report def evaluate(model, dataloader, device): model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in dataloader: images, labels images.to(device), labels.to(device) outputs model(images) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) acc accuracy_score(all_labels, all_preds) cm confusion_matrix(all_labels, all_preds) report classification_report(all_labels, all_preds, target_namesval_dataset.classes) return acc, cm, report5. 训练循环与超参数调优让模型真正“学”起来万事俱备只欠训练。一个完整的训练循环包括前向传播、损失计算、反向传播和参数更新。此外学习率调度和模型保存是提升效果的关键技巧。5.1 构建基础训练循环def train_one_epoch(model, dataloader, criterion, optimizer, device, epoch): model.train() running_loss 0.0 correct 0 total 0 for batch_idx, (images, labels) in enumerate(dataloader): images, labels images.to(device), labels.to(device) # 清零梯度 optimizer.zero_grad() # 前向传播 outputs model(images) loss criterion(outputs, labels) # 反向传播与优化 loss.backward() optimizer.step() # 统计 running_loss loss.item() * images.size(0) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() # 每N个batch打印一次日志 if batch_idx % 50 0: print(fEpoch [{epoch}], Batch [{batch_idx}/{len(dataloader)}], Loss: {loss.item():.4f}) epoch_loss running_loss / total epoch_acc 100. * correct / total return epoch_loss, epoch_acc5.2 学习率调度动态调整学习步伐固定的学习率可能不是最优的。常见的策略是“热身Warmup”“余弦退火Cosine Annealing”。Warmup在训练刚开始的少量步数内将学习率从0线性增加到初始学习率。这有助于稳定训练初期防止梯度爆炸。Cosine Annealing使学习率随着训练过程按照余弦函数的曲线从初始值衰减到0。这能让模型在后期更精细地收敛到最优点。我们可以使用torch.optim.lr_scheduler来实现组合调度from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR # 假设总epoch数为num_epochs warmup_epochs5 warmup_scheduler LinearLR(optimizer, start_factor0.01, end_factor1.0, total_iterslen(train_loader)*5) cosine_scheduler CosineAnnealingLR(optimizer, T_maxlen(train_loader)*(num_epochs-5), eta_min1e-6) # 在训练循环的每个epoch后调用 def scheduler_step(epoch): if epoch 5: warmup_scheduler.step() else: cosine_scheduler.step()5.3 模型保存与早停Early Stopping我们不仅要保存最终模型更要在验证集性能达到最佳时保存模型这通常称为“最佳检查点Best Checkpoint”。同时引入“早停”机制可以防止过拟合当验证集指标在连续多个epoch不再提升时自动停止训练。best_val_acc 0.0 patience 10 # 容忍多少个epoch性能不提升 counter 0 for epoch in range(num_epochs): train_loss, train_acc train_one_epoch(...) val_acc, _, _ evaluate(model, val_loader, device) # 学习率调度 scheduler_step(epoch) # 保存最佳模型 if val_acc best_val_acc: best_val_acc val_acc torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), val_acc: val_acc, }, best_convnextv2_checkpoint.pth) counter 0 # 重置计数器 print(f*** New best model saved with val_acc: {val_acc:.4f} ***) else: counter 1 # 早停判断 if counter patience: print(fEarly stopping triggered at epoch {epoch}) break避坑指南保存的检查点最好包含epoch、model_state_dict、optimizer_state_dict以及关键的指标。这样如果训练意外中断你可以从这个检查点恢复训练而不是从头开始这对于动辄训练几十个epoch的大模型至关重要。6. 可视化与调试用TensorBoard看清训练过程“黑箱”训练让人心里没底。TensorBoard是一个强大的可视化工具可以实时监控损失、准确率、学习率甚至图像样本。6.1 集成TensorBoard首先在代码中引入并配置SummaryWriterfrom torch.utils.tensorboard import SummaryWriter import os # 创建一个带有时间戳的日志目录方便区分不同实验 log_dir os.path.join(runs, fexp_{datetime.now().strftime(%Y%m%d_%H%M%S)}) writer SummaryWriter(log_dir)然后在训练循环的关键位置添加记录# 在每个epoch结束后记录 writer.add_scalar(Loss/train, train_loss, epoch) writer.add_scalar(Accuracy/train, train_acc, epoch) writer.add_scalar(Accuracy/val, val_acc, epoch) writer.add_scalar(Learning Rate, optimizer.param_groups[0][lr], epoch) # 可以记录一些图像样本例如第一个batch if epoch 0: images, _ next(iter(train_loader)) img_grid torchvision.utils.make_grid(images[:8]) # 取前8张 writer.add_image(Training images sample, img_grid, epoch)训练时在终端启动TensorBoardtensorboard --logdirruns然后在浏览器中打开http://localhost:6006你就能看到所有指标的实时曲线图。通过对比训练集和验证集的损失/准确率你可以轻松判断模型是欠拟合还是过拟合。如果训练损失持续下降但验证损失开始上升就是典型的过拟合信号需要加强正则化如增大weight_decay、添加Dropout或使用更多数据增强。6.2 常见问题排查清单训练过程中难免遇到问题这里提供一个快速排查清单问题现象可能原因排查步骤与解决方案Loss为NaN或突然变得巨大学习率过高数据中存在异常值如无效图像梯度爆炸。1. 大幅降低学习率如从1e-3降到1e-5。2. 检查数据加载过程确保图像被正确解码和归一化。可以在ToTensor()前添加transforms.Lambda(lambda x: x.float())。3. 使用梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)。训练准确率上升验证准确率停滞或下降过拟合。1. 增强数据增强如RandomRotation, RandomAffine。2. 增加权重衰减weight_decay。3. 在模型中添加Dropout层如果原模型没有。4. 获取更多训练数据或使用迁移学习。训练Loss几乎不下降学习率过低模型权重未正确更新如冻结了不该冻结的层数据标签错误。1. 尝试增大学习率。2. 打印模型参数检查requires_grad属性确保需要训练的层已解冻。3. 可视化一批训练数据及其标签确认数据与标签对应正确。GPU内存溢出OOMBatch Size过大模型或中间变量占用内存过多。1. 减小batch_size。2. 使用梯度累积每N个小batch进行一次optimizer.step()和zero_grad()模拟大batch效果。3. 使用混合精度训练AMP可显著减少内存占用并加速训练。7. 进阶技巧与优化从“能用”到“好用”当你的模型能够正常训练并收敛后下一步就是考虑如何让它更高效、更强大。7.1 混合精度训练Automatic Mixed Precision, AMPAMP通过使用半精度FP16进行计算和存储可以大幅减少GPU内存占用并可能加快训练速度尤其在现代Tensor Core GPU上效果显著。PyTorch内置了AMP支持使用非常简单from torch.cuda.amp import autocast, GradScaler scaler GradScaler() # 用于防止梯度下溢 def train_one_epoch_amp(model, dataloader, criterion, optimizer, device, epoch): model.train() running_loss 0.0 for images, labels in dataloader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() # 在autocast上下文管理器中进行前向传播 with autocast(): outputs model(images) loss criterion(outputs, labels) # 使用scaler进行反向传播和优化 scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() # ... 其余统计代码重要提示AMP并非万能。对于某些非常小的模型或特定的运算FP16可能导致数值不稳定如Loss变成NaN。如果遇到这种情况可以尝试调整GradScaler的初始化参数或者对模型中的某些模块如BatchNorm保持FP32精度。7.2 模型EMA指数移动平均EMA是一种在训练过程中维护模型权重滑动平均的技巧。在验证或测试时使用这个平均后的权重往往能获得比最终训练权重更稳定、泛化能力更好的模型。其原理是shadow_weights decay * shadow_weights (1 - decay) * model_weights。class ModelEMA: def __init__(self, model, decay0.9999): self.model model self.decay decay self.shadow {} self.backup {} self.register() def register(self): for name, param in self.model.named_parameters(): if param.requires_grad: self.shadow[name] param.data.clone() def update(self): for name, param in self.model.named_parameters(): if param.requires_grad: new_average (1.0 - self.decay) * param.data self.decay * self.shadow[name] self.shadow[name] new_average.clone() def apply_shadow(self): # 将EMA权重应用到模型 for name, param in self.model.named_parameters(): if param.requires_grad: self.backup[name] param.data param.data self.shadow[name] def restore(self): # 恢复原始权重 for name, param in self.model.named_parameters(): if param.requires_grad: param.data self.backup[name] self.backup {} # 在训练循环中使用 ema ModelEMA(model) for epoch in range(num_epochs): for batch in train_loader: # ... 训练步骤 ... optimizer.step() ema.update() # 在每个batch的optimizer.step()后更新EMA # 验证时使用EMA权重 ema.apply_shadow() val_acc evaluate(model, val_loader) # 此时model的权重已是EMA权重 ema.restore() # 验证完恢复训练权重7.3 超参数搜索的实用思路完全依赖手动调参效率低下。除了网格搜索和随机搜索一个更高效的实践是基于经验的手动迭代先找一个大致的范围学习率通常在1e-5到1e-2之间batch_size在能力范围内尽可能大32, 64, 128weight_decay在1e-4到1e-2之间。固定其他调学习率用一个较小的epoch数如5-10跑几个不同的学习率例如1e-4, 5e-4, 1e-3观察训练初期Loss的下降速度和稳定性。选择那个Loss下降稳定且速度合理的学习率。调整权重衰减固定学习率尝试不同的weight_decay观察验证集准确率防止过拟合。微调数据增强如果模型过拟合加强增强如CutMix, MixUp如果欠拟合减弱增强或使用更贴近真实测试数据的增强。对于资源充足的团队可以尝试使用更自动化的工具如Ray Tune或Optuna但理解上述手动过程背后的逻辑至关重要。8. 模型评估与结果分析不只是看准确率训练完成后在独立的测试集上评估模型是最后也是最重要的一步。不要只满足于一个整体的准确率数字。8.1 全面评估指标加载我们之前保存的最佳模型检查点在测试集上运行评估函数得到准确率、混淆矩阵和分类报告。# 加载最佳模型 checkpoint torch.load(best_convnextv2_checkpoint.pth) model.load_state_dict(checkpoint[model_state_dict]) model.eval() # 在测试集上评估 test_acc, test_cm, test_report evaluate(model, test_loader, device) print(fTest Accuracy: {test_acc:.4f}) print(Classification Report:) print(test_report)分类报告会给出每个类别的精确率Precision、召回率Recall和F1-score。这对于类别不平衡的数据集尤其有价值。例如在“森林图像分类”中如果“稀有树种”的样本很少模型可能倾向于将其预测为“常见树种”以获得更高的整体准确率。此时只看整体准确率会掩盖问题而F1-score能更好地反映模型对少数类的识别能力。8.2 混淆矩阵可视化混淆矩阵能直观地展示模型在哪里犯了错。import seaborn as sns import matplotlib.pyplot as plt plt.figure(figsize(10, 8)) sns.heatmap(test_cm, annotTrue, fmtd, cmapBlues, xticklabelstest_dataset.classes, yticklabelstest_dataset.classes) plt.ylabel(True Label) plt.xlabel(Predicted Label) plt.title(Confusion Matrix) plt.tight_layout() plt.savefig(confusion_matrix.png) plt.show()分析混淆矩阵中非对角线上的高值单元格。如果“类别A”经常被误判为“类别B”可能意味着这两个类别在视觉上本身就非常相似。训练数据中这两个类别的样本数量差异巨大。数据增强或预处理方式无意中模糊了这两个类别的区别。8.3 错误案例分析从失败中学习随机抽取一些被模型错误分类的样本进行可视化是提升模型和理解其局限性的最佳方式。def visualize_errors(model, dataloader, device, class_names, num_samples10): model.eval() errors [] with torch.no_grad(): for images, labels in dataloader: images, labels images.to(device), labels.to(device) outputs model(images) _, preds torch.max(outputs, 1) for i in range(images.size(0)): if preds[i] ! labels[i]: # 反归一化图像以便显示 img images[i].cpu().numpy().transpose((1, 2, 0)) mean np.array([0.485, 0.456, 0.406]) std np.array([0.229, 0.224, 0.225]) img std * img mean img np.clip(img, 0, 1) errors.append((img, class_names[labels[i]], class_names[preds[i]])) if len(errors) num_samples: break if len(errors) num_samples: break # 绘制错误样本 fig, axes plt.subplots(2, 5, figsize(15, 6)) axes axes.ravel() for idx in range(num_samples): axes[idx].imshow(errors[idx][0]) axes[idx].set_title(fTrue: {errors[idx][1]}\nPred: {errors[idx][2]}) axes[idx].axis(off) plt.tight_layout() plt.show() visualize_errors(model, test_loader, device, test_dataset.classes)通过观察这些错例你可能会发现一些规律是不是所有被误判的图片都光线很暗或者背景特别杂乱或者拍摄角度很特殊这些发现将直接指导你下一步的改进方向——是收集更多此类困难样本还是设计针对性的数据增强如模拟低光照、随机遮挡亦或是考虑引入更复杂的模型或训练技巧。本文还有配套的精品资源点击获取
返回列表