尧图精选

PyTorch CNN图像分类系统实战:从原理到部署的完整指南

🕒 发布时间:2026/10/1 3:03:34 📁 来源:尧图网络
简介一套面向毕业设计场景的Python卷积神经网络图像分类系统完整资料涵盖CNN原理讲解、模型实现与训练文档适合正在学习深度学习、需要快速上手图像分类项目的高校学生。压缩包共25个文件包含13个Python脚本、模型备份文件、编译缓存、数据集与训练好的权重以及Markdown和JSON格式的说明与配置文档整体仅62KB结构紧凑。已有64人学习使用。资料内提供LeNet-5、AlexNet、GoogLeNet、ResNet等经典网络的实现可对比TensorFlow与PyTorch两种框架的写法并配有README、模型调用入口、辅助脚本等方便读者理解从数据预处理到模型训练、评估的完整流程。借助源码与文档既能支撑毕业设计的代码复现也能帮助系统掌握CNN在图像分类中的实际应用。1. 图像分类系统从源码到可用模型先别急着敲训练命令手上拿到一套「Python 实现 CNN 卷积神经网络的图像分类系统」的源码和模型文档资料最常见的反应是打开终端直接跑训练脚本然后被报错、版本、路径问题劝退。这套东西本质上解决一个问题让你用 Python 完成从图片数据到分类模型的完整闭环中间覆盖数据整理、卷积网络搭建、训练调参、评估导出以及文档里那些只能靠踩坑才能理解的超参数。它适合刚好有点 Python 基础、第一次接触 CNN 图像分类的读者不适合只想抄一条命令就跑出 99% 准确率的人因为图像分类的坑全部藏在数据和参数里不在那一行命令里。2. CNN 图像分类的原理与模型选型为什么卷积核能学会“看”图像图像分类系统选什么网络结构取决于你手里有多少数据、多少显存、推理要跑在 CPU 还是 GPU 上。在动手改源码前先把 CNN 的基本结构过一遍后面调参才不会被一堆层名绕晕。2.1 卷积层在做什么局部感受野与参数共享图像在程序里是一个多维张量尺寸 224×224×3 的 RGB 图放进全连接网络会变成 150528 个输入节点第一层隐层如果也是 1024 个节点光这一层就有超过 1.5 亿个权重参数。这个规模在训练时既跑不动也极容易过拟合。卷积网络换了一种思路用一个小尺寸的卷积核比如 3×3 或 5×5在图像上滑动每次只看局部一个小区域这就是局部感受野。卷积核的权重在整张图上共享同一个核提取同一种特征比如水平边缘、颜色块、纹理方向。这样参数量从亿级降到几千级这也是 CNN 能在图像任务上站稳的基础。输出特征图的尺寸由输入尺寸、卷积核大小、padding 和 stride 共同决定公式是output floor((H 2*padding - kernel_size) / stride) 1用代码定义一个卷积层并验证输出尺寸是最快建立手感的方式import torch.nn as nn import torch # 输入: 1 张 224x224 的 RGB 图像 - [1, 3, 224, 224] x torch.randn(1, 3, 224, 224) # 3x3 卷积, padding1, stride1 时输出尺寸不变 conv1 nn.Conv2d(in_channels3, out_channels16, kernel_size3, padding1, stride1) y conv1(x) # [1, 16, 224, 224] print(y.shape)这里 in_channels3 对应 RGB 三个通道out_channels16 表示用 16 个卷积核得到 16 张特征图。kernel_size 越小局部感受野越小但可以通过堆叠多层扩大感受野padding1 是为了让边界像素也被卷积核覆盖。out_channels 是网络宽度的重要旋钮从 16 加到 64 会显著增加参数量和显存开销。2.2 池化与激活下采样与引入非线性卷积层是线性的滑动窗口加权求和多个卷积层叠在一起如果中间没有激活函数整个网络仍然是线性的表达不了复杂映射。现在一般用 ReLU 做激活f(x)max(0,x)计算简单梯度在正区间恒为 1能缓解深层网络的梯度消失。在 CNN 基础里激活层和卷积层通常成对出现“卷积 - BN - ReLU”几乎是所有现代主干的标准写法。池化层负责下采样最常见的是 MaxPooling。它的作用不是增加参数而是把 2×2 窗口里的最大值留下丢掉其他 75% 的信息换来平移不变性和更大的后续感受野。注意池化没有可学习参数但 kernel_size 和 stride 直接影响特征图缩小的比例比如 nn.MaxPool2d(2) 会让宽高各减半。源码里如果看到连续 3 个 224 输入每次池化后变成 112、56、28这就是一条清晰的下采样路径也是判断网络是否写到一半的思路。2.3 常见主干模型怎么选LeNet / AlexNet / ResNet 的适用边界搞图像分类系统通常有三种起点取决于数据量和算力模型深度参数量量级适用场景LeNet-55 层约 6 万MNIST 手写数字、教学演示AlexNet8 层约 6000 万百万级数据的开创者已少用ResNet-1818 层约 1100 万中小数据集迁移学习首选ResNet-5050 层约 2500 万数据量大、GPU 充裕时更好源码里如果自带训练好的模型文件先看它是什么结构的权重再决定怎么改分类头。常见做法是直接用 torchvision 的预训练 ResNet18 做迁移学习把最后一层全连接改成自己的类别数import torch.nn as nn import torchvision.models as models num_classes 10 # 按自己的类别数改 # weights 表示加载 ImageNet 上预训练权重替代旧版 pretrainedTrue model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # 替换最后一层分类头 model.fc nn.Linear(model.fc.in_features, num_classes)用迁移学习而不是从零训练是小数据集图像分类系统的默认方案。ImageNet 上训好的浅层特征边缘、颜色、纹理对大多数自然图像都通用需要重新学的只有后面的分类层。pretrained 参数在新版 torchvision 里推荐用 weights 方式指定直接写 pretrainedTrue 会收到弃用警告这个坑在模型文档资料里经常被忽略。如果你的源码版本旧用的还是 pretrainedTrue先看清版本再动手。挑选主干时还要考虑推理环境CPU 机器上 ResNet18 比 ResNet50 快一倍以上嵌入式设备可能更合适 MobileNet。深层模型在小数据集上不一定会更好反而更容易过拟合。先选一个跑得动的模型把全流程走通再换大模型这是最稳的推进节奏。3. 把源码跑起来数据准备、训练脚本与关键参数拿到源码后的第一件事是先把运行环境固定住。与其去翻 python 安装教程里那些版本选择细节不如直接装官方 Python 3.10 或 3.11然后为这个项目单独建一个虚拟环境避免把老项目和当前项目互相污染。环境稳定后再按数据、增强、训练、保存四步走。3.1 数据集目录约定与 ImageFolder 加载图像分类源码拿到手第一步永远不是改模型而是确认数据目录长什么样。PyTorch 的 ImageFolder 约定每个类别一个子文件夹训练数据按类别分目录这个约定一旦乱掉标签全部错位且报错不会告诉你哪里错。目录结构一般是data/ ├── train/ │ ├── cat/ # 所有猫的图片 │ ├── dog/ │ └── bird/ └── val/ ├── cat/ ├── dog/ └── bird/train 和 val 按同样类别划分。用 ImageFolder 加载from torchvision import datasets, transforms train_dataset datasets.ImageFolder( rootdata/train, transformtrain_transform, ) val_dataset datasets.ImageFolder( rootdata/val, transformval_transform, ) print(train_dataset.class_to_idx) # {cat: 0, dog: 1, bird: 2} print(train_dataset.classes)ImageFolder 会按文件夹名字母序分配标签索引也就是说 cat 是 0 不是因为它在数据里排第一而是因为它按字母序排第一。class_to_idx 这个字典一定要打印出来看一眼很多训练脚本里标签错位都是在这里悄悄发生的。如果数据不是按目录组织的而是一张 CSV 记录图片路径和类别就得自定义 Dataset常见做法是写一个类初始化时读 CSVgetitem里读图并返回样本和标签。训练集和验证集的划分最好在数据准备阶段就完成不要在训练代码里随机切否则每次跑出来的实验不可对比。没有现成 val 目录时用 sklearn 的 train_test_split 按 8:2 划分并保持每类比例stratify 参数是按标签分层抽样防止少数类别在验证集里消失。需要补数据时如果公开数据集不包含你的目标类别常见做法是用 python 爬虫去图片站抓一批图但这里提醒一句抓图前先确认版权和平台条款用来做学习 demo 问题不大商用要慎重这不是技术问题是风险问题。3.2 数据增强训练集和验证集为什么不能一样图像分类系统里最容易拉开差距的不是网络结构而是数据增强策略。训练集要做随机变换提升泛化性验证集只做尺寸调整和归一化保证评估结果稳定。如果验证集也做随机翻转每次评估的准确率都会抖动没法比较实验。from torchvision import transforms # 训练集增强 train_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomCrop(224), transforms.RandomHorizontalFlip(p0.5), transforms.ColorJitter(brightness0.2, contrast0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) # 验证集只做最小处理 val_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])Resize((256, 256)) 后再 RandomCrop(224) 是经典做法它比直接 Resize 到 224 多了一点随机裁剪的空间相当于免费的数据扩充。RandomHorizontalFlip 对左右对称的任务猫狗、车型有效但对文字识别这类左右不对称的任务要关掉。ColorJitter 的值不要给太大0.2 左右足够太大可能把颜色分布改到偏离真实场景。归一化的 mean 和 std用预训练模型时一定要用 ImageNet 这套标准值因为预训练权重是在这套归一化下训练出来的。如果用自己的数据从头训练mean/std 应该用训练集的统计值但多数图像分类系统都走迁移学习路线所以直接沿用 ImageNet 的数值。ToTensor 会把 0~255 的像素值缩放到 0~1Normalize 再按通道减均值除标准差这个顺序不能颠倒。3.3 训练循环损失、优化器、学习率调度训练循环是整套源码的核心。多数源码会有一个 train.py里面包含损失函数、优化器、学习率调度器和训练循环。一个能直接改的最小训练函数大约长这样import torch.optim as optim from tqdm import tqdm def train_one_epoch(model, dataloader, criterion, optimizer, device): model.train() running_loss, correct, total 0.0, 0, 0 for images, labels in tqdm(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() _, predicted outputs.max(1) total labels.size(0) correct predicted.eq(labels).sum().item() avg_loss running_loss / len(dataloader) acc 100.0 * correct / total return avg_loss, accoptimizer.zero_grad() 每步清空梯度backward 反传step 更新权重这三行顺序不能乱。outputs.max(1) 返回每个样本预测类别的值和索引predicted.eq(labels) 是逐一比较的布尔张量sum 后就是正确数。tqdm 只负责显示进度不影响结果不想装可以去掉。超参数的选择直接决定训练是否收敛下面是常见起点参数常用值说明batch_size16 / 32 / 64显存决定上限过小导致梯度震荡epochs30~50迁移学习通常不需要太多learning_rate1e-3 / 1e-4替换分类头用 1e-3微调主干用 1e-4weight_decay5e-4L2 正则抑制过拟合optimizerAdam 或 SGDAdam 上手快SGD 需要配 lr decay优化器选择上Adam 适合刚上手收敛快、对学习率不敏感SGD 配合 cos 或 step 学习率调度最终精度往往略高。常见做法是先把整体跑通用 Adam稳定后换 SGD momentum0.9。学习率调度有三种固定 lr、按 epoch 阶梯下降、余弦退火。阶梯下降里 milestone 一般设在总 epoch 的 1/2 和 3/4 处每次降为原来的 0.1。用 PyTorch 自带的 CosineAnnealingLR 也可以效果稳定且少一个需要调的数字。3.4 保存 checkpoint别只存 state_dict训练到一半服务器重启、GPU 被占、loss 跑飞想回退这些情况在图像分类项目里几乎都会遇到所以保存模型不能只存一个 state_dict要存完整 checkpoint。至少包含模型参数、优化器状态、当前 epoch、最佳验证准确率这样断点续训才有后悔药吃。import torch def save_checkpoint(state, filename): torch.save(state, filename) # 每个 epoch 结束后调用 best_acc 0.0 for epoch in range(1, total_epochs 1): train_loss, train_acc train_one_epoch(...) val_acc evaluate(model, val_loader, device) checkpoint { epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_acc: max(best_acc, val_acc), num_classes: num_classes, } if val_acc best_acc: best_acc val_acc save_checkpoint(checkpoint, checkpoints/best.pth) save_checkpoint(checkpoint, checkpoints/last.pth)best.pth 保留验证集表现最好的权重last.pth 保留 epoch 结束时权重两个文件的用途完全不同。断点续训时先 map_location 加载再恢复 optimizerckpt torch.load(checkpoints/last.pth, map_locationcpu) model.load_state_dict(ckpt[model_state_dict]) optimizer.load_state_dict(ckpt[optimizer_state_dict]) start_epoch ckpt[epoch] 1恢复 optimizer 状态很多人会漏漏掉之后 lr 会重新变回初始值前 10 个 epoch 的学习率调度就白做了。模型文档资料里如果提到“继续训练”却只给了 load_state_dict就要自己补上 optimizer 恢复。加载时推荐先 map_locationcpu 再 move 到 GPU省掉很多跨设备加载的报错。4. 模型评估与推理准确率之外还要看什么训练结束不等于项目结束。图像分类系统的交付物应该是可复现的评估结论和可用的推理接口而不是一个只会打印训练 acc 的脚本。这一章的评估函数、混淆矩阵和推理脚本几乎每个项目都要用到。4.1 验证集与测试集评估函数怎么写得稳训练完的模型不能只看训练集准确率。图像分类系统里通常有两条独立数据路径验证集用于挑模型和调参测试集用于最后一次性评估。如果一直在测试集上试来试去测试集就变成了验证集最后报告的指标会虚高。这个道理几乎所有文档资料都会写但源码里常见的问题是只有 val 没有 test这时至少要留出一部分 val 数据不参与调参。评估函数比训练函数简单核心是 model.eval() 和 torch.no_grad()def evaluate(model, dataloader, device): model.eval() correct, total 0, 0 with torch.no_grad(): for images, labels in dataloader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted outputs.max(1) total labels.size(0) correct predicted.eq(labels).sum().item() acc 100.0 * correct / total return accwith torch.no_grad() 关闭梯度计算推理速度和显存占用都更友好。model.eval() 会切换 BN 层和 Dropout 的行为训练时 BN 用 batch 统计量评估时用全局统计量Dropout 在评估时被关闭。只调用 evaluate 不写 model.eval() 是很多源码翻车的原因BN 层行为不一致会导致验证准确率忽高忽低。验证集准确率如果要和训练集对比还要统一输入预处理和 batch_size 的影响。batch_size 大时 BN 统计更稳定换 batch_size 后验证结果也可能轻微波动所以比较实验时尽量保持评估管线不变。4.2 混淆矩阵与每类召回率发现偏科总准确率 95% 的模型可能某一类只有 60% 召回率这类“偏科”在总指标里被掩盖了。图像分类系统的评估层面混淆矩阵比准确率更有诊断价值。用 sklearn 一行就能算from sklearn.metrics import confusion_matrix, classification_report import numpy as np def collect_preds(model, dataloader, device): all_preds, all_labels [], [] model.eval() with torch.no_grad(): for images, labels in dataloader: images images.to(device) outputs model(images) all_preds.extend(outputs.argmax(1).cpu().numpy()) all_labels.extend(labels.numpy()) return np.array(all_preds), np.array(all_labels) preds, labels collect_preds(model, val_loader, device) print(confusion_matrix(labels, preds)) print(classification_report(labels, preds, target_namesclass_names))confusion_matrix 的行是真实类别列是预测类别对角线越亮越好。看到某一类被大量分到邻近类通常是这两类在视觉上相似数据里样本数又不均衡。classification_report 里的 f1-score 是 precision 和 recall 的调和平均类别样本很少时只盯这一个指标就够了因为它不会被多数类的高准确率稀释。定位到偏科类别后常见做法有三种给少样本类别做过采样或复制用 class_weight 给损失函数里的少数类加权收集更多该类别的图片。前两种在源码里改动最小class_weight 的计算方式一般是 n_samples / (n_classes * class_count)传入 CrossEntropyLoss 的 weight 参数即可。4.3 单张图片推理与批量推理别重复加载模型训练完的模型最终要落到推理脚本。推理脚本的典型流程是加载 checkpoint、把模型送到设备、对单张图做预处理、前向、取 top-k。图片分类系统的推理脚本里最容易出现的坑是每张图都重复加载一遍模型或者把预处理写错导致输出结果和图片对不上。import torch import torch.nn as nn import torchvision.models as models from PIL import Image def load_model(checkpoint_path, num_classes, device): model models.resnet18(weightsNone) model.fc nn.Linear(model.fc.in_features, num_classes) ckpt torch.load(checkpoint_path, map_locationdevice) model.load_state_dict(ckpt[model_state_dict]) model.to(device).eval() return model def predict_one(model, image_path, transform, device, topk5): img Image.open(image_path).convert(RGB) x transform(img).unsqueeze(0).to(device) # [1, C, H, W] with torch.no_grad(): logits model(x) probs torch.softmax(logits, dim1) topk_probs, topk_idx probs.topk(topk, dim1) return [(idx.item(), prob.item()) for idx, prob in zip(topk_idx[0], topk_probs[0])]torch.softmax 把 logits 转成概率topk 取概率最高的前 k 个类别。要注意 Image.open 默认保留图片的原始色彩模式灰度图转 RGB 用 convert(RGB)否则通道数对不上预处理。推理时模型要放在同一个 device 上如果用 GPU 训练的 checkpoint 到 CPU 机器上跑load 时 map_locationcpu再传 device。批量推理时瓶颈往往在单张解码上常见做法是把一批图片文件做成 DataLoadernum_workers 开 2 到 4 个线程预读这样能跑满 GPU。文档资料里如果给出了推理帧率或单张耗时先确认这个数字是在什么硬件上测的CPU 和 GPU 能差一个数量级。4.4 模型导出pth、ONNX 与部署边界训练完整的模型要落地到业务通常要导出成部署格式。pth 是 PyTorch 训练专用部署服务大多不支持直接加载ONNX 是中间格式能被 ONNX Runtime、TensorRT 等推理引擎转换。图像分类系统的源码里建议保留一条导出脚本方便做跨平台推理。格式适合场景熟悉程度备注.pth/.pt继续训练、调试高依赖 PyTorch 环境.onnx服务端/边缘推理中跨框架、可转 TensorRT.pt 加 torch.jit.script纯 PyTorch 部署中保留动态控制流PyTorch 版本敏感导出 ONNX 的最小写法model.eval() dummy torch.randn(1, 3, 224, 224).to(device) torch.onnx.export( model, dummy, classifier.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}}, opset_version12, )dynamic_axes 声明 batch 维度可变这样导出的模型在推理时能接受任意 batch size。opset_version 决定了算子兼容性12 是保守选择新版推理引擎可能要求更高但设太高老的引擎读不了。导出后建议用 onnxruntime 快速验证一遍输出是否和 PyTorch 一致偶尔会因为 BN 层或 upsample 算子出现数值差异。5. CNN 图像分类复现避坑清单5 个高频翻车点这一章是血泪经验的汇总。前面几章把主流程走通后真正消耗时间的通常不是原理而是下面这几类问题。每一条都按现象、原因、解决三步写方便照着排查。5.1 训练集图片全对不上ImageFolder 的标签排序坑现象训练正常跑完验证准确率也有 80%但把测试图片单独拿出来预测结果完全不对像是标签整体错位了。原因ImageFolder 按文件夹名字母序分配标签而手工标注数据的人往往按自己习惯的类别名顺序贴标签两个顺序不一致。源码里如果用了 ImageFolder 却没有打印 class_to_idx 并核对标签错位一路带到训练和预测损失函数还在正常下降模型学到的其实是错位的映射。解决在加载数据集后立刻打印 class_to_idx并和项目文档里的类别清单逐项比对。最好在数据准备阶段写一个脚本从 Excel 或 CSV 生成固定映射存成 json训练和推理都从同一个 json 读标签避免两套顺序。5.2 CUDA out of memory显存不够不等于模型太大现象训练脚本一跑就报 CUDA out of memory然后被建议换 24G 显卡但同项目的同学用 6G 卡能跑。原因显存峰值往往不在模型参数而在中间特征图和优化器状态。batch_size64、输入 224、ResNet50 时中间特征图加起来比模型参数还多Adam 优化器还要额外存一阶二阶动量大约是参数量的两倍。新手最常见的翻车是开着验证集、tensorboard 和多个日志同时占显存。解决先把 batch_size 降到 16 或 8 试稳再逐项排查。设置 torch.cuda.empty_cache() 不能真正腾出显存治标不治本。低显存运行模型有几个实用手段启用混合精度AMP能减少约一半显存用 torch.no_grad() 包住验证把临时 Tensor 及时 del 并 detach。把迁移学习的 backbone 冻结requires_grad_(False)也能显著减少优化器动量占的显存效果接近显存减半。5.3 loss NaN 或训练集准确率一直不动学习率与初始化现象训练前几个 epoch loss 直接变成 nan或者 loss 很稳定但准确率不涨像是模型根本没在学。原因nan 最常见的原因是学习率过大梯度更新一步就冲爆如果用了预训练模型并替换分类头新初始化的分类头输出方差偏大与预训练特征尺度不匹配也会让 loss 一开始就很难看。准确率不动则通常是把 backbone 也设了 requires_gradFalse只训练分类头时又忘了把分类头放到 GPU 上。解决先打印每一步的 loss 和梯度的范数grad_norm如果 grad_norm 超过几百立刻降低学习率把 lr 从 1e-3 降到 1e-4 重试。换分类头时对新的 fc 层做更小的初始化nn.init.xavier_normal_(model.fc.weight) 或者直接单独给分类头设 10 倍小的学习率PyTorch 里可以给 optimizer 传不同参数组。数值稳定性上自定义 dataset 的标签要确认是 0 到 num_classes-1 的整数如果从 1 开始CrossEntropyLoss 很容易在边界上出问题。5.4 验证集精度高、测试集崩数据泄漏与随机种子现象验证集准确率 95%换一批测试图只剩 60%差距大到明显不合理。原因最常见的泄漏是数据划分前没有洗牌或者同一个类别的图片在 train 和 val 里来自同一批文件夹另一个泄漏来自数据增强里的 RandomCrop如果对同一张原始图在训练和验证里都做固定裁剪验证集其实早就被训练过程“见过”了。另一个隐蔽来源是随机种子没固定PyTorch 和 numpy 双随机源数据加载的 worker 顺序每次不一样导致 val 集合到训练集合的边界在每次 run 之间漂移。解决数据划分用 sklearn 的 train_test_split 并指定 random_state42在训练入口固定三个种子import random, numpy as np, torch def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark Falsecudnn.deterministicTrue 让卷积运算在浮点层面可复现但会牺牲一部分性能benchmarkFalse 关闭自动选择最优算法。这两个开关只在需要严格对比实验时开平时训练可以保持默认。5.5 文档资料和源码版本对不上先看环境再跑现象按 README 里的命令装依赖一跑就是 import 报错或者某个 API 已经不存在。原因模型文档资料写于 PyTorch 旧版本现在版本的 API 换了新写法。旧版本的 models.resnet18(pretrainedTrue) 写法弃用了DataLoader 的 pin_memory 行为也有变化甚至 Python 版本都可能不同比如老源码用 Python 3.6 的 typing 写法新版 Python 3.11 里部分语法已不兼容。解决先读文档里的 requirements 文件再看 torch 版本和 python 版本是否匹配。日常做法是为这个源码单独建一个 conda 或 venv 环境不要直接用全局环境避免把老项目和当前项目互相污染。遇到 API 弃用报错去 PyTorch 的 release note 查改动路径而不是看到一个旧写法就照抄。模型文件如果是旧版本保存的 .pth需要在 load 时指定 map_location 兼容这类问题在复现老源码时几乎必踩。6. 文档资料怎么读从模型文档反推训练细节拿到一套源码加文档资料最终能不能复现取决于你会不会读文档里那些“没写出来的内容”。模型文档资料里一般有 README、训练日志、checkpoint、配置文件。这里有一个实用心法先读 checkpoint再读 README。ckpt torch.load(checkpoints/best.pth, map_locationcpu) print(ckpt.keys()) for k, v in ckpt.items(): if not hasattr(v, shape): print(k, v)这个代码片段会把 checkpoint 里的非张量字段打印出来epoch、best_acc、超参记录、类别映射很多项目的文档没写全但 checkpoint 里留了。资料类型能反推出的信息训练日志loss/acc 曲线学习率是否衰减、是否过拟合、最佳 epochcheckpoint 里的超参字段batch size、epoch、输入尺寸、类别数模型文件的参数量主干结构是 ResNet 还是 VGG 级数据目录片段类别数量、标签顺序我的习惯是拿到文档先做三件事确认 PyTorch 版本、打印 checkpoint 的 keys、查预训练模型对应的类别数。如果文档自相矛盾以实际能跑为准先跑通再改文档。图像分类系统的源码复现八成精力花在数据准备和环境匹配上模型结构反而是最稳定的部分。希望帮到你。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →