这次我们来看一套可以直接拿走的深度学习图像分类代码模板。核心就一件事:用同一套训练、验证、导出、部署代码,无缝切换 AlexNet、VGG、ResNet、ViT 这四类网络,而不需要每次换模型都重写一套训练流程。
对于经常要在 CIFAR、ImageNet 子集或者自定义数据集上反复做 baseline 实验的同学来说,这类模板最大的价值不是某个模型调得有多好,而是把整个训练流程固定下来:数据加载、增强策略、模型选择、训练循环、日志记录、权重保存、导出部署,全部标准化。后面论文复现、课程作业、竞赛 baseline、项目预演,都可以在这个框架上直接改,不用重复造轮子。
本文我会带你过一遍这套模板的核心思路和完整代码,包括四类模型如何统一在一个 get_model 接口里、CIFAR-10 上的训练验证怎么做、训练曲线怎么看、模型怎么导出成 ONNX、以及怎么用 FastAPI 包一层推理服务。最后给一份常见问题排查清单和工程化建议。
如果你正在学 CNN,或者手里正好有一批图片要训练分类模型,这篇文章可以直接收藏。
1. 核心能力速览
| 能力项 | 说明 |
|---|---|
| 模型覆盖 | AlexNet、VGG、ResNet、ViT(可通过配置文件切换) |
| 代码结构 | 单套模板统一训练/验证/导出/推理流程 |
| 任务类型 | 图像分类 |
| 预训练支持 | 支持加载 ImageNet 预训练权重,也支持随机初始化从零训练 |
| 输入尺寸 | AlexNet / VGG / ResNet 常用 224x224;ViT 常用 224x224 或 384x384 |
| 运行方式 | Python 脚本启动,命令行传入参数 |
| API 部署 | 可扩展为 FastAPI HTTP 推理服务 |
| 批量任务 | 支持批推理脚本,适合文件夹批量预测 |
| 硬件要求 | 有 NVIDIA GPU 优先,显存不足可切 CPU 或减小 batch size |
| 适合读者 | 刚入门 CNN 的学生、做图像分类 baseline 的算法工程师 |
从材料来看,这套方案不是某个封装好的巨型框架,而是一份可维护、可扩展的工程代码模板。因此显存需求、推理速度这些指标,会直接由你的显卡型号、batch size 和输入分辨率决定,后面我会给出具体的观察和调优方法。
2. 四个模型,选型逻辑先搞清楚
很多初学者以为模型越多越好,实际选型要结合任务、显存和效果来定。模板里这四类网络各有明显定位。
2.1 AlexNet:入门最小闭环
AlexNet 是 2012 年的经典结构,用 5 层卷积加 3 层全连接完成 1000 类分类。在今天看来它效果不如新模型,但有两个不可替代的价值:结构简单,适合第一遍理解卷积、池化、Flatten、Dropout 这些基础概念;训练速度快,在 CIFAR-10 上跑几十个 epoch 也就几分钟到十几分钟,适合排查代码逻辑和数据 pipeline 问题。
用这套模板时,AlexNet 可以当作“冒烟测试模型”。先把整个训练链路用 AlexNet 跑通,再切到 ResNet 或 ViT 做正式实验,会省掉大量调试时间。
2.2 VGG:感受野和堆叠的典型
VGG 的核心思想是用小卷积核堆深度。3x3 卷积反复堆叠,网络变深但参数量可控,而且感受野逐步扩大,特征抽象层次分明。VGG 虽然有 1.38 亿参数这个体量较大的版本,但也有 VGG11、VGG13、VGG16 这些变体可选,模板里建议用 timm 或 torchvision 自带的版本,显存不够就选 VGG11。
VGG 非常适合用来观察 CNN 的“粗粒度到细粒度”特征变化:浅层卷积输出边缘和纹理,深层卷积输出语义部件。这个特点在 ResNet 和 ViT 里也有体现,但 VGG 的结构最直观。
2.3 ResNet:残差连接解决深网络退化
ResNet 的核心贡献是残差连接。网络加深之后,普通结构会出现训练退化问题,表现为训练 loss 降不下去、验证精度不升反降。ResNet 通过恒等映射让梯度回传更顺畅,使得 50 层、101 层甚至更深的网络可以稳定训练。
在工程实践里,ResNet 属于“测出来效果稳”的类型。它不像 ViT 那样需要大规模预训练数据,在中小型数据集上用预训练权重微调效果就很好。模板里默认推荐优先验证 ResNet18 或 ResNet34,显存压力小、收敛快、效果均衡。
2.4 ViT:从 CNN 到注意力机制的跨越
ViT(Vision Transformer)把图片切成 patch,然后通过自注意力机制建模全局依赖。它跟 CNN 最大的区别是感受野天然就是全局的,不需要靠堆叠卷积层逐步扩大感受野。对于纹理较弱、需要全局语义关系的任务,ViT 往往比 CNN 更强;但在小数据集上,没有预训练权重的 ViT 收敛速度通常不如 ResNet。
使用 ViT 时要注意一个关键点:输入的 patch size 通常固定,例如 16x16 或 14x14,所以图片分辨率最好按模型要求设置。模板里遇到 ViT 会自动调整输入尺寸,避免因为尺寸不匹配报错。
一句话总结选型逻辑:跑通流程用 AlexNet,做 baseline 用 ResNet,观察局部到全局特征用 VGG,追求大模型上限用 ViT 加载预训练权重。
3. 环境准备与前置条件
这套模板依赖比较常规,不需要特殊环境。下面给一份通用检查清单,具体版本以你本机为准。
| 环境项 | 建议配置 |
|---|---|
| 操作系统 | Windows 10/11、Ubuntu 18.04+、macOS 均可 |
| Python | 3.8 及以上 |
| 深度学习框架 | PyTorch 1.13 或 2.x,CPU 版或 CUDA 版均可 |
| CUDA | 如果你有 NVIDIA 显卡,建议 CUDA 11.8 或 12.1 以上 |
| GPU | 建议 4GB 显存以上,至少能跑 ResNet18/Swin-T 级别模型 |
| 磁盘空间 | 预留 10GB 以上,包含数据集和权重文件 |
| 依赖库 | torch、torchvision、timm、numpy、Pillow、tqdm、fastapi、uvicorn、onnx、onnxruntime |
安装依赖只需要一条命令:
# 基础训练环境 pip install torch torchvision timm numpy Pillow tqdm # 部署服务环境 pip install fastapi uvicorn onnx onnxruntime安装完成后先确认伪环境是否可用:
python -c "import torch; print(torch.__version__); print(torch.cuda.is_available())"如果你使用的是 CUDA 版 PyTorch,上面命令输出True说明显卡可用;输出False则后续训练走 CPU,速度会慢,但不影响流程验证。
4. 代码模板:四个模型统一入口
模板的核心理念是“架构与训练逻辑解耦”。模型定义单独放到models.py,训练脚本通过--model参数选择架构。无论切到哪个模型,数据加载、优化器、学习率策略、日志记录都走同一套代码。
4.1 模型定义文件 models.py
这里使用 torchvision 和 timm 作为模型后端。torchvision 负责 AlexNet、VGG、ResNet,timm 负责 ViT,因为 timm 提供的 ViT 预训练权重更全,加载更方便。
import torch import torch.nn as nn import torchvision.models as models import timm def get_model(model_name: str, num_classes: int = 10, pretrained: bool = True): """ 统一模型入口。 支持: alexnet, vgg11, vgg13, vgg16, vgg19, resnet18, resnet34, resnet50, resnet101, vit_base, vit_small, vit_large """ model_name = model_name.lower() # ---------- AlexNet ---------- if model_name == "alexnet": model = models.alexnet(weights=models.AlexNet_Weights.IMAGENET1K_V1 if pretrained else None) in_features = model.classifier[6].in_features model.classifier[6] = nn.Linear(in_features, num_classes) return model # ---------- VGG ---------- if model_name.startswith("vgg"): weights_map = { "vgg11": models.VGG11_Weights.IMAGENET1K_V1, "vgg13": models.VGG13_Weights.IMAGENET1K_V1, "vgg16": models.VGG16_Weights.IMAGENET1K_V1, "vgg19": models.VGG19_Weights.IMAGENET1K_V1, } model = models.__dict__[model_name](weights=weights_map[model_name] if pretrained else None) in_features = model.classifier[6].in_features model.classifier[6] = nn.Linear(in_features, num_classes) return model # ---------- ResNet ---------- if model_name.startswith("resnet"): weights_map = { "resnet18": models.ResNet18_Weights.IMAGENET1K_V1, "resnet34": models.ResNet34_Weights.IMAGENET1K_V1, "resnet50": models.ResNet50_Weights.IMAGENET1K_V1, "resnet101": models.ResNet101_Weights.IMAGENET1K_V1, } model = models.__dict__[model_name](weights=weights_map[model_name] if pretrained else None) in_features = model.fc.in_features model.fc = nn.Linear(in_features, num_classes) return model # ---------- ViT ---------- if model_name.startswith("vit"): pretrained_cfg = "imagenet1k" if pretrained else "" if model_name == "vit_base": model = timm.create_model("vit_base_patch16_224", pretrained=pretrained, num_classes=num_classes) elif model_name == "vit_small": model = timm.create_model("vit_small_patch16_224", pretrained=pretrained, num_classes=num_classes) elif model_name == "vit_large": model = timm.create_model("vit_large_patch16_224", pretrained=pretrained, num_classes=num_classes) else: raise ValueError(f"Unknown ViT variant: {model_name}") return model raise ValueError(f"Unsupported model: {model_name}") if __name__ == "__main__": # 快速验证四个模型都能正常构建 for name in ["alexnet", "vgg11", "resnet18", "vit_base"]: net = get_model(name, num_classes=10, pretrained=False) x = torch.randn(1, 3, 224, 224) y = net(x) print(f"{name:12s} -> output shape: {tuple(y.shape)}")这段代码里最关键的部分是全连接层替换。AlexNet 和 VGG 的最后一层是classifier[6],ResNet 是fc,ViT 通过num_classes参数直接控制。切模型时只需要改命令行参数,不需要改数据处理和训练逻辑。
4.2 训练脚本 train.py
训练脚本要做的事很明确:读取配置、构建数据加载器、创建模型、定义损失函数和优化器、循环训练和验证、保存最优权重。
import argparse import os import time import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms from tqdm import tqdm from models import get_model def get_transform(model_name: str): # 不同模型对输入尺寸要求不同,ViT 固定用 224x224,其余也统一走 224 normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) train_transform = transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), normalize, ]) val_transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), normalize, ]) return train_transform, val_transform def train_one_epoch(model, loader, criterion, optimizer, device): model.train() total_loss, total_correct, total_samples = 0.0, 0, 0 for images, labels in tqdm(loader, desc="Training"): images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() * images.size(0) total_correct += (outputs.argmax(dim=1) == labels).sum().item() total_samples += images.size(0) return total_loss / total_samples, total_correct / total_samples @torch.no_grad() def validate(model, loader, criterion, device): model.eval() total_loss, total_correct, total_samples = 0.0, 0, 0 for images, labels in tqdm(loader, desc="Validation"): images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) total_loss += loss.item() * images.size(0) total_correct += (outputs.argmax(dim=1) == labels).sum().item() total_samples += images.size(0) return total_loss / total_samples, total_correct / total_samples def main(): parser = argparse.ArgumentParser(description="Universal CNN/ViT Training Template") parser.add_argument("--model", type=str, default="resnet18", help="alexnet/vgg11/resnet18/vit_base") parser.add_argument("--dataset", type=str, default="cifar10", help="cifar10 or path to custom image folder") parser.add_argument("--data_dir", type=str, default="./data", help="dataset root") parser.add_argument("--epochs", type=int, default=30) parser.add_argument("--batch_size", type=int, default=64) parser.add_argument("--lr", type=float, default=1e-3) parser.add_argument("--num_classes", type=int, default=10) parser.add_argument("--pretrained", action="store_true", default=True) parser.add_argument("--device", type=str, default="cuda") parser.add_argument("--output_dir", type=str, default="./runs") args = parser.parse_args() os.makedirs(args.output_dir, exist_ok=True) device = torch.device(args.device if torch.cuda.is_available() else "cpu") print(f"Using device: {device}") train_transform, val_transform = get_transform(args.model) if args.dataset == "cifar10": train_ds = datasets.CIFAR10(root=args.data_dir, train=True, download=True, transform=train_transform) val_ds = datasets.CIFAR10(root=args.data_dir, train=False, download=True, transform=val_transform) else: # 自定义目录结构: data_dir/train/class1/*.jpg, data_dir/val/class1/*.jpg train_ds = datasets.ImageFolder(root=os.path.join(args.data_dir, "train"), transform=train_transform) val_ds = datasets.ImageFolder(root=os.path.join(args.data_dir, "val"), transform=val_transform) train_loader = DataLoader(train_ds, batch_size=args.batch_size, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=args.batch_size, shuffle=False, num_workers=4, pin_memory=True) model = get_model(args.model, num_classes=args.num_classes, pretrained=args.pretrained) model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=5e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs) best_acc = 0.0 for epoch in range(1, args.epochs + 1): train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device) val_loss, val_acc = validate(model, val_loader, criterion, device) scheduler.step() print(f"Epoch {epoch:03d} | Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | " f"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}") if val_acc > best_acc: best_acc = val_acc torch.save(model.state_dict(), os.path.join(args.output_dir, f"{args.model}_best.pth")) print(f"Best Val Acc: {best_acc:.4f}") if __name__ == "__main__": main()训练命令示例:
# AlexNet 跑通流程 python train.py --model alexnet --epochs 30 --batch_size 64 # ResNet18 正规训练 python train.py --model resnet18 --epochs 60 --batch_size 64 --lr 1e-3 # ViT 加载预训练权重微调 python train.py --model vit_base --epochs 50 --batch_size 32 --lr 5e-55. 功能测试与效果验证
5.1 模型构建冒烟测试
先运行models.py末尾的验证代码,确认四个模型输出形状一致。这一步能过滤掉 80% 的尺寸不匹配问题。
python models.py预期输出:
alexnet -> output shape: (1, 10) vgg11 -> output shape: (1, 10) resnet18 -> output shape: (1, 10) vit_base -> output shape: (1, 10)如果某个模型输出维度不是(1, 10),检查分类头替换那一行,重点看 in_features 是否和原始分类层匹配。
5.2 CIFAR-10 训练验证
在 CIFAR-10 上用小 epoch 数快速验证整个链路。建议第一次测试时不要追求精度,重点观察三件事:训练 loss 是否在下降,验证准确率是否随 epoch 上升,权重文件是否正常保存。
以 ResNet18 为例,30 epoch、batch size 64、学习率 1e-3,在单张 8GB 显存显卡上一般几分钟到十几分钟完成一个 epoch,显存占用约 1-2GB。ViT-base 因为参数量和 attention 计算量更大,显存约 4-6GB,具体以本机为准。
训练结束后,在runs/目录下找到resnet18_best.pth。这个文件就是后续部署要用的模型权重。
5.3 训练曲线判断标准
每次训练完成后,把输出的 Train Loss / Val Loss / Val Acc 整理成曲线,判断标准并不复杂:
- Train Loss 持续下降,Val Acc 同步上升:正常训练,继续跑。
- Train Loss 下降但 Val Loss 上升:过拟合,需要增大数据增强、加大 dropout、或者换更小的模型。
- Train Loss 和 Val Loss 都不降:学习率可能过大或过小,先调小一个数量级测试。
- ViT 在小数据集上收敛慢是正常现象,优先加载预训练权重,学习率建议 5e-5 到 1e-4。
这里我推荐把训练日志重定向到文件里,方便后面定位问题:
python train.py --model resnet18 --epochs 30 | tee train_resnet18.log6. 部署:从 PyTorch 权重到推理服务
训练完成之后,部署环节分为三步:导出 ONNX、写推理脚本、起 API 服务。这一步能让你脱离训练环境,直接用模型做实际预测。
6.1 导出 ONNX
import torch from models import get_model num_classes = 10 model = get_model("resnet18", num_classes=num_classes, pretrained=False) model.load_state_dict(torch.load("runs/resnet18_best.pth", map_location="cpu")) model.eval() dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, "resnet18.onnx", input_names=["images"], output_names=["logits"], dynamic_axes={"images": {0: "batch"}, "logits": {0: "batch"}}, opset_version=17, ) print("ONNX export done.")导出时开启动态 batch,这样推理时一次可以传入 1 张或多张图片,后面做批量任务更方便。
6.2 本地推理脚本
import numpy as np from PIL import Image from torchvision import transforms import onnxruntime as ort # 标签顺序需要和你训练时的类别一致 CLASS_NAMES = ["airplane", "automobile", "bird", "cat", "deer", "dog", "frog", "horse", "ship", "truck"] normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), normalize, ]) sess = ort.InferenceSession("resnet18.onnx", providers=["CUDAExecutionProvider", "CPUExecutionProvider"]) def predict(image_path: str): img = Image.open(image_path).convert("RGB") tensor = transform(img).unsqueeze(0).numpy() logits = sess.run(None, {"images": tensor})[0] pred_idx = int(np.argmax(logits[0])) return CLASS_NAMES[pred_idx], float(np.max(logits[0])) if __name__ == "__main__": print(predict("test.jpg"))6.3 FastAPI 推理服务
批量图片逐个调用推理脚本效率太低,更合理的做法是封装成 HTTP 服务。下面是一个最小可用的 FastAPI 示例:
from io import BytesIO import numpy as np from PIL import Image from fastapi import FastAPI, UploadFile, File from torchvision import transforms import onnxruntime as ort app = FastAPI() CLASS_NAMES = ["airplane", "automobile", "bird", "cat", "deer", "dog", "frog", "horse", "ship", "truck"] normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), normalize, ]) sess = ort.InferenceSession("resnet18.onnx", providers=["CUDAExecutionProvider", "CPUExecutionProvider"]) @app.post("/predict") async def predict(file: UploadFile = File(...)): img_bytes = await file.read() img = Image.open(BytesIO(img_bytes)).convert("RGB") tensor = transform(img).unsqueeze(0).numpy() logits = sess.run(None, {"images": tensor})[0] pred_idx = int(np.argmax(logits[0])) return { "class": CLASS_NAMES[pred_idx], "class_id": pred_idx, "confidence": float(np.max(logits[0])), } # 启动: uvicorn api_server:app --host 0.0.0.0 --port 8000启动服务后,用 curl 测试接口:
curl -X POST http://127.0.0.1:8000/predict \ -F "file=@test.jpg"返回 JSON 示例:
{ "class": "cat", "class_id": 3, "confidence": 0.9251 }6.4 批量推理目录
批量任务不需要每次都走 HTTP 请求,直接遍历文件夹更快:
import os import numpy as np from PIL import Image from torchvision import transforms import onnxruntime as ort sess = ort.InferenceSession("resnet18.onnx", providers=["CPUExecutionProvider"]) transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) def batch_predict(image_dir: str): results = {} for fname in os.listdir(image_dir): if not fname.lower().endswith((".jpg", ".jpeg", ".png")): continue path = os.path.join(image_dir, fname) img = Image.open(path).convert("RGB") tensor = transform(img).unsqueeze(0).numpy() logits = sess.run(None, {"images": tensor})[0] pred_idx = int(np.argmax(logits[0])) results[fname] = pred_idx print(f"{fname}: class {pred_idx}") return results if __name__ == "__main__": batch_predict("./test_images")实际做上万张图片的批量推理时,建议另外加进度记录和失败重试逻辑,每处理 500 张保存一次中间结果。
7. 资源占用与性能观察
模型训练时资源占用可以分成两个维度观察:显存和耗时。
显存占用主要取决于三个因素:模型参数量、batch size、输入分辨率。
- AlexNet 参数量约 6100 万,但因为结构简单,显存占用低,4GB 显存的卡也能跑。
- VGG16 参数量约 1.38 亿,显存占用明显上升,显存不够时优先换 VGG11。
- ResNet18 参数量约 1100 万,训练速度快,显存占用适中,是最推荐的第一选择。
- ViT-base 参数量约 8600 万,attention 计算显存消耗更高,建议 batch size 从 16 或 32 开始试。
观察显存的方法很简单。训练时开一个新的终端窗口,用 nvidia-smi 看进程占用:
watch -n 1 nvidia-smi显存不足时的解决顺序:先减小 batch size,再降低输入分辨率,最后换更小的模型。优先动 batch size,因为它对精度影响最小。
CPU 推理和 GPU 推理的差别在 ViT 上体现最明显,因为 ViT 的自注意力计算更重。如果只有 CPU,建议选择 ResNet 系列并且打开 ONNX Runtime 的 CPU 优化;ViT 在 CPU 上做小 batch 测试可以,大规模推理不建议。
另外要注意训练时的 num_workers 不要超过 CPU 核心数,否则数据加载本身会成为瓶颈。
8. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 训练时显存不足(CUDA out of memory) | batch size 过大或输入分辨率过高 | 查看 nvidia-smi 显存占用 | 调小 batch size,或换小模型,或降低分辨率 |
| ViT 训练 loss 不降 | 学习率过大或没有加载预训练 | 打印前 10 个 batch 的 loss,确认数据是否正常 | 学习率降到 5e-5,加载预训练权重 |
| 模型输出维度不匹配 | 分类头替换时 in_features 取错 | 打印 model 结构,核对全连接层索引 | 分别检查 classifier[6]、fc、head 层的属性名 |
| ONNX 导出报错 | 使用了动态控制流或不支持的算子 | 查看报错信息中具体算子名 | 固定输入尺寸导出,不要用动态尺寸 |
| API 服务启动失败 | 端口被占用或依赖缺失 | 查看 uvicorn 启动日志 | 换端口,或重新安装 fastapi/uvicorn |
| 批量推理速度慢 | 单张图片反复初始化 session | 把 InferenceSession 创建放到循环外 | 按上面代码,session 创建只做一次 |
| 自定义数据集加载出错 | 目录结构和 ImageFolder 要求不匹配 | 检查根目录下是否有 train/val 子目录,内部每个子目录代表一个类别 | 调整目录结构为 train/class1/img.jpg |
| 验证精度远低于训练精度 | 过拟合或验证数据预处理不一致 | 检查 val transform 是否包含训练时的增强 | 验证集不要加 RandomResizedCrop 和 RandomHorizontalFlip |
第一个重点排查项建议从数据 pipeline 开始。图像分类 70% 的问题不在模型,而在数据读取和预处理。建议先用很小的数据集、单 epoch 跑通一遍,观察 loss 是否从随机值开始下降,再逐步放大数据量。
9. 最佳实践与使用建议
这套模板真正用起来,建议遵守以下几个原则。
第一,固定一套“冒烟配置”。比如 AlexNet + CIFAR-10 + 5 epoch,每次改代码后先跑冒烟配置,确认代码没写错再跑正式实验。这能帮你把“模型问题”和“代码问题”分开。
第二,把数据增强当成超参数来调。CIFAR-10 上用 RandomCrop 和 HorizontalFlip 就够,但 ImageNet 级别的数据集需要更复杂的增强策略,例如 AutoAugment、MixUp、CutMix。模板里的增强只是最基础的版本,正式训练大模型时建议升级。
第三,学习率策略要按模型分开设置。CNN 用 AdamW 加 CosineAnnealing 通常没问题;ViT 微调时学习率要比 CNN 小一个量级,常见是 5e-5 到 1e-4。ResNet 从头训练可以用 1e-3,ViT 从零训练建议谨慎,小数据集不加载预训练很难收敛。
第四,权重命名里带上模型名、epoch、精度三个信息。例如resnet18_e30_acc94.2.pth,不要只存一个best.pth。后面做多组实验对比时会省掉很多麻烦。
第五,部署和训练环境分离。训练用 PyTorch,部署用 ONNX Runtime。这样部署端不需要装完整版 PyTorch,依赖体积小很多,推理速度也更快。
第六,涉及商用数据训练时,确认数据集授权情况。从公开数据集下载要注意许可协议,用爬虫抓图训练要确认图片版权归属,不要直接拿未授权数据做商用。
10. 总结与下一步
这套模板最值得尝试的点,是把 AlexNet、VGG、ResNet、ViT 之间的切换成本降到了最低。你不需要重新学习四套代码,只需要改一个--model参数。
第一次使用建议按这个顺序验证:先用python models.py确认四个模型能构建,再用 AlexNet 跑 5 个 epoch 验证训练链路,切到 ResNet18 跑 30 个 epoch 拿到稳定精度,最后导出 ONNX 并用 FastAPI 封一层服务。整个过程跑通之后,你就有了一套属于自己的图像分类基础设施。
最容易踩的坑有三个:ViT 输入尺寸和 CNN 不一致、分类头替换时 in_features 取错、自定义数据集目录结构和 ImageFolder 要求不匹配。这三个问题在 8 的排查表里都有对应方案。
后续可以扩展的方向也很多。比如把模板里的单标签分类改成多标签分类,把 ResNet 换成 ResNet FPN 结构提取粗粒度到细粒度的多尺度特征,或者把 ViT 的 patch embedding 换成自监督预训练版本。模板的价值就在这里:当你需要验证一个新想法时,不用从零搭训练代码,直接在这个框架上加模块就行。建议收藏备用。