先记住一句话
Adam 每个可训练参数通常至少多两份 moment state;低精度 model weights 并不会自动让这些 state 也变成低精度。
1. 每参数内存账本
一个常见 mixed-precision Adam 配置可能包含:
| 对象 | 典型精度 | 字节/参数 |
|---|---|---|
| model parameter | BF16/FP16 | 2 |
| gradient | BF16/FP16 或 FP32 | 2 或 4 |
| FP32 master parameter | FP32(实现相关) | 4 |
| Adam m | FP32 | 4 |
| Adam v | FP32 | 4 |
因此常见估算是 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 sharding | m/v 等 state |
| + gradient sharding | state + gradients |
| + parameter sharding | state + 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。