从零开始:用 PyTorch 官方代码训练 ResNet18 图像分类模型(迁移学习实战)

发布时间:2026/8/10 17:17:22

从零开始:用 PyTorch 官方代码训练 ResNet18 图像分类模型(迁移学习实战)
适用人群有一定 Python 基础想入门深度学习图像分类的开发者本文目标用 PyTorch 官方脚本 预训练权重训练一个自己的图像分类模型并完成测试验证前言很多初学者第一次接触深度学习时都希望训练一个属于自己的图像分类模型。但自己从零搭网络、从零训练往往又慢又容易梯度爆炸。其实有一个更聪明的做法迁移学习Transfer Learning——借用已经在 ImageNet 千万级图片上训练好的 ResNet18只需要微调最后几层就能快速迁移到自己的任务上。本文会用PyTorch 官方开箱即用的训练代码一步步带你完成环境准备 → 数据集准备 → 理解代码 → 开始训练 → 测试验证。整个流程不依赖任何私有环境你在自己的电脑或云服务器上都能复现。目录环境准备数据集准备下载官方代码什么是迁移学习核心概念代码关键点讲解开始训练测试与验证总结与常见问题一、环境准备1.1 安装 Anaconda如果还没有去官网下载 Anaconda 并安装。它自带 Python还能方便地创建隔离的环境避免不同项目依赖冲突。1.2 创建一个虚拟环境打开终端Windows 用 Anaconda PromptLinux/macOS 用自带终端执行conda create-nresnetpython3.10-yconda activate resnet1.3 安装 PyTorchPyTorch 官网会根据你的系统生成对应的安装命令。关键选择点是你的机器有没有 NVIDIA 显卡有显卡GPU装 GPU 版训练速度快很多没有显卡CPU装 CPU 版能跑但慢# GPU 版以 CUDA 12.x 为例具体见官网生成的命令pipinstalltorch torchvision --index-url https://download.pytorch.org/whl/cu121# CPU 版pipinstalltorch torchvision安装后验证是否成功importtorchprint(torch.__version__)# 版本号print(torch.cuda.is_available())# 有 GPU 应输出 True1.4 安装其他依赖pipinstallpillow numpy二、数据集准备2.1 数据长什么样图像分类任务的数据推荐按下面的目录结构存放这是torchvision.datasets.ImageFolder的标准格式数据集根目录/ ├── train/ # 训练集 │ ├── 类别1/ # 每个子文件夹是一个类别 │ │ ├── 图片1.jpg │ │ ├── 图片2.jpg │ │ └── ... │ ├── 类别2/ │ └── ... ├── val/ # 验证集结构同上每类放少量图片 └── test/ # 测试集结构同上用于最终评估关键规则每个类别一个文件夹文件夹名就是类别名如cat、dogtrain和val里每个类放多少张图一般来说训练集每类越多越好最少也要几十张验证集每类放 10 张左右即可图片建议统一为 JPG/PNG 格式RGB 三通道2.2 数据量多大合适迁移学习的好处之一就是不需要海量数据。因为预训练模型已经学会了通用特征你只需要教会它区分你自己的类别。即使是每类几十张的小数据集也能获得不错的效果。三、下载官方代码PyTorch 官方在pytorch/vision仓库里维护了一套可以直接使用的分类训练代码包含训练、数据增强、评估等完整功能。这是最规范、最值得学习的版本。在你的项目目录下执行mkdirofficialcdofficialforfintrain.py utils.py transforms.py sampler.py presets.py;docurl-Ohttps://raw.githubusercontent.com/pytorch/vision/v0.25.0/references/classification/$fdone如果下载慢也可以直接在浏览器打开上面链接手动保存或从国内镜像站获取。下载完成后应该有 5 个文件train.py主训练脚本、utils.py工具函数、presets.py数据增强预设、transforms.py、sampler.py。四、什么是迁移学习核心概念在加载代码前先理解本文最重要的概念。ResNet18 的网络结构大致是图片 → conv1 → layer1 → layer2 → layer3 → layer4 → fc(全连接层) → 1000类输出前半部分conv1~layer4卷积骨干负责提取图片特征边缘、纹理、形状、物体部件最后一部分fc 全连接层根据特征做分类输出各类别的概率帮别人训练好的权重骨干部分已经非常擅长提取通用特征。而fc层的输出维度是 1000ImageNet 的类别数跟你的任务类别数对不上。迁移学习的做法就是保留预训练好的卷积骨干把最后一层fc换成匹配你自己类别数的新层随机初始化用你自己的数据训练整个网络或只训练 fc 层这样新层从零学你的类别骨干层只需微调训练又快效果又好。五、代码关键点讲解5.1 获取预训练权重torchvision 提供了官方的预训练权重。有两种方式加载方式一在线下载torchvision 自动拉取importtorchvision weightstorchvision.models.get_weight(ResNet18_Weights.IMAGENET1K_V1)state_dictweights.get_state_dict(progressTrue)方式二本地文件加载推荐避免网络不稳定先手动下载权重文件resnet18.pth约 46MB放到项目目录然后importtorch state_dicttorch.load(resnet18.pth,map_locationcpu,weights_onlyFalse)ifisinstance(state_dict,dict)andstate_dictinstate_dict:state_dictstate_dict[state_dict]是在没有我这里可以提供resnet18.pth后台私信我即可5.2 迁移学习改造核心代码官方脚本里get_model传预训练权重时会强制把类别数覆盖成 1000无法适配你自己的类别数。需要在「创建模型」处做改造importtorchvision.modelsasmodels num_classes102# 换成你自己的类别数# ① 新建一个 num_classes 类别的模型fc 层随机初始化modelmodels.resnet18(num_classesnum_classes)# ② 加载预训练权重state_dicttorch.load(resnet18.pth,map_locationcpu,weights_onlyFalse)ifisinstance(state_dict,dict)andstate_dictinstate_dict:state_dictstate_dict[state_dict]# ③ 删掉最后一层 fc 的权重因为类别数不同形状不匹配state_dict.pop(fc.weight,None)state_dict.pop(fc.bias,None)# ④ strictFalse只加载卷积骨干层fc 层保持随机初始化missing,unexpectedmodel.load_state_dict(state_dict,strictFalse)print(missing:,len(missing),unexpected:,len(unexpected))# 输出 missing:2 正是被替换的 fc 层符合预期5.3 数据增强官方 presets.pypresets.py里定义了两套预处理训练集增广防止过拟合RandomResizedCrop(224)随机裁剪缩放到 224×224RandomHorizontalFlip(0.5)随机水平翻转Normalize用 ImageNet 均值/方差归一化验证/测试集只做固定尺寸调整Resize(256)→CenterCrop(224)→Normalize5.4 训练主循环官方 train.py每个 epoch 做三件循环往复的事1. train_one_epoch(model, ...) # 训练前向传播 → 算损失 → 反向传播 → 更新权重 2. evaluate(model, ...) # 在 val 集评估准确率 3. 保存 checkpoint # 保存 model、optimizer、lr_scheduler 等损失函数用交叉熵CrossEntropyLoss优化器用 SGD带动量学习率用 StepLR 调度。六、开始训练6.1 训练命令python official/train.py\--data-path 数据集根目录\--modelresnet18\--pretrained-path resnet18.pth\--devicecuda\--epochs50\--batch-size16\--lr0.01\--output-dir output\--workers0⚠️ Windows 必加--workers 0Windows 上 PyTorch 默认会用多进程加载数据--workers默认为 16与pin_memory叠加容易触发CUDA error: resource already mapped报错。把 worker 数设为 0主进程直接加载即可规避小数据集加载速度影响可忽略。详见「八、常见问题 Q7」。6.2 参数解释参数含义建议--data-path数据集根目录含 train/ 和 val/必填--model模型名这里用 resnet18resnet18--pretrained-path本地预训练权重路径填你下载的权重--device用 cuda 还是 cpu有 GPU 填 cuda--epochs训练轮数小数据集 30~50 轮效果更好--batch-size每批图片数GPU 显存允许就 16 或更大--lr学习率微调用 0.01比从零训练的 0.1 小--output-dircheckpoint 保存目录自定义--workers数据加载进程数Windows 建议 0Linux 可设 4~8为什么微调要用更小的学习率因为预训练权重已经接近最优点学习率太大容易把学好的特征破坏掉小学习率只做精细调整。6.3 观察训练输出训练过程中会实时打印Epoch: [9] Total time: 0:00:18 Acc1 72.549 Acc5 92.353 Test: Acc1 69.96 Acc5 91.17loss训练损失整体应逐渐下降acc1 / acc5当前 batch 的 Top-1 和 Top-5 准确率Test每轮结束在 val 集上的评估结果训练结束后output/目录下会生成model_0.pth~model_49.pth以及checkpoint.pth每个文件对应一个 epoch 的模型。七、测试与验证官方脚本的--test-only只评估 val 集。如果还有独立的 test 集建议写一个测试脚本加载训练好的模型在 test 集上评估并支持单张图片推理。7.1 测试脚本test.py 用法: 1) 评估整个 test 集准确率: python test.py --checkpoint output/model_9.pth --data-path 数据集/test --mode eval 2) 对单张图片推理: python test.py --checkpoint output/model_9.pth --image path/to/img.jpg --mode predict importargparseimporttorchimporttorchvisionfromtorchvisionimporttransformsfromtorch.utils.dataimportDataLoaderfromtorchvision.datasetsimportImageFolderfromPILimportImage MEAN(0.485,0.456,0.406)STD(0.229,0.224,0.225)defbuild_model(checkpoint_path,device):从 checkpoint 恢复模型, 自动匹配类别数ckpttorch.load(checkpoint_path,map_locationdevice,weights_onlyFalse)num_classesckpt[model][fc.weight].shape[0]modeltorchvision.models.resnet18(num_classesnum_classes)model.load_state_dict(ckpt[model])model.to(device)model.eval()returnmodel,num_classesdefeval_testset(checkpoint,data_path,batch_size64):在 test 集上评估 Top-1 / Top-5 准确率devicecudaiftorch.cuda.is_available()elsecpumodel,_build_model(checkpoint,device)transformtransforms.Compose([transforms.Resize(256),transforms.CenterCrop(224),transforms.ToTensor(),transforms.Normalize(MEAN,STD),])datasetImageFolder(data_path,transformtransform)loaderDataLoader(dataset,batch_sizebatch_size,num_workers0,pin_memoryTrue)print(ftest 集:{len(dataset)}张,{len(dataset.classes)}类)correct1correct5total0withtorch.no_grad():forimages,targetsinloader:images,targetsimages.to(device),targets.to(device)outputsmodel(images)_,predoutputs.topk(5,1,True,True)correct1(pred[:,0]targets).sum().item()correct5(predtargets.unsqueeze(1)).any(dim1).sum().item()totaltargets.size(0)print(fTop-1 Acc:{correct1/total*100:.2f}%)print(fTop-5 Acc:{correct5/total*100:.2f}%)defpredict_image(checkpoint,image_path,topk3):对单张图片做推理, 输出 Top-k 类别devicecudaiftorch.cuda.is_available()elsecpumodel,_build_model(checkpoint,device)transformtransforms.Compose([transforms.Resize(256),transforms.CenterCrop(224),transforms.ToTensor(),transforms.Normalize(MEAN,STD),])imgImage.open(image_path).convert(RGB)tensortransform(img).unsqueeze(0).to(device)withtorch.no_grad():outputsmodel(tensor)probstorch.softmax(outputs,dim1)[0]top_probs,top_idxprobs.topk(topk)print(f图片:{image_path})foriinrange(topk):print(f Top{i1}:{dataset_class_names[top_idx[i]]}置信度{top_probs[i].item()*100:.2f}%)if__name____main__:parserargparse.ArgumentParser()parser.add_argument(--checkpoint,requiredTrue)parser.add_argument(--mode,choices[eval,predict],defaulteval)parser.add_argument(--data-path,helptest 集目录(评估模式用))parser.add_argument(--image,help单张图片路径(推理模式用))parser.add_argument(--batch-size,typeint,default64)argsparser.parse_args()ifargs.modeeval:assertargs.data_path,--data-path 必填eval_testset(args.checkpoint,args.data_path,args.batch_size)else:assertargs.image,--image 必填predict_image(args.checkpoint,args.image)7.2 评估 test 集准确率python test.py\--checkpointoutput/model_49.pth\--modeeval\--data-path 数据集/test输出类似test 集: 6149 张, 102 类 Top-1 Acc: 69.96% Top-5 Acc: 91.17%7.3 单张图片推理python test.py--checkpointoutput/model_49.pth--modepredict--image某张图片.jpg输出图片的类别和置信度图片:classification/test/class_0/image_06734.jpg Top1: class_0 置信度84.63% Top2: class_48 置信度 11.51% Top3: class_62 置信度2.10%八、总结与常见问题8.1 流程回顾环境准备 → 数据整理 → 下载官方代码 → 加载预训练权重 → 迁移学习改造 → 训练 → 测试核心就一句话直接用官方代码加载预训练权重替换最后一层全连接层然后用小学习率微调。这是目前最主流、最省力的图像分类落地方案。8.2 常见问题Q1没有 GPU 能跑吗能。把--device cpu训练会慢但小数据集也能跑通。Q2准确率上不去怎么办先检查数据每类图片是否太少、类别是否均衡。再考虑增加训练轮数、减小学习率、增强数据增广。Q3missing keys: 2正常吗正常。那 2 个 missing 的 key 正是被替换的 fc 层权重说明迁移学习改造成功。Q4预训练权重下载太慢或失败手动下载resnet18.pth放本地用--pretrained-path指定避免在线下载。Q5训练 loss 不下降检查是否正确加载了预训练权重、学习率是否过小、数据是否已正确归一化。Q6如何换其他模型代码支持很多模型把--model resnet18换成resnet50、mobilenet_v3等即可同时换对应的预训练权重。Q7训练时报CUDA error: resource already mapped怎么办这是 Windows 上 PyTorch 多进程数据加载默认--workers 16与pin_memory叠加引发的已知冲突不是代码或数据问题。解决训练命令加--workers 0主进程直接加载数据小数据集几乎无速度损失。若仍报错可再降--batch-size。Q8训练时--batch-size大会怎样显存占用更高但每轮迭代数更少。显存不够时减小 batch-size或开启--amp混合精度省显存加速。参考资料PyTorch 官方代码pytorch/vision references/classificationPyTorch 官方文档torchvision 预训练权重说明希望这篇教程能帮你迈出深度学习图像分类的第一步。动手跑通一遍比看十遍都管用。祝你训练愉快

相关新闻

java的深拷贝和浅拷贝

java的深拷贝和浅拷贝

2026/8/10 17:17:22

总结:深拷贝:基本类型复制其值,引用类型都会创建新的实例。浅拷贝:对于基本类型就是复制其值,对于引用类型则是复制了指向这些数据类型的内存地址。浅拷贝(Shallow Copy)浅拷贝是指在创建新对象…

Micrometer 系列【33】Spring Boot Micrometer Metrics 自动配置模块解析

Micrometer 系列【33】Spring Boot Micrometer Metrics 自动配置模块解析

2026/8/10 17:17:22

文章目录1. 概述2. 包结构与模块划分2.1 整体分层2.2 根包:全局Filter 故障诊断2.3 actuate/endpoint 指标端点2.3 autoconfigure 核心自动配置包2.3.1 核心自动配置类2.3.2 autoconfigure/export 导出到监控系统2.3.3 内置指标 Binder 模块2.4 OTLP 开发测试辅助2…

类似小鹅通的知识付费小程序平台有哪些

类似小鹅通的知识付费小程序平台有哪些

2026/8/10 17:17:22

知识付费小程序平台怎么选?拆解码云数智、小鹅通、有赞教育三大主流工具随着私域流量持续发展,知识付费已经成为讲师、自媒体博主、培训机构变现的重要路径。搭建一套属于自己的知识付费小程序,不用依附第三方内容平台,学员数据、…

元景万悟:企业级AI智能体开发架构与高可用微服务实践深度解析

元景万悟:企业级AI智能体开发架构与高可用微服务实践深度解析

2026/8/10 18:17:24

元景万悟:企业级AI智能体开发架构与高可用微服务实践深度解析 【免费下载链接】wanwu China Unicoms Yuanjing Wanwu Agent Platform is an enterprise-grade, multi-tenant AI agent development platform. It helps users build applications such as intelligent…

Volume Cloud配置全解析:天气图与高度密度图塑造千变万化的云层

Volume Cloud配置全解析:天气图与高度密度图塑造千变万化的云层

2026/8/10 18:17:24

Volume Cloud配置全解析:天气图与高度密度图塑造千变万化的云层 【免费下载链接】VolumeCloud Volume cloud for Unity3D 项目地址: https://gitcode.com/gh_mirrors/vo/VolumeCloud Volume Cloud是Unity3D中一款强大的体积云渲染工具,通过天气图…

webpack-simple-starter vs 框架脚手架:谁更适合中小型前端项目?

webpack-simple-starter vs 框架脚手架:谁更适合中小型前端项目?

2026/8/10 18:17:24

webpack-simple-starter vs 框架脚手架:谁更适合中小型前端项目? 【免费下载链接】webpack-simple-starter A simple webpack starter without framework (Like Vue, React, Angular, etc.) 项目地址: https://gitcode.com/gh_mirrors/we/webpack-simp…

为什么选择 reshade-steam-proton?Linux 平台 ReShade 工具横向对比评测

为什么选择 reshade-steam-proton?Linux 平台 ReShade 工具横向对比评测

2026/8/10 18:17:24

为什么选择 reshade-steam-proton?Linux 平台 ReShade 工具横向对比评测 【免费下载链接】reshade-steam-proton Easy setup and updating of ReShade on Linux for games using wine or proton. 项目地址: https://gitcode.com/gh_mirrors/re/reshade-steam-prot…

ACE-Step UI终极指南:免费开源的AI音乐创作利器

ACE-Step UI终极指南:免费开源的AI音乐创作利器

2026/8/10 18:17:24

ACE-Step UI终极指南:免费开源的AI音乐创作利器 【免费下载链接】ace-step-ui 🎵 The Ultimate Open Source Suno Alternative - Professional UI for ACE-Step 1.5 AI Music Generation. Free, local, unlimited. Stop paying for Suno! 项目地址: ht…

Spektrum核心功能大揭秘:从频率扫描到光标测量的完整教程

Spektrum核心功能大揭秘:从频率扫描到光标测量的完整教程

2026/8/10 18:07:24

Spektrum核心功能大揭秘:从频率扫描到光标测量的完整教程 【免费下载链接】spektrum rtl-sdr spectrum analyzer 项目地址: https://gitcode.com/gh_mirrors/sp/spektrum Spektrum是一款功能强大的rtl-sdr频谱分析仪,能够帮助用户轻松实现频率扫描…

比较好的亚太EMBA,问了6位校友师资差别真的挺大

比较好的亚太EMBA,问了6位校友师资差别真的挺大

2026/8/10 5:58:32

比较好的亚太EMBA核心差异先看什么?对于希望兼顾工作与系统管理能力提升的亚太区高管而言,筛选匹配度高的EMBA项目时,师资配置是决定学习体验与实际收获的核心要素之一。我们结合3-4个公开信息透明、办学历史较长的亚太区主流EMBA项目特点&am…

备考3个月对比6份资料 海外游学的亚洲EMBA面试注意点

备考3个月对比6份资料 海外游学的亚洲EMBA面试注意点

2026/8/10 7:54:12

备考海外游学的亚洲EMBA面试,核心要围绕项目国际化设计逻辑、个人跨文化管理经验匹配度两个维度准备,避免把游学模块等同于普通旅游参访的认知偏差。不少备考者花3个月对比6份资料,却容易忽略面试官对“国际视野落地能力”的考察——比如香港…

比较好的国内EMBA,问了二十位校友聊透人脉价值

比较好的国内EMBA,问了二十位校友聊透人脉价值

2026/8/10 7:19:21

比较好的国内EMBA核心差异体现在哪些方面?比较好的国内EMBA的核心长期价值,很大程度上依托于校友网络的连接质量与资源生态的活跃度,这也是不少高管在择校时优先考量的因素。我们结合3-4个市场关注度较高的项目公开信息,从课程、师…

Prometheus 监控体系深度部署:选型别只看功能清单

Prometheus 监控体系深度部署:选型别只看功能清单

2026/8/10 0:06:33

Prometheus 监控体系深度部署:选型别只看功能清单 选型场景:小规模集群直接部署 Thanos 的代价 如果为解决 15 天本地存储限制,直接部署 Thanos Sidecar、Store Gateway、Querier、Compactor、Ruler、Bucket Web 并接入 S3,就需…

ELK 日志分析平台与全链路追踪:代码评审该盯住哪些细节

ELK 日志分析平台与全链路追踪:代码评审该盯住哪些细节

2026/8/10 0:06:33

ELK 日志分析平台与全链路追踪:代码评审该盯住哪些细节 场景示例:一条 2MB 日志影响 Elasticsearch 写入 一个上传接口若执行 log.Info("Request dumped: ", r.Body),会将 2MB 的二进制 Body 写入日志。高并发下,这类超…

从零到一构建开源项目的完整历程:代码评审该盯住哪些细节

从零到一构建开源项目的完整历程:代码评审该盯住哪些细节

2026/8/10 0:06:33

从零到一构建开源项目的完整历程:代码评审该盯住哪些细节 项目进入稳定版本后,外部 Pull Request(PR)会带来新的协作成本。大范围改动混入风格重构,或修复局部问题时修改公共函数签名,都可能扩大评审和兼容…

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

2026/8/8 5:07:31

一天写完毕业论文在2026年已不再是天方夜谭。2026年最炸裂、实测能大幅提速的AI论文写作工具,覆盖选题构思、文献整理、内容生成、格式排版等核心场景,真正帮你高效搞定论文难题。 一、全流程王者:一站式搞定论文全链路(一天定稿首…

导师推荐!2026最新AI论文工具测评与实用推荐

导师推荐!2026最新AI论文工具测评与实用推荐

2026/8/9 13:42:46

2026年真正好用的AI论文工具,核心看生成的论文质量、低AI味、格式正确、学术适配四大指标。综合实测,千笔AI、ThouPen、豆包、DeepSeek、Grammarly 是当前最值得推荐的梯队,覆盖从免费到付费、从中文到英文、从文科到理工的全场景需求。 一、…

告别游戏崩溃:XCOM 2模组管理器的智能革命

告别游戏崩溃:XCOM 2模组管理器的智能革命

2026/8/8 2:30:15

告别游戏崩溃:XCOM 2模组管理器的智能革命 【免费下载链接】xcom2-launcher The Alternative Mod Launcher (AML) is a replacement for the default game launchers from XCOM 2 and XCOM Chimera Squad. 项目地址: https://gitcode.com/gh_mirrors/xc/xcom2-lau…