torch.optim.Optimizer的注意事项?如何实现自定义的分布式优化器?继承torch.optim.Optimizer的注意事项?
torch.optim.Optimizer 要求子类:
__init__ 里调 super().__init__(params, defaults),并保存 self.param_groups(及每个 group 的 param 列表和超参)。step(closure=None):遍历 self.param_groups 里每个 param,用其 param.grad 按你的算法更新 param.data;若支持 closure(如 LBFGS),可执行 closure 再更新。zero_grad();由用户在 backward 后、step 前自己调,或你在 step 里可选地调。分布式下:梯度通常已由 DDP/FSDP 等做 all-reduce,每个 rank 上的 param.grad 已是「全局梯度」;自定义优化器只需在 本 rank 上按梯度更新本 rank 的参数(FSDP 下参数是分片的,每卡只更新自己的分片)。所以「分布式」部分多数由 DDP/FSDP 解决,优化器侧主要是正确读 grad、写 data。
param.data 和 param.grad 的 device/dtype 要一致;混合精度时 grad 可能是 fp16,若优化器用 fp32 状态,要转成 fp32 再更、写回时再转(或用 master weight)。step(closure),应先执行 closure() 再更新,并返回 closure 的返回值(若需要);多数优化器 closure 可为 None。backward()(DDP 会 all-reduce),再 optimizer.step();若用 AMP,先 scaler.unscale_(optimizer)、检查 inf/nan,再 optimizer.step()、scaler.update()。__init__(params, defaults) 和 step();step 里遍历 param_groups、用 grad 更新 data。| 返回模块 | 返回总览 |