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),方便对比不同架构的效果。

坑与注意

  1. requirements.txt 版本可能过时:示例的 PyTorch 版本多为 1.x,2.x 版本可能有不兼容改动;建议在独立虚拟环境中运行。
  2. 强化学习需要 gymnasium:gym 已停止维护,新代码请用 pip install gymnasium;REINFORCE/Actor-Critic 示例可能需要适配 gymnasium API。
  3. CUDA / MPS 加速不是默认开启--accel 标志需要显式指定,否则默认 CPU 运行。
  4. 超分辨率示例的数据集需要手动下载:BSD300 数据集不会自动下载,按 README 说明放置到 dataset/BSDS300/ 目录后使用。
  5. 语言模型Wikitext-2:首次运行会自动下载;若网络不通,可手动下载放到 data/ 目录。
  6. MNIST_HOGWILD 多线程:在共享内存环境下(不是 Docker 容器)效果最好,多进程通信开销可能反而拖慢训练。
  7. 各子目录相互独立:不要在一个子目录里 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 写自己项目之间的最佳桥梁。