☰
图像分类实战代码拆解:从数据准备到训练评估的完整指南
2026/9/28 12:14:47 网站建设 项目流程

“第五节课—图像分类实战代码分析”,这个标题光是看起来就让人想起一个很现实的场景:很多同学学了卷积神经网络、看了图像分类的理论、背下了各种模型的名字,但一打开训练代码就发懵——这么长一段到底在干什么?挑了某个环节复制进自己的项目里,结果不是报错就是效果完全不对。

我这几年带人做视觉项目,见得最多的情况倒不是算法不懂,而是代码看不懂、改不对。理论层面的 ResNet、Transformer、注意力机制大家都聊得头头是道,一到实战环节就开始挠头。之所以出现这种问题,是因为网上能找到的“图像分类代码”要么是几百行的官方文档源码,要么是别人项目的魔改版,中间缺乏一层“把代码逐段拆开,告诉你每一行到底为什么这么写”的讲解。

这篇文章做的是最笨也最扎实的事:从数据集准备开始,把一段完整的图像分类训练代码按步骤拆开分析。不追求把代码写得花哨,而是让每个读了这篇文章的人,回头看到自己手里那段图像分类代码时,能看出结构、看懂逻辑、知道出了问题去哪里查。

1. 数据集是前提,但代码的重点不在下载而在组织——从零看数据准备脚本

很多人拿到一个图像分类任务,第一反应就是去搜数据集下载地址。但对于代码分析来说,下载数据只是起点,真正决定后面训练能不能顺跑的,是数据进入模型之前那些“看不见”的代码逻辑。这一类代码在教科书中通常被一笔带过,在实战里却恰恰是翻车最密集的区域。

1.1 数据下载与目录规划:分类任务的第一步是“把数据摆放整齐”

以经典的图像分类数据集为例,通常你下载完成后会得到一个压缩包,解压后里面可能是按类别分好的文件夹,也可能是一堆平铺的图片外加一个标注文件。无论哪种情况,在写代码之前,我强烈建议先统一成下面这种目录结构:

data/ train/ cat/ 001.jpg 002.jpg dog/ 001.jpg 002.jpg val/ cat/ 001.jpg dog/ 001.jpg

这种按类别建文件夹的方式,配合 PyTorch 的ImageFolder类可以直接加载,不需要自己写复杂的标注解析逻辑。要是你手里的数据集是类似 CIFAR-10 那类已经封装好的下载接口,那就更省事,直接调用接口下载就行。

但这里我要多提醒一句:不要因为代码里能直接下载数据,就放弃检查数据本身。下载完成之后,最好写个几行的统计脚本,看一下训练集和验证集各自的图片总数、每个类别的图片数、图片尺寸分布。我遇到过一个挺常见的坑:某个类别的图片在验证集中只有个位数,训练时整体准确率挺高,一输出每个类别的结果才发现这个类基本全错了——这种问题就是数据分布不均造成的,跟模型本身没关系。

1.2 自定义 Dataset 类的三个关键点:init、len、getitem

有的数据集不能用ImageFolder直接吃进来,那时候就需要自己写一个Dataset类。很多初学的同学一听到“自定义 Dataset”就紧张,其实拆开看,核心就三个函数:

from torch.utils.data import Dataset from PIL import Image import os class MyDataset(Dataset): def __init__(self, root_dir, transform=None): self.root_dir = root_dir self.transform = transform self.classes = sorted(os.listdir(root_dir)) self.samples = [] for cls in self.classes: cls_dir = os.path.join(root_dir, cls) for img_name in os.listdir(cls_dir): self.samples.append((os.path.join(cls_dir, img_name), cls)) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label = self.samples[idx] image = Image.open(img_path).convert("RGB") if self.transform: image = self.transform(image) label_idx = self.classes.index(label) return image, label_idx

这里最容易被忽略的是convert("RGB")这行。很多图片是 RGBA 四通道或者灰度单通道的,如果不统一转成三通道,模型第一层卷积可能直接报维度错误,或者更隐蔽地——不报错但训练效果变差。这种错误在写代码时不容易发现,因为你不一定会注意到某个数据集里混进了一张黑白图。

1.3 数据增强与 DataLoader 的节奏感

数据增强这段代码值得你单独挑出来分析,因为它的配置很大程度上决定了模型的泛化能力。常见的训练集增强组合是随机翻转、随机裁剪、归一化;验证集则只做归一化和尺寸缩放,不做随机增强。

from torchvision import transforms train_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) val_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ])

我经常看到有人把训练集和验证集的数据增强配成一模一样,这其实是个误区。验证集想要的是模型在“干净”数据上的表现,叠加随机扰动会让评估结果不稳定——同一个模型很有可能因为某次验证时随机翻转了一下就把准确率波动了一两个点。

再说DataLoader这个环节。参数看着不多,但num_workers和pin_memory这两个值得认真调一调。num_workers是数据加载的进程数,一般取 4 到 8 比较常见,但如果你用的是 Windows 系统,这里踩过头的概率较大,建议先从 0 开始调,跑通了再往上加。pin_memory建议设为True,尤其是在 GPU 训练的时候,数据从 CPU 内存搬到 GPU 显存的效率会好一些。

from torch.utils.data import DataLoader train_loader = DataLoader(dataset=train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(dataset=val_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True)

顺带提一个非常容易踩的坑:验证集和测试集的shuffle必须设为False。因为评估阶段你不需要打乱数据顺序,而且乱序之后你自己想对比某几张图的实际预测结果时,很难对应上样本。训练集则相反,一定要shuffle=True,否则每个 epoch 输入到模型的样本顺序完全一致,会严重影响梯度更新的随机性。

2. 模型定义代码:为什么大部分项目都在用“加载预训练模型+换分类头”

看图像分类代码,你会发现一个规律:不管课程里讲了多少种模型结构,实战代码里最常见的写法就是——把别人训练好的权重拿过来,然后把最后一层分类器换掉,接着去训练自己的数据。这段代码要分析清楚,得先弄明白几个关键选择背后的逻辑。

2.1 从 ResNet 到 Transformer:代码选型的现实逻辑

前几年图像分类代码的模型部分基本都是 ResNet 系列,近两年 Transformer 结构的模型在分类代码中出现得越来越频繁。热搜词里的“transformer图像分类”也不是什么新鲜玩意了,ViT(Vision Transformer)和各类变体已经成了项目任务里相当常见的 baseline。

但这里有个反直觉的现象:在普通中规模数据集上,从头训练一个 ViT 的效果往往不如直接用 ResNet 预训练权重做迁移学习。原因在于 Transformer 结构的模型对数据量的要求远高于传统的卷积神经网络。你看很多图像分类实战代码里,虽然模型名写的是 ViT,但加载的还是 ImageNet 上预训练好的权重。这一点要在代码分析时重点讲清楚——选型不是追求模型越大越好,而是权衡自己的数据量、算力、任务复杂度之后做出的妥协。

从实战角度,我给你的建议是:

  • 数据量小(几千到几万张):用 ResNet 预训练模型做迁移学习,稳扎稳打。
  • 数据量中等(几万到几十万张):可以考虑 EfficientNet 或 ViT 的轻量版本。
  • 数据量很大且算力充足:直接上大模型训练全套参数。

2.2 迁移学习代码里的几个关键细节

典型的加载预训练模型代码长这样:

import torchvision.models as models model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) num_features = model.fc.in_features model.fc = torch.nn.Linear(num_features, num_classes)

第二行那个.fc.in_features在代码分析时值得多说两句。很多人想换成结构更强的 ResNet50,但拿到代码后只知道把resnet18改成resnet50,而fc那一行不用动——因为不管主干是几层的 ResNet,最后一层输出的特征维数都是一样的,故in_features取出来再用,代码的通用性会好很多。

但如果你换成 Transformer 结构的模型,这个“换分类头”的写法就不太一样了。ViT 的最终分类输出头通常叫heads.head,特征维度也和 ResNet 不一样,修改时要看清楚模型内部结构,别再用model.fc硬套,那段代码会直接报错。

2.3 参数量与计算量怎么看

分析模型代码时,可以顺手加两行统计:

def count_parameters(model): return sum(p.numel() for p in model.parameters() if p.requires_grad) print(f"Trainable parameters: {count_parameters(model) / 1e6:.2f} M")

这个数字直接决定了你的显存需求和训练一个 epoch 所花的时间。ResNet18 参数量约 1100 万左右,ResNet50 约 2500 万左右,ViT-Base 是 8600 万左右。参数多不代表一定好,但显存占用和训练时间是实打实的成本,这也解释了为什么很多实战项目至今仍坚持用 ResNet 系列当骨架。

3. 训练循环的代码没有那么玄乎,但坑全藏在细节里

图像分类代码中最核心也最长的一段,通常就是训练循环。它的结构其实高度固定,但为什么每个人跑出来的结果差异这么大?原因在于细节处理。

3.1 模型、数据、损失、优化器:一套标准循环的骨架

一段典型的 PyTorch 训练循环长这样:

import torch import torch.nn as nn import torch.optim as optim device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=1e-4) num_epochs = 20 for epoch in range(num_epochs): model.train() running_loss = 0.0 correct = 0 total = 0 for images, labels in train_loader: 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, 1) total += labels.size(0) correct += (predicted == labels).sum().item() epoch_loss = running_loss / total epoch_acc = correct / total print(f"Epoch [{epoch+1}/{num_epochs}], Loss: {epoch_loss:.4f}, Acc: {epoch_acc:.4f}")

这段代码逐行分析下来,有经验的人会注意optimizer.zero_grad()这个位置。我见过不少初学者写的代码把这个过程放在loss.backward()之后,这是一个经典的逻辑错误——梯度是累加的,不清零的话每个 batch 的梯度会叠加上一个 batch 的梯度,更新方向完全乱掉。

再注意loss.item()的用法。loss这个张量如果直接参与打印或者累加,它所在的计算图并不会自动释放,会带来显存缓慢增长的问题。用.item()取出 Python 数值,一是切断计算图,二也能防止训练时显存越占越多直到内存溢出。

3.2 训练与验证交替的逻辑,以及一个常见的“shuffle翻车”现场

一套合格的分类训练代码,训练循环内部一般会在训练完一个 epoch 之后接一个验证过程。它的作用是看模型在没见过的数据上的表现,代码大致如下:

model.eval() val_loss = 0.0 val_correct = 0 val_total = 0 with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) val_loss += loss.item() * images.size(0) _, predicted = torch.max(outputs, 1) val_total += labels.size(0) val_correct += (predicted == labels).sum().item() val_epoch_loss = val_loss / val_total val_epoch_acc = val_correct / val_total print(f"Val Loss: {val_epoch_loss:.4f}, Val Acc: {val_epoch_acc:.4f}")

这个验证循环有两点要注意:一是model.eval()必须配上with torch.no_grad(),二是这段代码里绝对不能有optimizer.step()。前者是为了关掉 dropout 和 batch normalization 的训练行为,同时不保存中间梯度;后者则是评估的底线——如果验证循环里混入了梯度更新,评估结果就成了“开卷考试”,模型瞥到了正确答案,后续的数据对比会整体失真。

有的同学会问:为什么训练时每个 batch 都要跑一次model.train()和model.eval()来回切换,不能训练完再验证吗?这个问题其实出在十行以内就能解释清楚的责任模型身上——BatchNorm 层的统计数据是跟着前向传播不断更新的,验证阶段如果不切到eval(),统计的是训练 batch 的分布,那个分布恰好又被验证代码计算了一遍,自然就是两个不统一的口径。之前我遇到过一个团队在验证时忘记切model.eval(),又开着 dropout,结果验证准确率和训练准确率始终差六七个点,查了半天最后定位到这行。

3.3 学习率调度和每多少轮保存模型:最容易偷懒但也最值钱的配置

训练循环中有两个部分在很多课程 demo 里被省略了,但对实战来说极有价值:一个是学习率调度,一个是模型保存。

学习率调度我常用的是ReduceLROnPlateau,它的思路是:当验证集 loss 连续几个 epoch 不下降时,就把学习率乘一个系数往下调。代码很简单:

scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3) # 每个 epoch 验证完之后调用: scheduler.step(val_epoch_loss)

这种调度策略相当实用,因为手动调学习率谁也没那么好的手感,机器根据 loss 的下降趋势自动调整,比人拍脑袋来得稳定得多。

模型保存则建议配合验证准确率来决定存哪个版本:

best_acc = 0.0 if val_epoch_acc > best_acc: best_acc = val_epoch_acc torch.save(model.state_dict(), "best_model.pth")

这样能保证你最后拿到的不是最后一个 epoch 的模型,而是验证集上表现最好的那一个。很多人在训练结束时直接保存最后一轮,结果后面做测试时发现过拟合已经开始抬头,模型效果并不理想。

4. 从打印 loss 到画混淆矩阵:训练结束之后代码才真正开始

训练循环跑完,屏幕上出现一个“准确率 98%”的数字,很多人的分析就到此为止了。但真正有价值的代码分析,往往体现在训练结束后那几段不起眼的评估代码里。

4.1 准确率只是第一层,混淆矩阵和分类报告才是调优入口

98% 的整体准确率听起来很不错,但一个类别占了数据的 90% 的时候,这个 98% 一点意义也没有。要看清模型在每个类别上的真实表现,需要用以下代码生成混合矩阵和分类报告:

from sklearn.metrics import confusion_matrix, classification_report import numpy as np all_preds = [] all_labels = [] model.eval() with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, predicted = torch.max(outputs, 1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm = confusion_matrix(all_labels, all_preds) print(cm) print(classification_report(all_labels, all_preds, target_names=class_names))

从这份报告里你能很清晰地看到哪些类别的 precision、recall 偏低,这时候调优方向就明确了——不是盲目调学习率,而是考虑这类样本是不是太少,需不需要做类别加权,或者用数据增强补足这一类。迷茫的时候,多看一眼混淆矩阵比什么都有用。

4.2 预测结果可视化:让你不是“感觉模型还行”而是“看到模型行”

还有一种代码值得加到分析里——把预测结果画到图上。它本身不参与训练,但能帮你快速定位模型为什么分错。

import matplotlib.pyplot as plt def visualize_predictions(model, val_loader, class_names, num_images=8): model.eval() images, labels = next(iter(val_loader)) images, labels = images.to(device), labels.to(device) with torch.no_grad(): outputs = model(images) _, preds = torch.max(outputs, 1) fig, axes = plt.subplots(2, 4, figsize=(12, 6)) axes = axes.ravel() for idx in range(num_images): img = images[idx].cpu().permute(1, 2, 0).numpy() img = img * [0.229, 0.224, 0.225] + [0.485, 0.456, 0.406] img = np.clip(img, 0, 1) axes[idx].imshow(img) axes[idx].set_title(f"True: {class_names[labels[idx]]}\nPred: {class_names[preds[idx]]}") axes[idx].axis('off') plt.tight_layout() plt.show()

注意代码里的反归一化那段,很多人直接把图片 tensor 丢给plt.imshow,出来的画面是花的,原因在于 ImageNet 均值和标准差被叠加进去了。把增强的参数给它“倒回去”,视觉效果就正常了。

4.3 显存清理:调参过程中被忽略的 Python 进程陷阱

图像分类代码分析中还有一个内容几乎不会写进课程大纲,但实际做项目时极其磨人——GPU 显存释放的问题。训练跑到一半显存不够报错是很常见的。除了前面提到的用.item()切断计算图之外,还有两个常见原因:

第一是DataLoader的num_workers开的进程没有及时回收,尤其是 Windows 环境下,训练结束之后进程还占着显存。你往往会看到已经关闭了 Python 窗口,但python.exe进程还在后台。这种时候不要犹豫,直接进行进程清理,把残留的 Python 进程结束掉,再重新启动训练。

第二是代码里某个位置不小心把不需要的 tensor 保留了引用,导致它无法被 GC 释放。排查这种问题可以用torch.cuda.memory_summary()打印显存分配情况,看看哪一行代码把显存弄爆了。这类经验属于“没遇到过就不会提前写”的坑,但一旦遇到,你会在项目里比别人省出半天时间。

5. 把课程代码搬到真实场景:从森林图像分类聊起

从热搜词里我注意到“森林图像分类”这个方向。拿它做例子特别合适,因为这一类真实场景任务和课程里的 CIFAR 数据集差异很大,正好用来说明“课程代码”与“实战代码”之间的距离在哪里。

5.1 真实数据集的三个特征:不均衡、细粒度、拍摄环境复杂

CIFAR-10 这种数据集的类别都是猫、狗、飞机、汽车这种差别很大的物体,类别之间边界清晰,背景简单,甚至每张图片都已经被居中裁剪好了。但森林图像分类面对的是什么?同样是“树”这个类别,底下还有树种之分、季节之分、光照之分。同样一片森林,晴天和雨天拍出来的模型输入完全不一样。

真实项目里最常见的三个问题:类别样本数量差距大(像某些珍稀树种的图像数量可能只有常见树种的十分之一)、类别间差异小(不同树种的叶子纹理很接近)、同类别内差异大(同一棵树不同角度拍出来的图可能看起来不像同一类)。这些都会让直接从课程里抄来的代码效果大打折扣。

5.2 从任务出发修改模型代码:多标签/多类别的边界判断

图像分类代码里有个隐含假设:一张图只属于一个类别。森林图像分类如果是做“这个区域主要是什么植被类型”,单标签没问题;但如果你想识别“这张图里有哪几类树木”,这就是多标签分类任务,代码要做相应的调整。

多标签分类在代码层面的改动主要有三处:

  • 损失函数从CrossEntropyLoss换成BCEWithLogitsLoss。
  • 最后一层输出维度从 N 改成 N 个独立的二分类输出。
  • 预测时不再用torch.max找最大索引,而是对每个输出节点做 sigmoid 之后和阈值(比如 0.5)比较,大于阈值的类别都视为存在。

这个改动看起来不大,但如果你没意识到任务本身就出了变化,直接用单标签代码硬跑,模型在训练时就会被 loss 的数值带偏——它只能输出一个类别,而真实标注里可能同时有两个类别都对,那模型学到的“正确”就成了“每次选一个可能性最大的”,其余全错。

5.3 代码之外:图像分类实战中“数据代码”往往占了一半工作量

最后聊一个真实但不常出现在课程里的比例。很多人以为图像分类实战代码的重头戏是模型结构和训练循环,但在实际项目里,数据清洗、数据标注检查、数据划分、类别映射关系维护这段代码,通常要占整个项目代码量的一半以上。

就拿数据划分来说,课程里常见的是按文件夹随机划分训练验证集。但在真实场景中,如果你的数据不是一个时间点采集的,分训练集和验证集时必须考虑时间维度——比如前三个月的照片做训练,后一个月的照片做验证。如果随机划分,同一个时间拍摄的同场景照片可能同时出现在训练集和验证集里,验证指标虚高得离谱。这种代码逻辑不涉及任何深度学习理论,但它在真实项目里的价值比换一个更强模型大得多。

你去看任何一份森林图像分类或其他真实场景项目的代码,最值得精读的往往不是那个几十行的训练循环,而是前处理部分那些大量处理“脏数据”的函数——这些才决定了模型在天花板内能飞多高。

写到这里突然想多说一句,每次我在实际项目里碰壁,回过来反思代码里问题根源,最常发现的就是“想当然”三个字。图像分类代码分析的意义就在于此——它不教你新算法,而是在锻炼一种把理论代码和真实数据对接起来的判断力。把今天拆过的这些细节吃透,回去再看你手里那份图像分类代码,真的会感觉哪里都亮堂了不少。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询