• 人工智能
  • 深度学习
  • 分布式训练
  • 模型优化

【免费下载链接】apex

A PyTorch Extension: Tools for easy mixed precision and distributed training in Pytorch

项目地址: https://gitcode.com/gh_mirrors/ap/apex
点击查看 免费下载

本文以 Apex 仓库中的 examples/imagenet/README.md 及其配套脚本 main_amp.py 为线索,系统讲解如何用 Apex 的 Automatic Mixed Precision(Amp)API 在 ImageNet 上训练 ResNet、AlexNet、VGG 等经典模型。读完本文,你将掌握 amp.initialize 三行接入模式、O0~O3 四种优化级别(opt_level)的取舍、静态与动态 Loss Scaling 的用法、基于 torch.distributed.launch 的多进程分布式训练流程,以及 --deterministic--sync_bn--prof 等调试与剖析选项,并能直接复用文中全部命令行完成实战训练。

一、示例定位:Apex 中的 ImageNet 混合精度参考实现

examples/imagenet/ 目录下只有两个文件:README.md(即本文依据的官方说明)与 main_amp.py(完整可运行的训练脚本)。该示例基于 PyTorch 官方 pytorch/examples 中的 ImageNet 训练代码改造而来,其核心目的只有一个——用最少的代码改动演示 Apex Amp 的接入方式,并展示如何通过命令行参数在多种"纯精度 / 混合精度"模式之间自由切换。

文档原文用三行代码概括了 Amp 的全部接入成本:

# 在模型和优化器构造完成之后加入
model, optimizer = amp.initialize(model, optimizer, flags...)
...
# loss.backward() 改写为:
with amp.scale_loss(loss, optimizer) as scaled_loss:
    scaled_loss.backward()

这句话还包含一个关键承诺:使用新版 Amp API 后,你完全不需要手动把模型或输入数据显式转换成 .half()——精度管理与梯度缩放全部由 Amp 接管。这正是混合精度训练对开发者最友好的形态,也是该示例反复强调的设计理念。

二、环境准备与数据集组织

运行该示例前需要完成两部分准备:

  1. 下载 ImageNet 数据集,并把验证集图片移动到按类别命名的子文件夹中(ImageFolder 格式)。官方说明推荐使用 valprep.sh 脚本完成验证集整理。
  2. 创建软链接,把训练集与验证集链接到当前目录,命令如下:
$ ln -sf /data/imagenet/train-jpeg/ train
$ ln -sf /data/imagenet/val-jpeg/ val

之后脚本以 ./(或任意数据集根路径)作为位置参数传入,脚本内部会拼出 ./train./val 两个子目录(见 main_amp.pytraindir/valdir 的构造逻辑)。

硬件与资源前提

  • 示例默认假设 GPU 显存 ≥ 16GB,因此批量大小 --b 224 才能安全运行;文档提示 --b 256 "很接近临界值",不同 PyTorch 版本下可能直接 OOM(out-of-memory)。
  • 所有示例命令均使用 --workers 4(4 个 DataLoader 子进程),用于缓解 CPU 数据加载瓶颈。
  • 脚本在启动时会断言 torch.backends.cudnn.enabled,即 Amp 要求启用 cuDNN 后端,否则直接报错。

三、优化级别(opt_level)与 Loss Scaling 速查

--opt-level 是接入 amp.initialize 的最重要标志,它决定整套精度与缩放策略。文档给出三条默认规则:

优化级别本质默认 Loss Scaling说明
O0纯 FP32 训练不用强制给 O0 加 loss scaling 会触发警告,因为纯 FP32 下缩放无意义
O3纯 FP16 训练不用"纯"模式,并非真正的混合精度
O1官方混合精度配方(推荐日常使用)动态缩放除非手动覆盖
O2"Almost FP16" 混合精度,比 O1 更激进动态缩放除非手动覆盖

也就是说:O1/O2 默认启用动态 Loss Scaling,O0/O3 默认关闭O0/O3 也可通过 --loss-scale 手动开启缩放,但文档明确警告对 O0 使用 loss scaling "没有实际意义,且会触发警告"。

速查命令集

文档给出的 Summary 命令完整如下(覆盖四种 opt_level、静态缩放覆盖、多进程分布式):

$ python main_amp.py -a resnet50 --b 128 --workers 4 --opt-level O0 ./
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O3 ./
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O3 --keep-batchnorm-fp32 True ./
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O1 ./
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O1 --loss-scale 128.0 ./
$ python -m torch.distributed.launch --nproc_per_node=2 main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O1 ./
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O2 ./
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O2 --loss-scale 128.0 ./
$ python -m torch.distributed.launch --nproc_per_node=2 main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O2 ./

注意:--loss-scale 128.0 会把动态缩放覆盖为静态缩放(128 是经典示例值,也可换为 1024 等 2 的幂)。所有命令末尾的 ./ 是数据集路径位置参数。

O0O3:纯精度训练,确立性能基线

  • 纯 FP32(O0)
$ python main_amp.py -a resnet50 --b 128 --workers 4 --opt-level O0 ./
  • 纯 FP16(O3)
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O3 ./
  • FP16 训练 + FP32 批归一化(O3)
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O3 --keep-batchnorm-fp32 True ./

文档特别解释了 --keep-batchnorm-fp32 True 的两个收益:其一,BN 保持 FP32 能提升训练稳定性;其二,它允许 PyTorch 使用 cuDNN 的批归一化实现,而 cuDNN BN 在 ResNet50 上能显著提速。

需要如实指出的是:O3 相关配置可能无法收敛,因为它不是真正的混合精度——没有 FP32 权重主副本、没有缩放保护。它的价值在于为你的模型建立"光速"(speed of light)性能基线,供 O1/O2 对比参考。对 ResNet50 而言,--opt-level O3 --keep-batchnorm-fp32 True 即最优基线组合;去掉 --keep-batchnorm-fp32 反而更慢(因为退出了 cuDNN BN 路径)。

O1:官方混合精度配方(推荐典型使用)

O1 的工作机制是对 Torch 函数打补丁,按"白名单-黑名单"模型自动决定输入精度:对 Tensor Core 友好的算子(GEMM、卷积)输入走 FP16,对 FP32 更受益的算子(batchnorm、softmax)保持 FP32,同时默认启用动态 Loss Scaling。这种"函数级"打补丁方式意味着你既不需要改模型定义,也不需要手动 cast 数据。

$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O1 ./

覆盖为静态缩放:

$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O1 --loss-scale 128.0 ./

两进程分布式(每进程 1 卡):

$ python -m torch.distributed.launch --nproc_per_node=2 main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O1 ./

性能建议:--nproc_per_node 设为节点上的 GPU 总数,以利用全部可用资源。

O2:"Almost FP16"(比 O1 更危险)

O2 的动作更激进:把整个模型 cast 到 FP16、batchnorm 保持 FP32、权重维护 FP32 主副本,并默认启用动态 Loss Scaling。与 O1 不同,O2 不打补丁 Torch 函数。文档明确指出,O2 主要服务于某些内部用例,日常使用请优先选择 O1

$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O2 ./
$ python main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O2 --loss-scale 128.0 ./
$ python -m torch.distributed.launch --nproc_per_node=2 main_amp.py -a resnet50 --b 224 --workers 4 --opt-level O2 ./

四、分布式训练:DDP 包装与 amp.initialize 的先后顺序

main_amp.py 可选使用 apex.parallel.DistributedDataParallel(Apex DDP)做"每进程一张卡"的多进程训练。Apex DDP 是 Torch DDP 的直接替换

# Apex 写法
model = apex.parallel.DistributedDataParallel(model)

# 等价于 Torch 写法
model = torch.nn.parallel.DistributedDataParallel(model,
                                                  device_ids=[arg.local_rank],
                                                  output_device=arg.local_rank)

两者差异在于:Torch DDP 允许单进程管理多卡,因此必须手动指定运行设备与输出设备;而 Apex DDP 默认只用当前设备(current device),无需显式指定。

文档特别澄清了几个容易混淆的点:

  1. DDP 包装器的选择与 Amp 及其他 Apex 工具正交——apex.amp 与 Torch DDP、Apex DDP 均安全兼容。
  2. amp.initialize 必须先于 DDP(model) 执行
model, optimizer = amp.initialize(model, optimizer, flags...)   # 先行
model = DDP(model)                                              # 后行

如果先做 DDP 包装再调 amp.initialize,会直接抛错。 3. 无论用哪种 DDP,每个进程都必须在创建模型或任何张量之前调用 torch.cuda.set_device(args.local_rank)。 4. 启动方式统一使用 PyTorch 的多进程启动器:

python -m torch.distributed.launch --nproc_per_node=NUM_GPUS main_amp.py args...

其中 NUM_GPUS 应小于等于节点可见 GPU 数。torch.distributed.launch 的使用与 DDP 包装器选择无关,两种 DDP 都可用它启动。

此外,若希望跨进程使用同步批归一化,只需在参数末尾追加 --sync_bn。从源码看,该标志触发的是 apex.parallel.convert_syncbn_model(model)(见 main_amp.py)。

五、确定性训练与性能剖析

--deterministic:位级可复现

加上 --deterministic 后,无论搭配什么其他选项,多次运行都应产出逐位相同的输出。其实现代价是禁用 torch.backends.cudnn.benchmark(源码中同时设置 cudnn.deterministic = True 并按 local_rank 播种 torch.manual_seed),因此可能带来一定的性能下降——它更适合调试与回归验证场景,而非追求吞吐的训练场景。

--prof 与 NVTX 剖析

如果你关心网络在 CPU/GPU 时间线上的真实形态(例如整体利用率如何、prefetcher 是否真正重叠了数据搬运),可以对 main_amp.py 做性能剖析。源码提供了 --prof N 参数:从第 N 次迭代开始执行 cudaProfilerStart(),并围绕 forward/backward/optimizer.step/prefetcher.next 等关键段打入 torch.cuda.nvtx.range_push/pop 标记,运行 10 个迭代后自动 cudaProfilerStop() 并退出。配合 Nsight 等工具即可直观观察每个阶段的时间占比。

六、源码级深入:main_amp.py 的实现细节与参数速查

main_amp.py(共 513 行)在官方 PyTorch 示例基础上做了多处面向混合精度与吞吐的改造,以下事实均可在 main_amp.py 中直接核对。

6.1 命令行参数总表

参数默认值作用
data(位置参数)必填数据集根目录,内部拼出 train/val/
-a/--archresnet18模型架构,可选值为 torchvision 全部小写模型名
-j/--workers4DataLoader 子进程数
--epochs90总训练轮数
--start-epoch0断点续训起始轮
-b/--batch-size256每个进程的 mini-batch 大小
--lr0.1初始学习率,会按全局 batch size 自动缩放
--momentum0.9SGD 动量
--weight-decay1e-4权重衰减
--print-freq10日志打印间隔
--resume断点文件路径
-e/--evaluate关闭仅做验证集评测
--pretrained关闭使用预训练权重
--prof-1性能剖析:仅运行 10 个迭代
--deterministic关闭确定性训练
--local_rank环境变量 LOCAL_RANK多进程本地 rank
--sync_bn关闭启用 Apex 同步 BN
--opt-levelAmp 优化级别(O0/O1/O2/O3)
--keep-batchnorm-fp32None是否保持 BN 在 FP32
--loss-scaleNone覆盖为静态 Loss Scaling
--channels-lastFalse使用 NHWC 内存格式

6.2 学习率缩放与 warmup 调度

脚本按"全局 batch size"(分布式进程数 × 每进程 batch size)缩放学习率:

args.lr = args.lr * float(args.batch_size * args.world_size) / 256.

调度器目标是在 batch size 256 下收敛到约 76% 的 Top-1 精度:以 epoch // 30 为衰减因子执行 lr * 0.1**factor(epoch ≥ 80 额外 +1),且前 5 个 epoch 执行线性 warmup(见 adjust_learning_ratemain_amp.py)。

6.3 为吞吐优化的数据管线

  • fast_collate:跳过 transforms.ToTensor() 的逐像素 PIL→float 转换,直接把 uint8 图像堆叠进预分配张量(支持 channels_last NHWC 格式),并同步打包 int64 标签。
  • data_prefetcher:在独立 CUDA stream 上以 non_blocking=True 预取下一批数据,GPU 端用 mean/std(已乘 255)完成归一化,再通过 record_stream 保证内存生命周期正确——这正是文档所说"prefetcher 真正重叠数据搬运"的落点。注释中还保留了一行重要提示:"With Amp, it isn't necessary to manually convert data to half."(数据无需手动转半精度)。

6.4 关于 Amp API 的演进说明

需要向读者如实交代一个细节:README 描述的是 Apex 经典的 amp.initialize / amp.scale_loss 接入模式,而当前仓库的 main_amp.py 训练主循环已迁移到 PyTorch 原生 AMP 接口:

scaler = torch.amp.GradScaler("cuda")
...
with torch.autocast(device_type="cuda"):
    output = model(input)
    loss = criterion(output, target)
scaler.scale(loss).backward()
...
scaler.step(optimizer)
scaler.update()

脚本仍解析并打印 --opt-level--keep-batchnorm-fp32--loss-scale 三个兼容性标志,但实际前向/反向由 torch.autocastGradScaler 驱动。从源码结构可以推断,该示例正在从 Apex Amp API 向 PyTorch 内置 AMP 过渡;阅读本文第一部分的三行接入模式时,应把它理解为 Amp 的经典编程模型,而运行当前仓库代码时则以 torch.amp 的实际行为为准。

6.5 其他工程细节

  • 断点保存为 checkpoint.pth.tar,最优模型另存 model_best.pth.tar,仅在 local_rank == 0 的进程执行。
  • 评估阶段在 torch.no_grad() 下运行,返回 Top-1 平均精度;分布式下日志指标通过 dist.all_reduce 跨进程求平均(reduce_tensor)。
  • inception_v3 架构当前不被支持,脚本会直接抛出 RuntimeError

七、仓库内的测试佐证

该示例并非孤立脚本,仓库的 L1 测试体系复用了同样的训练模式:tests/L1/common/main_amp.py 是基于 examples/imagenet 思路的测试版训练脚本,而 tests/L1/common/run_test.sh 以如下命令驱动单卡与双卡回归:

python main_amp.py -a resnet50 --b 128 --workers 4 --deterministic --prints-to-process 5
python -m torch.distributed.launch --nproc_per_node=2 main_amp.py -a resnet50 --b 128 --workers 4 --deterministic --prints-to-process 5

tests/L1/cross_product/run.sh 的注释中则直接指定数据集目录为 examples/imagenet/bare_metal_train_val/,印证了该示例目录同时充当 L1 交叉测试(不同 opt_level 与分布式组合)的数据与代码底座。如果你想在改动示例后做快速冒烟验证,可以参照上述测试命令。

八、实践建议小结

  • 日常训练首选 O1:白名单-黑名单打补丁 + 动态 Loss Scaling 的组合在精度与速度间最均衡。
  • 性能基线用 O3 --keep-batchnorm-fp32 True:它代表"光速"上限,但不保证收敛,仅作对比基准。
  • 显存不足时降 --b224 已要求 ≥16GB 显存,128 更稳妥。
  • 多卡必守两条铁律amp.initialize 先于 DDP 包装;每个进程先 torch.cuda.set_device(args.local_rank)
  • 调试用 --deterministic,调优用 --prof:前者保证逐位可复现,后者通过 NVTX 标记定位 CPU/GPU 时间线瓶颈。

本文所有命令与结论均可对照 examples/imagenet/README.mdexamples/imagenet/main_amp.py 逐行验证;若要进一步了解 Apex 的其他能力(融合优化器、融合 LayerNorm、同步 BN 等),可继续阅读 apex/init.py 暴露的模块入口与 docs/source/index.rst 的 API 索引。

  • 人工智能
  • 深度学习
  • 分布式训练
  • 模型优化

【免费下载链接】apex

A PyTorch Extension: Tools for easy mixed precision and distributed training in Pytorch

项目地址: https://gitcode.com/gh_mirrors/ap/apex
点击查看 免费下载
Logo

火山引擎视频云技术社区,是面向 AI 音视频开发者的技术交流平台。这里汇聚源自抖音、豆包等亿级 DAU 产品的 RTC、直播、点播、AI 媒体处理、音视频互动技术,提供接入指南、最佳实践、性能调优、场景案例、Demo 代码、开源项目、白皮书和 API 文档。社区汇聚官方工程师与一线开发者,为 AI 视频通话、数字人、AI 视频处理等应用的开发与落地提供技术支持。

更多推荐