先记住一句话

Adam 每个可训练参数通常至少多两份 moment state;低精度 model weights 并不会自动让这些 state 也变成低精度。

列 tensor 账本选 precision估 activation选 shard stage量通信验证 resume
ModelBF16 params
+
Backwardgrads
+
OptimizerFP32 m,v
+
Runtimeactivations/buffers
PeakGPU memory
先按对象算理论下限,再用 profiler 找 temporary、fragmentation 与通信峰值。

1. 每参数内存账本

一个常见 mixed-precision Adam 配置可能包含:

对象典型精度字节/参数
model parameterBF16/FP162
gradientBF16/FP16 或 FP322 或 4
FP32 master parameterFP32(实现相关)4
Adam mFP324
Adam vFP324

因此常见估算是 12–16 bytes/trainable parameter,尚未算 activation、temporary buffers、allocator fragmentation 和 communication buckets。具体实现可能没有 master copy 或使用不同 state dtype,必须实测。

2. BF16 与 FP16

BF16 exponent range 接近 FP32,较少 overflow/underflow,但 mantissa 精度低;FP16 mantissa 更多、range 小,常需 dynamic loss scaling。optimizer states 通常保留 FP32,以免小 update 在低精度累积中丢失。

3. Loss scaling 的顺序

scaled_loss.backward()
unscale_(optimizer)
check non-finite gradients
clip_grad_norm_()
optimizer.step()
scaler.update()

若出现 Inf/NaN,应跳过 update,不让污染 gradient 写入 m/v。BF16 常不需 loss scaling,但计算图中的某些 operations 仍可能 overflow。

4. for-loop、foreach 与 fused

  • for-loop:逐 tensor 发 kernel,简单但 launch 多;
  • foreach:multi-tensor operation 降 launch overhead,可能用额外 tensorlist memory;
  • fused:把 update 多步融合进少量 kernel,通常最快,但 dtype/device/feature 支持有限。

adamw_torch_fused 主要改变 implementation throughput,不改变 AdamW 数学目标;仍要验证数值与 checkpoint compatibility。

5. 8-bit optimizer state

把 m/v block-wise quantize 可显著省 state memory,通常对大 dense tensors 最有效;小 tensor/outlier 可能保留高精度。它减少存储/带宽,不等于 model forward 变 8-bit;quantization error 和 kernel support 需要实测。

6. ZeRO/FSDP 在 shard 什么

级别跨 data-parallel ranks 分片
Optimizer-state shardingm/v 等 state
+ gradient shardingstate + gradients
+ parameter shardingstate + gradients + parameters

省显存换来 all-gather/reduce-scatter、bucket 与 overlap 复杂度。通信量、network topology 和小 tensor latency 会决定实际速度。

7. Gradient accumulation 省什么

它允许小 microbatch forward/backward,主要减少 activation 峰值;不直接减少每个参数的 optimizer state,也不让一次 optimizer step 更便宜。activation checkpointing 用重算换 activation memory;两者与 state sharding 是不同维度。

8. Checkpoint 为什么巨大

完整 resume 需要 model、optimizer states、scheduler、scaler、RNG 和 data position。Adam checkpoint 可比纯 weights 大数倍。分布式 checkpoint 还要记录 sharding metadata;换 world size/optimizer 时需明确 reshard 或只做 weights-only initialization。

9. 该同时优化哪些指标

  • peak allocated/reserved memory 与 fragmentation;
  • tokens/s、optimizer-step time、communication overlap;
  • validation loss per token 与 per wall-clock;
  • checkpoint save/load time、size 与 resume equivalence;
  • NaN skip、loss-scale trajectory 与 state dtype。

10. 四个系统账本计算

例 1:十亿参数 Adam 账本

1B 参数若有 BF16 weight 2 GB、BF16 gradient 2 GB、FP32 master 4 GB、FP32 m/v 各 4 GB,总计 2+2+4+4+4=16 GB,还没算 activation 与临时 buffer。

例 2:loss scaling 的数值

真实 FP32 gradient=2×10−8,loss scale=32768,backward 中变成 2e−8×32768=6.5536e−4;optimizer 前再除 32768 回到 2e−8,然后才做 clipping。

例 3:optimizer-state sharding

m/v 共 8 GB,在 8 ranks 完全均分时理论上每卡只放 8/8=1 GB;但 parameter、gradient、gather buffer 和通信 bucket 另算,所以峰值不会简单变成单卡的 1/8。

例 4:accumulation 改 batch 不改 state

4 GPUs、microbatch=2、accumulate=8,global batch=4×2×8=64。每次 forward 只保留 2 samples 的 activation 峰值,但 m/v 仍为全部参数各一份(或各自 shard)。

常见误解:

“fused AdamW 比 AdamW 收敛更快”要分清含义:它通常是每 step wall-clock 更快,不应被描述成需要更少 optimization steps,除非有对齐数值的实验。

自测

1. 为什么 Adam state 常比 BF16 weights 大很多?

m/v 往往各用一份 FP32 tensor,可能还有 FP32 master parameter。

2. accumulation 会减少 m/v 内存吗?

不会;它减少每个 microbatch activation 峰值,optimizer state 仍按全部 trainable parameters 保存。

3. sharding 的主要代价是什么?

额外通信、临时 gather buffer、调度复杂度和 checkpoint reshard。

官方资料

PyTorch Optimizer 文档说明 per-parameter options、foreach/fused implementation;PyTorch FSDP说明参数、gradient 与 optimizer state sharding。