pytorch/examples · 上手攻略
- 仓库:pytorch/examples
- 链接:https://github.com/pytorch/examples
- 分类:ai · pytorch · tutorials
- 作者:Tom
- 更新:2026-08-18
是什么
pytorch/examples 是 PyTorch 官方维护的示例仓库,汇集了 PyTorch 官方团队编写的各领域高质量参考实现。与 PyTorch Tutorials 不同,这些示例追求代码简洁、独立可运行、少依赖,每个子目录对应一个完整任务,读者可以直接把代码复制到自己的项目中使用或改写。覆盖领域包括:MNIST 手写数字识别(含多种变体)、VAE、图像超分辨率、强化学习(REINFORCE / Actor-Critic)、词语言模型(RNN / Transformer)、Siamese Network、时间序列预测等。
解决什么问题
学习 PyTorch 时最常见的困境是"文档看懂了,但不知道从哪里开始写自己的代码"。pytorch/examples 解决了这个问题:
- 从零到一:每个示例都是完整的可运行项目,涵盖数据加载、模型定义、训练循环、推理全链路。
- 官方权威性:代码由 PyTorch 团队编写或审核,代表了"官方推荐写法"。
- 少依赖:每个子目录有独立的
requirements.txt,避免引入整个 PyTorch 生态的复杂依赖。 - 对号入座:无论是图像、NLP、强化学习还是生成模型,都能快速找到对应示例作为起点。
快速安装
# 克隆仓库
git clone https://github.com/pytorch/examples.git
cd examples
# 进入子目录,安装该示例的专属依赖后运行
cd mnist
pip install -r requirements.txt
python main.py
⚠️ PyTorch 需单独安装(>= 1.7.1,建议用最新稳定版):
# CPU
pip install torch torchvision
# GPU(CUDA 11.8 为例)
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118
核心用法(分子目录讲解)
MNIST 系列(手写数字识别)
MNIST 是最经典的入门数据集,pytorch/examples 提供了多个变体:
基础 MLP 版本
cd mnist
pip install -r requirements.txt
python main.py
# 指定 GPU
CUDA_VISIBLE_DEVICES=0 python main.py
含异步 SGD(Hogwild)版本
cd mnist_hogwild
pip install -r requirements.txt
python main.py
RNN 版本
cd mnist_rnn
pip install -r requirements.txt
python main.py
Forward-Forward 版本(神经网络新训练范式)
cd mnist_forward_forward
pip install -r requirements.txt
python main.py
⚠️ 注意:不同 MNIST 变体使用不同的模型架构(MLP / CNN / RNN / FF),效果和训练速度差异明显,建议按需选择。
图像超分辨率(Super Resolution)
基于 Shi et al. 的 Efficient Sub-Pixel CNN 论文,在 BSD300 数据集上训练。
cd super_resolution
# 训练
python main.py \
--upscale_factor 3 \
--batchSize 4 \
--nEpochs 30 \
--lr 0.001 \
--accel
# 推理(对单张图片超分辨率)
python super_resolve.py \
--input_image dataset/BSDS300/images/test/16077.jpg \
--model model_epoch_30.pth \
--output_filename out.png \
--accel
常用参数:
| 参数 | 说明 | 默认值 |
|------|------|--------|
| --upscale_factor | 上采样倍率(2/3/4) | 必须指定 |
| --batchSize | 训练 batch size | 4 |
| --nEpochs | 训练轮数 | 30 |
| --lr | 学习率 | 0.01 |
| --accel | 启用 GPU/MPS 加速 | CPU |
强化学习(REINFORCE & Actor-Critic)
在 gym 环境中实现两种策略梯度算法。
cd reinforcement_learning
pip install -r requirements.txt
# REINFORCE 算法
python reinforce.py
# Actor-Critic 算法
python actor_critic.py
⚠️ 需要 gymnasium(原 gym)环境,默认使用 CartPole-v1。强化学习收敛对超参数敏感,遇到训练不收敛时可尝试减小学习率或增加 epoch 数。
词语言模型(RNN / GRU / LSTM / Transformer)
在 Wikitext-2 数据集上训练语言模型,支持生成文本。
cd word_language_model
# 基础 LSTM(推荐先跑这个)
python main.py --accel --epochs 6
# LSTM + 权重绑定(tied weights,降低参数量)
python main.py --accel --epochs 6 --tied
# 更大模型(650隐藏层,65 epoch)
python main.py --accel --emsize 650 --nhid 650 --dropout 0.5 --epochs 40 --tied
# Transformer 版本
python main.py --accel --epochs 6 --model Transformer --lr 5
# 生成文本
python generate.py --accel
⚠️ Wikitext-2 是小型数据集(~2M tokens),模型不要太复杂;若用 Transformer,建议同时加 --use-optimizer --lr 0.001(用 AdamW 替代默认优化器)。
变分自编码器(VAE)
实现 Kingma & Welling 的 VAE 论文(ReLU + Adam,替代原版 Sigmoid + Adagrad,收敛更快)。
cd vae
pip install -r requirements.txt
# 默认配置
python main.py
# 显式指定参数
python main.py --batch-size 128 --epochs 10 --no-accel --seed 42 --log-interval 100
时序预测(Time Sequence Prediction)
用 LSTM 预测时序数据(如天气、股价等序列)。
cd time_sequence_prediction
pip install -r requirements.txt
python main.py
⚠️ 该示例默认使用随机合成的时序数据;如用于真实数据,需修改 main.py 中的数据加载部分。
其他子目录一览
| 子目录 | 任务 | 适合场景 |
|---|---|---|
regression |
线性/多项式回归 | 入门 PyTorch 训练循环 |
siamese_network |
Siamese Network(双塔相似度) | 人脸验证/签名验证 |
dcgan |
深度卷积 GAN | 生成模型入门 |
word_language_model |
语言模型 | NLP 入门,文本生成 |
典型适用场景
- PyTorch 入门:第一次学 PyTorch,从
mnist开始,照着跑一遍就理解训练循环的全貌。 - 快速验证论文方法:读了某篇论文,想快速复现结果做消融实验,直接找对应示例改写。
- 作为自己项目的基础骨架:把
mnist的结构迁移到自己的数据集,替换数据加载器和模型即可。 - 对比不同模型:同一任务(如 MNIST)有多个变体(MLP / RNN / CNN),方便对比不同架构的效果。
坑与注意
- requirements.txt 版本可能过时:示例的 PyTorch 版本多为 1.x,2.x 版本可能有不兼容改动;建议在独立虚拟环境中运行。
- 强化学习需要 gymnasium:gym 已停止维护,新代码请用
pip install gymnasium;REINFORCE/Actor-Critic 示例可能需要适配 gymnasium API。 - CUDA / MPS 加速不是默认开启:
--accel标志需要显式指定,否则默认 CPU 运行。 - 超分辨率示例的数据集需要手动下载:BSD300 数据集不会自动下载,按 README 说明放置到
dataset/BSDS300/目录后使用。 - 语言模型Wikitext-2:首次运行会自动下载;若网络不通,可手动下载放到
data/目录。 - MNIST_HOGWILD 多线程:在共享内存环境下(不是 Docker 容器)效果最好,多进程通信开销可能反而拖慢训练。
- 各子目录相互独立:不要在一个子目录里
pip install -r requirements.txt后跳到另一个子目录直接跑,两个子目录的依赖可能不同。
与同类对比
| 资源 | 特点 | 适合人群 |
|---|---|---|
| pytorch/examples(官方) | 简洁、官方权威、少依赖 | 所有 PyTorch 用户 |
| PyTorch Tutorials | 覆盖面广、含概念讲解 | 零基础入门 |
| PyTorch Hub | 预训练模型直接加载 | 需要特定模型的人 |
| Hugging Face Transformers | 预训练 NLP 模型生态 | NLP 研究者 |
| DeepLearningExamples(NVIDIA) | GPU 优化版 | 生产部署 |
| labml.ai Annotated Papers | 论文 + 代码对照 | 研究者 |
一句话推荐结论
pytorch/examples 是 PyTorch 官方出的"代码字典"——每个示例都短小精悍、可以直接跑,是从学会 PyTorch 到能用 PyTorch 写自己项目之间的最佳桥梁。