ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

一套模板搞定AlexNet/VGG/ResNet/ViT图像分类训练与部署

一套模板搞定AlexNet/VGG/ResNet/ViT图像分类训练与部署 这次我们来看一套可以直接拿走的深度学习图像分类代码模板。核心就一件事用同一套训练、验证、导出、部署代码无缝切换 AlexNet、VGG、ResNet、ViT 这四类网络而不需要每次换模型都重写一套训练流程。对于经常要在 CIFAR、ImageNet 子集或者自定义数据集上反复做 baseline 实验的同学来说这类模板最大的价值不是某个模型调得有多好而是把整个训练流程固定下来数据加载、增强策略、模型选择、训练循环、日志记录、权重保存、导出部署全部标准化。后面论文复现、课程作业、竞赛 baseline、项目预演都可以在这个框架上直接改不用重复造轮子。本文我会带你过一遍这套模板的核心思路和完整代码包括四类模型如何统一在一个 get_model 接口里、CIFAR-10 上的训练验证怎么做、训练曲线怎么看、模型怎么导出成 ONNX、以及怎么用 FastAPI 包一层推理服务。最后给一份常见问题排查清单和工程化建议。如果你正在学 CNN或者手里正好有一批图片要训练分类模型这篇文章可以直接收藏。1. 核心能力速览能力项说明模型覆盖AlexNet、VGG、ResNet、ViT可通过配置文件切换代码结构单套模板统一训练/验证/导出/推理流程任务类型图像分类预训练支持支持加载 ImageNet 预训练权重也支持随机初始化从零训练输入尺寸AlexNet / VGG / ResNet 常用 224x224ViT 常用 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 到注意力机制的跨越ViTVision 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 均可Python3.8 及以上深度学习框架PyTorch 1.13 或 2.xCPU 版或 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、ResNettimm 负责 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(weightsmodels.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](weightsweights_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](weightsweights_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, pretrainedpretrained, num_classesnum_classes) elif model_name vit_small: model timm.create_model(vit_small_patch16_224, pretrainedpretrained, num_classesnum_classes) elif model_name vit_large: model timm.create_model(vit_large_patch16_224, pretrainedpretrained, num_classesnum_classes) else: raise ValueError(fUnknown ViT variant: {model_name}) return model raise ValueError(fUnsupported model: {model_name}) if __name__ __main__: # 快速验证四个模型都能正常构建 for name in [alexnet, vgg11, resnet18, vit_base]: net get_model(name, num_classes10, pretrainedFalse) x torch.randn(1, 3, 224, 224) y net(x) print(f{name:12s} - output shape: {tuple(y.shape)})这段代码里最关键的部分是全连接层替换。AlexNet 和 VGG 的最后一层是classifier[6]ResNet 是fcViT 通过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, descTraining): 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(dim1) 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, descValidation): 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(dim1) labels).sum().item() total_samples images.size(0) return total_loss / total_samples, total_correct / total_samples def main(): parser argparse.ArgumentParser(descriptionUniversal CNN/ViT Training Template) parser.add_argument(--model, typestr, defaultresnet18, helpalexnet/vgg11/resnet18/vit_base) parser.add_argument(--dataset, typestr, defaultcifar10, helpcifar10 or path to custom image folder) parser.add_argument(--data_dir, typestr, default./data, helpdataset root) parser.add_argument(--epochs, typeint, default30) parser.add_argument(--batch_size, typeint, default64) parser.add_argument(--lr, typefloat, default1e-3) parser.add_argument(--num_classes, typeint, default10) parser.add_argument(--pretrained, actionstore_true, defaultTrue) parser.add_argument(--device, typestr, defaultcuda) parser.add_argument(--output_dir, typestr, default./runs) args parser.parse_args() os.makedirs(args.output_dir, exist_okTrue) device torch.device(args.device if torch.cuda.is_available() else cpu) print(fUsing device: {device}) train_transform, val_transform get_transform(args.model) if args.dataset cifar10: train_ds datasets.CIFAR10(rootargs.data_dir, trainTrue, downloadTrue, transformtrain_transform) val_ds datasets.CIFAR10(rootargs.data_dir, trainFalse, downloadTrue, transformval_transform) else: # 自定义目录结构: data_dir/train/class1/*.jpg, data_dir/val/class1/*.jpg train_ds datasets.ImageFolder(rootos.path.join(args.data_dir, train), transformtrain_transform) val_ds datasets.ImageFolder(rootos.path.join(args.data_dir, val), transformval_transform) train_loader DataLoader(train_ds, batch_sizeargs.batch_size, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_sizeargs.batch_size, shuffleFalse, num_workers4, pin_memoryTrue) model get_model(args.model, num_classesargs.num_classes, pretrainedargs.pretrained) model model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lrargs.lr, weight_decay5e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxargs.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(fEpoch {epoch:03d} | Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | fVal 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(fBest 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 导出 ONNXimport torch from models import get_model num_classes 10 model get_model(resnet18, num_classesnum_classes, pretrainedFalse) model.load_state_dict(torch.load(runs/resnet18_best.pth, map_locationcpu)) 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_version17, ) 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 filetest.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 memorybatch 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-3ViT 从零训练建议谨慎小数据集不加载预训练很难收敛。第四权重命名里带上模型名、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 换成自监督预训练版本。模板的价值就在这里当你需要验证一个新想法时不用从零搭训练代码直接在这个框架上加模块就行。建议收藏备用。
返回列表