news 2026/9/8 21:18:25

一条命令跑通PyTorch图像分类:从训练到存模型

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
一条命令跑通PyTorch图像分类:从训练到存模型

一条命令跑通PyTorch图像分类:从训练到存模型

【免费下载链接】pytorch-deep-learningMaterials for the Learn PyTorch for Deep Learning: Zero to Mastery course.项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-deep-learning

如果你看 PyTorch 的文档和零散示例,却始终拼不出一个"能训练、能测试、能存模型"的完整项目,这篇文章帮你在 10 分钟内用 pytorch-deep-learning 这个课程仓库跑通一条完整的图像分类流水线:克隆仓库、解压数据、执行一条命令,几分钟后得到一个训练好的.pth模型文件。

环境搭建和第一次训练 🚀

依赖只有三个包:torchtorchvisiontqdm(进度条)。仓库里没有requirements.txt,但训练脚本用到这些库的地方不多,装完就能跑。

git clone https://gitcode.com/GitHub_Trending/py/pytorch-deep-learning cd pytorch-deep-learning unzip data/pizza_steak_sushi.zip -d data/ pip install torch torchvision tqdm python going_modular/going_modular/train.py

最后这条命令就是全部。运行后终端会逐行打印每个 epoch 的train_losstrain_acctest_losstest_acc,跑完自动把模型权重存进models/目录。仓库里已经附带训练好的模型(如going_modular/models/05_going_modular_script_mode_tinyvgg_model.pth),想先看看效果可以直接拿它做预测。

两个注意点:解压后训练数据必须落在data/pizza_steak_sushi/train、测试数据在data/pizza_steak_sushi/test,路径不对会在建数据集时报错;另外务必在仓库根目录执行最后那条命令,原因下一节讲。

拆解:engine.py 的训练循环在干什么 🔬

整个going_modular/going_modular/目录把项目拆成 5 个各司其职的文件:data_setup.py负责把目录变成 DataLoader,model_builder.py负责网络结构(一个照搬 CNN Explainer 的 TinyVGG 卷积网络),utils.py负责保存模型,train.py只是配置页,而真正干活的训练逻辑集中在engine.pytrain_step里,核心就这几行:

for batch, (X, y) in enumerate(dataloader): X, y = X.to(device), y.to(device) y_pred = model(X) # 前向传播 loss = loss_fn(y_pred, y) # 算损失 optimizer.zero_grad() # 清空梯度 loss.backward() # 反向传播 optimizer.step() # 更新参数

它为什么这样设计:训练循环本身没有任何"魔法",难点在于它要同时管数据搬运、设备(CPU/GPU)切换、梯度和指标统计。engine.py把这些封成train_steptest_step两个函数,外面再用train()按 epoch 调度,所以train.py才薄到只有 30 行——你换数据集、换模型、换损失函数时,只动配置页,循环代码一行不碰。这也是你在真实 PyTorch 项目里最常看到的那种工程结构。test_steptrain_step的区别一句话带过:测试时用model.eval()关掉 Dropout 类行为,并用torch.inference_mode()省掉不保存计算图。

调参对照表:改哪个数字有什么用 🎛️

所有关键参数都集中在going_modular/going_modular/train.py开头(NUM_EPOCHS = 5BATCH_SIZE = 32HIDDEN_UNITS = 10LEARNING_RATE = 0.001),图像尺寸则在同文件的transforms.Resize((64, 64))里:

参数默认值调大的效果调小的效果
NUM_EPOCHS5更充分收敛,但更慢,可能过拟合快速试参,欠拟合风险
BATCH_SIZE32单 epoch 更快,显存占用更高省显存,梯度更"抖"
HIDDEN_UNITS10模型容量更大、参数更多模型更小更好训,但表达力受限
LEARNING_RATE0.001下降更快,容易震荡甚至不收敛更稳,但收敛慢
图像尺寸 (64, 64)64×64细节更多,训练更慢更快,细节丢失

另外data_setup.pyNUM_WORKERS默认取os.cpu_count(),多核机器上数据加载会自动并行。判断调参方向的标尺是损失曲线:训练集和测试集损失都下降且差距稳定是理想状态,训练损失持续降而测试损失开始回升就是过拟合,这时候先加 epoch 之外的手段——降学习率、调小 HIDDEN_UNITS 都比硬堆数据更直接。

进阶玩法:torch.compile 一行提速 🧪

训练跑通之后,如果你的 torch 是 2.0+,可以在本地克隆的going_modular/going_modular/train.py里、创建模型之后加一行:

model = torch.compile(model)

第一个 epoch 会因为编译开销变慢,之后逐 epoch 提速。仓库的extras/pytorch_2_intro.ipynb专门讲了 PyTorch 2.0 的这些新特性,extras/pytorch_2_results/里还有 ResNet50 在 CIFAR10 上 compiled 与非 compiled 的实测 CSV 和曲线图,比如下面这张 RTX 4080 上逐 epoch 训练耗时的对比,可以直接对照自己的机器验证提速幅度。

踩坑记录:第一次跑最常撞的四个错 ⚠️

  • FileNotFoundError/ 数据集为空→ 解压后目录层级和data/pizza_steak_sushi/train对不上 → 检查解压结果,必要时把 train、test 两层的子目录挪到正确位置。
  • ModuleNotFoundError: No module named 'data_setup'→ 在错误的目录层级启动脚本 → 在仓库根目录执行python going_modular/going_modular/train.py,Python 会自动把脚本所在目录加入搜索路径。
  • 训练奇慢,日志显示跑在 cpu 上→ 装的是 CPU 版 torch → 按SETUP.md重装带 CUDA 的版本,或直接换 Colab 免费 GPU 跑。
  • CUDA out of memory→ BATCH_SIZE 或图像尺寸太大 → 先把 BATCH_SIZE 从 32 降到 16 再重试。

下一步往哪走 📍

这个仓库能覆盖的范围是:从张量基础到图像分类训练、实验跟踪再到模型部署的完整代码教学,但不提供开箱即用的模型 API,也不涉及 NLP 和时间序列方向。建议的延伸路径:extras/pytorch_cheatsheet.ipynb当速查表随手查,06_pytorch_transfer_learning.ipynb学用预训练模型提精度,extras/pytorch_extra_resources.md里有按方向的延伸资源清单。

【免费下载链接】pytorch-deep-learningMaterials for the Learn PyTorch for Deep Learning: Zero to Mastery course.项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-deep-learning

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/8 21:15:49

AI Agent的Skill是什么?一文讲透Skill原理、与Prompt/Tool/Agent的边界

最近一段时间,我身边的开发者圈子几乎被“Skill”这个词刷屏了。从Claude Code到Codex,再到Trae这类集成开发环境,更新日志里高频出现Skills功能;社区里铺天盖地都是“skill推荐”“skill creator”“某某场景skill下载”&#xf…

作者头像 李华
网站建设 2026/9/8 21:14:03

MATLAB实现工业机器人DH参数辨识:从建模到0.5mm精度补偿实战

简介:面向工业机器人标定的DH参数辨识Matlab程序,实测精度可达0.5 毫米,适合需要提升机械臂绝对定位精度的研发、调试与工程人员。代码采用模块化设计,不仅覆盖旋转矩阵、DH建模、雅可比求解、工具坐标系粗标定与DH精标定等关键环…

作者头像 李华
网站建设 2026/9/8 21:12:26

4 个翻译引擎实测:划词翻译怎么选

4 个翻译引擎实测:划词翻译怎么选 【免费下载链接】pot-desktop 🌈一个跨平台的划词翻译和OCR软件 | A cross-platform software for text translation and recognition. 项目地址: https://gitcode.com/GitHub_Trending/po/pot-desktop 深夜读英…

作者头像 李华
网站建设 2026/9/8 21:11:42

Ollama本地部署实战:从安装到接入IDE与Web API的全流程指南

最近半年,我把自己主力机器的模型部署方案彻底切换到了 Ollama,从最早只是拿它跑跑 Qwen 玩,到现在 IDE 补全、项目里的 Web 小工具、脚本里的批量任务全部走本地模型,整个链路已经稳定跑了很久。这篇东西不是官方文档的复述&…

作者头像 李华