Apex Amp 混合精度 ImageNet 训练实战:examples/imagenet/main_amp.py 与 O0/O1/O2/O3 优化级别全解析
【免费下载链接】apex
A PyTorch Extension: Tools for easy mixed precision and distributed training in Pytorch
本文以 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 接管。这正是混合精度训练对开发者最友好的形态,也是该示例反复强调的设计理念。
二、环境准备与数据集组织
运行该示例前需要完成两部分准备:
- 下载 ImageNet 数据集,并把验证集图片移动到按类别命名的子文件夹中(ImageFolder 格式)。官方说明推荐使用
valprep.sh脚本完成验证集整理。 - 创建软链接,把训练集与验证集链接到当前目录,命令如下:
$ ln -sf /data/imagenet/train-jpeg/ train
$ ln -sf /data/imagenet/val-jpeg/ val
之后脚本以 ./(或任意数据集根路径)作为位置参数传入,脚本内部会拼出 ./train 与 ./val 两个子目录(见 main_amp.py 中 traindir/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 的幂)。所有命令末尾的 ./ 是数据集路径位置参数。
O0 与 O3:纯精度训练,确立性能基线
- 纯 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),无需显式指定。
文档特别澄清了几个容易混淆的点:
- DDP 包装器的选择与 Amp 及其他 Apex 工具正交——
apex.amp与 Torch DDP、Apex DDP 均安全兼容。 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/--arch | resnet18 | 模型架构,可选值为 torchvision 全部小写模型名 |
-j/--workers | 4 | DataLoader 子进程数 |
--epochs | 90 | 总训练轮数 |
--start-epoch | 0 | 断点续训起始轮 |
-b/--batch-size | 256 | 每个进程的 mini-batch 大小 |
--lr | 0.1 | 初始学习率,会按全局 batch size 自动缩放 |
--momentum | 0.9 | SGD 动量 |
--weight-decay | 1e-4 | 权重衰减 |
--print-freq | 10 | 日志打印间隔 |
--resume | 空 | 断点文件路径 |
-e/--evaluate | 关闭 | 仅做验证集评测 |
--pretrained | 关闭 | 使用预训练权重 |
--prof | -1 | 性能剖析:仅运行 10 个迭代 |
--deterministic | 关闭 | 确定性训练 |
--local_rank | 环境变量 LOCAL_RANK | 多进程本地 rank |
--sync_bn | 关闭 | 启用 Apex 同步 BN |
--opt-level | 无 | Amp 优化级别(O0/O1/O2/O3) |
--keep-batchnorm-fp32 | None | 是否保持 BN 在 FP32 |
--loss-scale | None | 覆盖为静态 Loss Scaling |
--channels-last | False | 使用 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_rate,main_amp.py)。
6.3 为吞吐优化的数据管线
fast_collate:跳过transforms.ToTensor()的逐像素 PIL→float 转换,直接把uint8图像堆叠进预分配张量(支持channels_lastNHWC 格式),并同步打包 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.autocast 与 GradScaler 驱动。从源码结构可以推断,该示例正在从 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:它代表"光速"上限,但不保证收敛,仅作对比基准。 - 显存不足时降
--b:224已要求 ≥16GB 显存,128更稳妥。 - 多卡必守两条铁律:
amp.initialize先于 DDP 包装;每个进程先torch.cuda.set_device(args.local_rank)。 - 调试用
--deterministic,调优用--prof:前者保证逐位可复现,后者通过 NVTX 标记定位 CPU/GPU 时间线瓶颈。
本文所有命令与结论均可对照 examples/imagenet/README.md 与 examples/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
火山引擎视频云技术社区,是面向 AI 音视频开发者的技术交流平台。这里汇聚源自抖音、豆包等亿级 DAU 产品的 RTC、直播、点播、AI 媒体处理、音视频互动技术,提供接入指南、最佳实践、性能调优、场景案例、Demo 代码、开源项目、白皮书和 API 文档。社区汇聚官方工程师与一线开发者,为 AI 视频通话、数字人、AI 视频处理等应用的开发与落地提供技术支持。
更多推荐
所有评论(0)