用 PyTorch 搭建一个可复用的 CNN 图像分类训练闭环

举报
优质中转haerapi 发表于 2026/08/02 11:23:13 2026/08/02
【摘要】 很多开发者第一次学习 CNN 时,容易停留在“卷积层提特征、池化层降维、全连接层分类”的概念层面;真正写训练代码时,却会遇到一组更工程化的问题:数据目录怎么组织、输入尺寸如何统一、训练和验证如何拆分、模型参数怎样保存、推理脚本如何复用训练时的预处理逻辑。本文不追求刷榜精度,也不编造某个数据集上的测试结果,而是搭建一个可复用的最小训练闭环。你可以先用自己的小型图片数据跑通流程,再替换模型结构、...

很多开发者第一次学习 CNN 时,容易停留在“卷积层提特征、池化层降维、全连接层分类”的概念层面;真正写训练代码时,却会遇到一组更工程化的问题:数据目录怎么组织、输入尺寸如何统一、训练和验证如何拆分、模型参数怎样保存、推理脚本如何复用训练时的预处理逻辑。

本文不追求刷榜精度,也不编造某个数据集上的测试结果,而是搭建一个可复用的最小训练闭环。你可以先用自己的小型图片数据跑通流程,再替换模型结构、增强策略或部署方式。

典型目录如下:

cnn-demo/
  config.py
  model.py
  train.py
  predict.py
  data/
    train/
      cat/
      dog/
    val/
      cat/
      dog/
  checkpoints/

这里使用 ImageFolder 约定:每个类别一个子目录,目录名就是类别名。真实项目中,建议将训练集和验证集提前固定下来,避免每次随机划分导致结果不可复现。

CNN 的核心原理

CNN 的优势来自局部连接和参数共享。普通全连接层会让每个输入像素都连接到每个输出神经元,参数量随图片尺寸快速膨胀;卷积层只在局部窗口内计算,并让同一个卷积核在整张图上滑动,因此能用较少参数捕捉边缘、纹理、局部形状等视觉模式。

一个基础图像分类 CNN 通常包含四类组件:

  • 卷积层:提取局部特征,例如边缘、颜色块、纹理组合。
  • 激活函数:引入非线性,常用 ReLU
  • 池化层:降低空间尺寸,减少计算量,并提高一定的位置鲁棒性。
  • 分类头:将高维特征映射为类别 logits,再交给损失函数计算误差。

需要注意,训练时模型输出通常不是概率,而是 logits。使用 nn.CrossEntropyLoss 时,不需要在模型末尾手动加 Softmax,因为该损失函数内部会处理对数概率计算。推理阶段如果要展示置信度,再对 logits 做 softmax 即可。

环境与配置

先安装依赖。具体版本应以你的项目环境为准,如果使用 GPU,还需要安装与你 CUDA 环境匹配的 PyTorch 构建包。

pip install torch torchvision pillow

把可变参数集中到 config.py,便于后续调整:

from pathlib import Path

ROOT = Path(__file__).resolve().parent
DATA_DIR = ROOT / "data"
TRAIN_DIR = DATA_DIR / "train"
VAL_DIR = DATA_DIR / "val"
CKPT_DIR = ROOT / "checkpoints"
CKPT_PATH = CKPT_DIR / "cnn_best.pt"

IMAGE_SIZE = 128
BATCH_SIZE = 32
EPOCHS = 10
LR = 1e-3
NUM_WORKERS = 2

如果你的项目需要访问私有对象存储或远程服务,不要把密钥写进代码,应从环境变量读取,例如:

import os

access_key = os.environ.get("APP_ACCESS_KEY")
if not access_key:
    raise RuntimeError("APP_ACCESS_KEY is required")

本文示例本身不需要任何密钥。

定义模型

下面是一个小型 CNN,适合用来验证训练链路。它不是面向生产精度优化的结构,但层次清晰,便于理解和修改。

import torch
from torch import nn

class SmallCNN(nn.Module):
    def __init__(self, num_classes: int):
        super().__init__()
        self.features = nn.Sequential(
            nn.Conv2d(3, 32, kernel_size=3, padding=1),
            nn.BatchNorm2d(32),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2),

            nn.Conv2d(32, 64, kernel_size=3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2),

            nn.Conv2d(64, 128, kernel_size=3, padding=1),
            nn.BatchNorm2d(128),
            nn.ReLU(inplace=True),
            nn.AdaptiveAvgPool2d((1, 1)),
        )
        self.classifier = nn.Linear(128, num_classes)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = self.features(x)
        x = torch.flatten(x, 1)
        return self.classifier(x)

这里使用 AdaptiveAvgPool2d((1, 1)),可以让分类头不依赖固定的中间特征图尺寸。只要输入图片经过预处理后尺寸一致,模型结构就更容易维护。

训练与验证流程

训练脚本要完成五件事:加载数据、构建模型、定义损失和优化器、循环训练、保存验证集表现最好的权重。

import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

from config import TRAIN_DIR, VAL_DIR, CKPT_DIR, CKPT_PATH, IMAGE_SIZE, BATCH_SIZE, EPOCHS, LR, NUM_WORKERS
from model import SmallCNN

def build_loaders():
    train_tf = transforms.Compose([
        transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
        transforms.RandomHorizontalFlip(),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
    ])
    val_tf = transforms.Compose([
        transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
    ])

    train_set = datasets.ImageFolder(TRAIN_DIR, transform=train_tf)
    val_set = datasets.ImageFolder(VAL_DIR, transform=val_tf)

    train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)
    val_loader = DataLoader(val_set, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)
    return train_loader, val_loader, train_set.classes

def evaluate(model, loader, criterion, device):
    model.eval()
    total_loss, correct, total = 0.0, 0, 0
    with torch.no_grad():
        for images, labels in loader:
            images, labels = images.to(device), labels.to(device)
            logits = model(images)
            loss = criterion(logits, labels)
            total_loss += loss.item() * images.size(0)
            preds = logits.argmax(dim=1)
            correct += (preds == labels).sum().item()
            total += labels.size(0)
    return total_loss / total, correct / total

def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    train_loader, val_loader, classes = build_loaders()

    model = SmallCNN(num_classes=len(classes)).to(device)
    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.Adam(model.parameters(), lr=LR)

    CKPT_DIR.mkdir(parents=True, exist_ok=True)
    best_acc = 0.0

    for epoch in range(1, EPOCHS + 1):
        model.train()
        running_loss = 0.0
        for images, labels in train_loader:
            images, labels = images.to(device), labels.to(device)
            optimizer.zero_grad()
            logits = model(images)
            loss = criterion(logits, labels)
            loss.backward()
            optimizer.step()
            running_loss += loss.item() * images.size(0)

        train_loss = running_loss / len(train_loader.dataset)
        val_loss, val_acc = evaluate(model, val_loader, criterion, device)
        print(f"epoch={epoch} train_loss={train_loss:.4f} val_loss={val_loss:.4f} val_acc={val_acc:.4f}")

        if val_acc > best_acc:
            best_acc = val_acc
            torch.save({"model": model.state_dict(), "classes": classes}, CKPT_PATH)

if __name__ == "__main__":
    main()

执行训练:

python train.py

如果你的机器没有 GPU,代码会自动使用 CPU,只是训练速度可能较慢。示例中的准确率输出只能反映当前数据、划分、增强方式和训练轮数,不能作为通用性能结论。

推理脚本

推理阶段必须复用验证阶段的尺寸调整和归一化逻辑,否则训练和推理的数据分布会不一致。

import sys
import torch
from PIL import Image
from torchvision import transforms

from config import CKPT_PATH, IMAGE_SIZE
from model import SmallCNN

def main(image_path: str):
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    checkpoint = torch.load(CKPT_PATH, map_location=device)
    classes = checkpoint["classes"]

    model = SmallCNN(num_classes=len(classes)).to(device)
    model.load_state_dict(checkpoint["model"])
    model.eval()

    tf = transforms.Compose([
        transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
    ])

    image = Image.open(image_path).convert("RGB")
    tensor = tf(image).unsqueeze(0).to(device)

    with torch.no_grad():
        logits = model(tensor)
        probs = torch.softmax(logits, dim=1)[0]
        idx = int(probs.argmax().item())

    print({"class": classes[idx], "confidence": float(probs[idx].item())})

if __name__ == "__main__":
    if len(sys.argv) != 2:
        raise SystemExit("usage: python predict.py path/to/image.jpg")
    main(sys.argv[1])

执行:

python predict.py ./sample.jpg

可执行改造建议

跑通最小闭环后,可以按优先级逐步改造:

  1. 先检查数据质量:类别目录是否正确、是否存在损坏图片、训练集和验证集是否混入重复样本。
  2. 再调整输入尺寸和 batch size:显存不足时优先降低 batch size,而不是盲目删模型层。
  3. 引入更强的数据增强:例如随机裁剪、颜色扰动,但验证集不要使用随机增强。
  4. 替换骨干网络:可以用 torchvision.models 中的预训练模型做迁移学习,但要确认输入归一化和分类头修改正确。
  5. 增加日志与配置管理:生产项目建议记录参数、代码版本、数据版本和模型文件路径。

这些改造的前提是先有稳定的训练、验证、保存和推理闭环。没有闭环时直接堆复杂模型,往往只会增加排查难度。

常见问题

1. 为什么训练集准确率升高,验证集不升反降?

常见原因是过拟合、训练验证分布不一致、数据量过小或验证集标注质量差。可以先减少模型容量、增加数据增强、固定划分方式,并人工抽查错误样本。

2. 为什么 CrossEntropyLoss 前不要加 Softmax

因为 CrossEntropyLoss 期望输入 logits,并在内部组合了对数 softmax 与负对数似然损失。提前加 Softmax 可能带来数值稳定性和梯度表达问题。

3. 为什么推理结果类别对不上?

ImageFolder 会按类别目录名生成类别索引。保存模型时应同时保存 classes,推理时读取同一份类别列表,避免手写类别顺序导致错位。

4. 小数据集是否适合从零训练 CNN?

可以用于学习流程,但未必适合获得稳定泛化能力。真实业务中,如果数据量有限,通常优先考虑迁移学习、冻结部分骨干层和更严格的数据清洗。

5. 多进程 DataLoader 在 Windows 上报错怎么办?

确保训练入口放在 if __name__ == "__main__": 下;如果仍不稳定,可以先把 NUM_WORKERS 改为 0 验证主流程。

总结

一个可维护的 CNN 项目不只是模型结构本身,还包括数据约定、预处理一致性、训练验证拆分、权重保存和推理复用。本文给出的 PyTorch 示例刻意保持简单,目标是让训练闭环清晰可运行。后续无论替换为 ResNet、MobileNet,还是加入更复杂的增强和部署逻辑,都应保留这条主线:输入可追踪,训练可复现,验证可解释,推理与训练保持同一套数据处理规则。

【声明】本内容来自华为云开发者社区博主,不代表华为云及华为云开发者社区的观点和立场。转载时必须标注文章的来源(华为云社区)、文章链接、文章作者等基本信息,否则作者和本社区有权追究责任。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱: cloudbbs@huaweicloud.com
  • 点赞
  • 收藏
  • 关注作者

评论(0

0/1000
抱歉,系统识别当前为高风险访问,暂不支持该操作

全部回复

上滑加载中

设置昵称

在此一键设置昵称,即可参与社区互动!

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。