mindspore-ai / mindspore-ai/hyper-parallel
支持FP32优化器
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
AutoModels Trainer 支持 FP32 main_param 优化器
本文按功能方案(Request for Comments,RFC)组织,分为“测试方案”和“开发方案”。
实现依据:合并请求 !1262。本能力属于 PyTorch AutoModels Trainer 路径。
0. 范围和术语
模型使用 16 位浮点(float16)或脑浮点 16 位(bfloat16)参数训练时,如果优化器直接修改低精度参数,较小的更新量可能在写回时被舍入掉。本功能为低精度模型参数维护一份 32 位浮点(float32,下文简称 FP32)的 main_param。
- 模型计算参数:模型前向和反向实际使用的参数,可以是
float16、bfloat16或float32。 main_param:优化器实际更新的 FP32 参数。main_grad:全分片数据并行归约后交给main_param的 FP32 梯度。- 全分片数据并行(Fully Sharded Data Parallel,FSDP):把参数和梯度分片到多个设备的并行方式。
- 分布式检查点(Distributed Checkpoint,DCP):保存和恢复模型、优化器及训练状态的机制。
本功能支持 AdamW 和 Muon。使用 Muon 时,不适合 Muon 更新的参数仍按原有规则交给 AdamW;下文简称“Muon + AdamW”。
测试方案
1. 最终可见配置
下面展示 Muon + AdamW 场景。只使用 AdamW 时,将 _target_ 改为 AdamW 并保留 fp32_main_params 和 adamw_config 即可。
model_init_dtype: bfloat16
optimizer:
_target_: hyper_parallel.auto_models.components.optim.optimizer.optimizer.Muon
fp32_main_params: true
muon_config:
muon_lr: 0.02
muon_weight_decay: 0.0
muon_momentum: 0.95
muon_nesterov: true
muon_ns_steps: 5
muon_ns_variant: asym5
adamw_config:
adamw_lr: 0.001
adamw_weight_decay: 0.01
fsdp_config:
mix_precision:
param_dtype: bfloat16
reduce_dtype: float32
output_dtype: bfloat16
cast_forward_inputs: true
关键配置如下:
model_init_dtype:可配置为float16、bfloat16、float32或null,默认null。它决定模型构建或加载权重完成后的浮点参数和浮点缓冲区精度;整数和布尔缓冲区不转换。optimizer.fp32_main_params:默认false。设置为true后,低精度模型参数使用独立的 FP32main_param;原生 FP32 模型参数直接把自身作为main_param。fsdp_config.mix_precision.reduce_dtype:开启fp32_main_params时必须显式配置为float32,系统不会自动补齐。param_dtype、output_dtype和cast_forward_inputs:分别控制 FSDP 计算参数精度、输出精度和浮点输入转换行为。
开启后的训练流程为:
低精度模型参数完成前向和反向
→ FSDP 使用 float32 归约梯度并写入 main_grad
→ 梯度裁剪
→ Muon + AdamW 更新 main_param
→ main_param 转换并回写模型计算参数
→ 清理 grad 和 main_grad
2. 核心看护用例
2.1 用例一:bfloat16 模型的 Muon + AdamW 训练精度
验证场景
model_init_dtype: bfloat16
optimizer:
_target_: hyper_parallel.auto_models.components.optim.optimizer.optimizer.Muon
fp32_main_params: true
fsdp_config:
mix_precision:
param_dtype: bfloat16
reduce_dtype: float32
output_dtype: bfloat16
使用固定模型初始化种子和固定输入数据,连续执行多个训练步骤。参考路径使用相同模型、相同 Muon + AdamW 参数分组和相同 main_param 包装,但不进行分布式切分;被测路径启用实际并行和 FSDP。
效果
低精度模型负责前向和反向,Muon 与 AdamW 都使用 FP32 main_param 和 main_grad 完成更新。分布式切分不改变训练语义,更新后的模型参数与未切分参考路径在设定容差内一致。
校验点
- 模型浮点参数和浮点缓冲区为
bfloat16。 - 每个低精度可训练参数都有独立的 FP32
main_param。 - FSDP 归约后
main_grad为float32,普通grad为空。 - Muon 参数组和 AdamW 参数组都实际执行更新,且引用的是
main_param。 - 每步比较损失、全局梯度范数、逐参数
main_grad和更新后的模型参数。 - 每步更新后,模型参数等于对应
main_param转换为bfloat16后的值。
2.2 用例二:bfloat16 模型保存并恢复模型与优化器状态
验证场景
使用与用例一相同的 bfloat16、fp32_main_params: true 和 Muon + AdamW 配置。保存端配置 checkpoint.save_optimizer: true,加载端通过 checkpoint.restore_from 指向本次保存路径,并配置 checkpoint.restore_optimizer: true:
- 使用固定种子训练 1 步。
- 保存模型权重、Muon 状态、AdamW 状态和
main_param。 - 新建相同配置的 Trainer,加载本次保存的模型和优化器状态。
- 使用相同的下一批数据继续训练。
- 与不经过保存和加载、连续训练相同步数的参考路径比较。
效果
检查点精确恢复低精度模型无法表示的 FP32 main_param 尾数,同时恢复 Muon 动量、AdamW 一阶和二阶矩以及优化器步骤。加载后的下一步训练与连续训练路径保持一致。
校验点
- 模型状态按
bfloat16保存和恢复。 - 优化器状态包含
_mixed_precision_optimizer.fp32_from_fp16_params,其中覆盖全部低精度可训练参数,值为float32。 - Muon 动量、AdamW 一阶和二阶矩、参数组超参数及步骤计数均完成恢复。
- 加载后的
main_param与保存时逐参数一致,不能先经过bfloat16再重建。 - 续训下一步的损失、梯度、
main_param和模型参数与连续训练参考路径在设定容差内一致。
2.3 用例三:float32 模型不保存冗余 main_param
验证场景
model_init_dtype: float32
optimizer:
_target_: hyper_parallel.auto_models.components.optim.optimizer.optimizer.Muon
fp32_main_params: true
fsdp_config:
mix_precision:
param_dtype: float32
reduce_dtype: float32
output_dtype: float32
使用 Muon + AdamW 训练 1 步,保存模型权重和优化器状态;保存端配置 checkpoint.save_optimizer: true。随后新建相同配置的 Trainer,通过 checkpoint.restore_from 和 checkpoint.restore_optimizer: true 加载保存状态并继续训练,与连续训练参考路径比较。
效果
模型参数本身已经是 FP32,因此直接作为 main_param,不创建冗余副本。检查点仍能使用统一的混合精度优化器状态结构,但不会额外保存任何 main_param 张量。加载后可恢复 Muon + AdamW 状态并继续训练。
校验点
- 所有浮点模型参数均为
float32,且model_param.main_param is model_param。 - 不存在由模型参数复制出的额外
main_param对象或存储。 _mixed_precision_optimizer.fp32_from_fp16_params映射为空,不包含额外参数张量。- 模型权重、Muon 状态、AdamW 状态和步骤计数均完成恢复。
- 续训下一步的损失、梯度和模型参数与连续训练参考路径在设定容差内一致。
2.4 用例四:普通优化器检查点兼容与不完整检查点保护
验证场景
分别构造以下两种输入:
- 检查点完全不包含
_mixed_precision_optimizer,但包含完整模型和 Muon + AdamW 原有状态。 - 检查点包含
_mixed_precision_optimizer.fp32_from_fp16_params,但缺少部分参数。
效果
普通优化器检查点可以加载,main_param 从已加载模型参数刷新,并输出精度兼容警告;只包含部分 main_param 的检查点被视为损坏数据,直接停止加载。
校验点
- 普通检查点能够恢复 Muon + AdamW 原有状态,并从当前模型参数建立
main_param。 - 普通检查点路径输出明确警告,说明低精度模型已丢失的 FP32 尾数无法恢复。
- 部分
main_param检查点明确列出缺失参数,不允许静默补齐或回退。
2.5 用例五:配置约束
验证场景
保持 optimizer.fp32_main_params: true,分别不配置 reduce_dtype,或将其配置为 bfloat16。
效果
配置在训练启动前失败,不进入模型训练和检查点流程。
校验点
- 两种配置都返回明确错误。
- 错误信息同时指出
optimizer.fp32_main_params和fsdp_config.mix_precision.reduce_dtype。 - 系统不能静默把
reduce_dtype改成float32。
3. 现有用例覆盖说明
当前 8 卡昇腾系统测试使用小型 Qwen2 混合专家模型、bfloat16 模型参数、FP32 梯度归约和 AdamW,组合数据并行、上下文并行、张量并行、专家并行与 FSDP,连续训练 20 步。用例比较未切分参考路径和分布式路径的损失、梯度范数、逐参数 main_grad 及模型参数回写结果。
该系统测试已经覆盖 FSDP 交付 main_grad、AdamW 更新 main_param 和多维并行数值对齐。Muon 参数替换、Muon/AdamW 状态、main_param 保存恢复、普通检查点兼容和部分键缺失目前已有单元测试覆盖;上述核心看护用例用于补齐 Muon + AdamW 组合及完整 Trainer 保存、加载、续训场景。
基线统一使用相同模型结构、初始权重、输入数据、优化器超参数和 main_param 包装。训练精度用例比较未切分参考路径与分布式路径;保存恢复用例比较连续训练路径与“训练 1 步后保存、加载并续训”的路径。
开发方案:设计与实现
4. 设计目标与边界
设计目标是在不修改 AdamW 和 Muon 更新公式的前提下,为低精度模型提供 FP32 main_param 更新能力,并保持现有参数路由、学习率调度、梯度裁剪和分布式检查点接口不变。
核心设计原则如下:
- 模型计算参数继续用于前向和反向;低精度参数由包装器建立独立的 FP32
main_param。 - AdamW 和 Muon 只操作传入参数,不分别实现
main_param分支。包装器直接替换叶子优化器参数组中的参数对象。 - FSDP 根据同一个
optimizer.fp32_main_params策略把归约梯度写入main_grad,避免产生低精度梯度后再转换。 - 原生 FP32 模型参数直接作为
main_param,不创建副本,也不在优化器检查点中重复保存。 - 检查点显式保存低精度参数对应的 FP32
main_param,同时兼容不包含该状态的普通优化器检查点。 - AdamW、Muon 和检查点继续使用模型参数完整名称组织状态,不因内部参数替换改变对外键名。
本方案属于 PyTorch AutoModels Trainer 路径,支持 AdamW,以及由 Muon 和回退 AdamW 组成的组合优化器。
5. 配置模型与运行时传递
5.1 优化器配置解析
fp32_main_params 属于 Trainer 管理的精度策略,不作为 AdamW 或 Muon 构造参数。解析器先从 optimizer 节点取出该字段,再按原 _target_ 规则解析其余配置:
def resolve_optimizer(optimizer_yaml):
enable_fp32_main_params = optimizer_yaml.pop("fp32_main_params", False)
optimizer_target = resolve_target(optimizer_yaml)
return OptimizerConfig(
target=optimizer_target,
fp32_main_params=enable_fp32_main_params,
)
解析后的关键结构为:
TrainerConfig
├── model_init_dtype
├── optimizer
│ ├── target
│ └── fp32_main_params
└── fsdp_config.mix_precision
├── param_dtype
├── reduce_dtype
├── output_dtype
└── cast_forward_inputs
OptimizerConfig.to_dict() 会把 fp32_main_params 放回优化器配置视图,使运行日志能展示默认值和最终值。
5.2 跨配置校验
TrainerConfig.__post_init__() 负责跨配置约束:
if optimizer.fp32_main_params and fsdp_mix_precision.reduce_dtype != "float32":
raise ValueError(
"optimizer.fp32_main_params=true requires "
"fsdp_config.mix_precision.reduce_dtype='float32'"
)
系统不自动改写 reduce_dtype。这样可以保证配置文件、解析日志和实际通信精度一致,也能在创建模型和通信组之前报告错误。
5.3 传入 FSDP
创建分布式运行环境时,cfg.optimizer.fp32_main_params 会写入 DistributedSetup,再传给 FSDP2Manager。管理器构建混合精度策略时,同时设置:
policy = MixedPrecisionPolicy(
param_dtype=resolve_dtype(config.param_dtype),
reduce_dtype=resolve_dtype(config.reduce_dtype),
output_dtype=resolve_dtype(config.output_dtype),
cast_forward_inputs=config.cast_forward_inputs,
apply_grad_on_fp32_main_grad=fp32_main_params,
)
因此 Trainer、优化器包装器和 FSDP 使用同一个开关,不会出现梯度存储方式与优化器读取方式不一致。
6. 模型精度转换设计
6.1 正常构建路径
Trainer 的模型和优化器构建顺序为:
构建模型并完成初始化或权重加载
→ 完成并行切分和 FSDP 包装
→ apply_model_init_dtype(model, model_init_dtype)
→ 构建 AdamW 或 Muon + AdamW
→ 按 fp32_main_params 增加 main_param 包装器
优化器看到的是已经转换到最终配置精度的模型参数,因此包装器可以仅根据参数实际 dtype 决定建立副本还是直接复用。
6.2 检查点加载路径
分布式检查点把模型权重写入当前模型后,Trainer 再调用一次 apply_model_init_dtype():
DCP 加载模型权重
→ model.load_state_dict()
→ apply_model_init_dtype()
→ 加载优化器状态
模型最终精度由当前训练配置决定,与检查点内模型张量的保存精度解耦。优化器检查点中独立保存的 FP32 main_param 保持 FP32,不受模型精度转换影响。
6.3 转换约束
apply_model_init_dtype() 使用模块精度转换接口处理所有浮点参数和浮点缓冲区,同时保持整数和布尔缓冲区不变。转换前后会校验:
Parameter对象身份没有变化;- 分布式张量的设备网格和分片布局没有变化;
- 浮点参数和浮点缓冲区均达到目标精度;
- 延迟分配实际存储的
meta参数可以在不替换对象的情况下转换精度; - 已完成 FSDP 包装的参数刷新分片存储及精度元数据。
这些约束保证优化器参数路由、参数共享关系和分布式布局不会因精度转换失效。
7. 优化器组合与参数替换
7.1 组合关系
外层组合关系如下:
Float16OptimizerWithFloat16Params
└── ChainedOptimizer
├── AdamW
└── Muon(使用 Muon 配置时存在)
只配置 AdamW 时,ChainedOptimizer 中只有 AdamW。配置 Muon 时,矩阵参数按原有规则交给 Muon,其余参数交给 AdamW;外层包装器在参数路由完成后统一处理两类参数组。
7.2 初始化顺序
inner_optimizer = optimizer_target.build(model=model).get_optimizer()
if not config.optimizer.fp32_main_params:
return inner_optimizer
return Float16OptimizerWithFloat16Params(inner_optimizer, model)
包装器遍历每个叶子优化器参数组,并维护三类分组:
float16_groups
模型中的 float16 或 bfloat16 计算参数
fp32_from_float16_groups
为上述低精度参数建立的 FP32 main_param
fp32_from_fp32_groups
模型中原本就是 float32 的参数
低精度参数处理逻辑为:
main_param = Parameter(
model_param.detach().to(dtype=float32),
requires_grad=model_param.requires_grad,
)
model_param.main_param = main_param
optimizer_group["params"][index] = main_param
if model_param in leaf_optimizer.state:
leaf_optimizer.state[main_param] = leaf_optimizer.state.pop(model_param)
原生 FP32 参数不复制:
model_param.main_param = model_param
fp32_from_fp32_group.append(model_param)
不属于 float16、bfloat16 或 float32 的可训练参数会直接报错,避免优化器静默处理不支持的参数精度。
7.3 参数映射和 Muon 缓存重建
参数替换后,ChainedOptimizer 维护双向映射:
optimizer_param → model_param
model_param → optimizer_param
该映射用于:
- 让学习率调度和日志继续使用模型参数完整名称;
- 保存或加载优化器状态时,临时把
main_param注册到原模型参数位置; - 让 AdamW 和 Muon 状态继续按模型参数完整名称组织。
Muon 在构建时会按参数对象建立分组、分片信息、批次信息、广播信息和步骤分类缓存。替换参数后,reset_optimizer_parameters() 会重建这些缓存,确保运行时不再引用已被替换的模型计算参数。
AdamW 不依赖上述分布式缓存。包装器替换其参数组并迁移已有状态后,AdamW 后续自然读取 main_param.grad 并更新 main_param。
8. main_grad 与每步更新生命周期
8.1 FSDP 梯度写入
关闭 fp32_main_params 时,FSDP 保持普通流程,把归约梯度写入 model_param.grad。
开启后,FSDP 使用配置的 FP32 归约精度,并执行:
reduced_grad = reduce_gradient_in_float32()
if model_param.main_grad is None:
model_param.main_grad = reduced_grad
else:
model_param.main_grad += reduced_grad
model_param.grad = None
多个微批次的梯度直接累积到 main_grad,不会先保存为低精度 grad 再转换。分布式梯度裁剪也优先读取 main_grad。
8.2 优化器步骤
Trainer 仍调用统一的 optimizer.step() 和 optimizer.zero_grad()。包装器内部流程如下:
def step():
prepare_grads()
return step_with_ready_grads()
def prepare_grads():
for model_param, main_param in low_precision_parameter_pairs:
gradient = model_param.main_grad
if gradient is None:
gradient = model_param.grad
main_param.grad = cast_to_float32(gradient)
model_param.grad = None
for model_param in native_float32_parameters:
if model_param.main_grad is not None:
model_param.grad = model_param.main_grad
def step_with_ready_grads():
chained_optimizer.step()
copy_main_params_to_model_params()
def zero_grad():
chained_optimizer.zero_grad()
clear_model_grad_and_main_grad()
叶子优化器更新完成后,包装器将 FP32 main_param 转换为模型参数精度并回写,供下一次前向计算使用。
8.3 反向刷新
reload_model_params() 执行相反方向的复制,用当前模型参数刷新低精度参数对应的 FP32 main_param:
for model_param, main_param in low_precision_parameter_pairs:
main_param.copy_(model_param.to(dtype=float32))
该接口只用于检查点没有提供完整 main_param,或配置为不恢复优化器状态的场景。检查点已经包含完整 main_param 时不能调用,否则会丢失只能由 FP32 表示的尾数。
9. 分布式检查点设计
9.1 保存结构
包装器在原优化器状态中追加:
optimizer
├── AdamW 原有状态,例如 exp_avg、exp_avg_sq
├── Muon 原有状态,例如 momentum_buffer 和 step
└── _mixed_precision_optimizer
└── fp32_from_fp16_params
└── <模型参数完整名称>: <float32 main_param>
只保存由 float16 或 bfloat16 模型参数建立的 FP32 main_param。原生 FP32 模型参数已经存在于模型状态中,因此 fp32_from_fp16_params 保持为空,不重复保存参数张量。运行时 main_grad 不写入检查点。
ChainedOptimizer.state_dict() 在收集 AdamW 和 Muon 状态时,使用参数双向映射临时恢复模型参数完整名称;Muon 分布式状态完成必要的复制组同步后,再与 AdamW 状态合并。包装器最后追加 main_param,由同一次 DCP 保存操作写出。
9.2 加载前初始化优化器状态
新建 AdamW 或 Muon 时,一阶和二阶矩、动量等延迟状态尚不存在。DCP 只能把检查点张量填入当前已经准备好的状态结构,因此加载前需要创建对应位置。
初始化过程对每个叶子优化器分别执行:
save_lr_and_weight_decay()
set_lr_and_weight_decay_to_zero()
set_zero_grad_for_all_params()
optimizer.step()
restore_lr_and_weight_decay()
optimizer.zero_grad()
学习率和权重衰减临时为零,因此该步骤只创建优化器状态,不改变参数。外层存在 Float16OptimizerWithFloat16Params 时,初始化逻辑会找到内部 ChainedOptimizer,再逐个处理 AdamW 和 Muon。
9.3 加载分支
DCP 加载规划器先比较当前状态结构需要的 main_param 键与检查点实际键,然后选择以下分支。
分支一:检查点包含全部 main_param
DCP 加载模型、AdamW/Muon 状态和全部 main_param
→ model.load_state_dict()
→ 按 model_init_dtype 转换模型状态
→ wrapper.load_state_dict()
→ 恢复 AdamW/Muon 状态和检查点中的 FP32 main_param
该分支不调用 reload_model_params(),因此能够精确恢复无法经低精度模型参数往返的 FP32 尾数。
分支二:检查点完全没有混合精度优化器状态
加载规划器移除当前接收结构中的 mixed-precision 分支
→ DCP 加载模型和 AdamW/Muon 原有状态
→ 按 model_init_dtype 转换模型状态
→ wrapper.load_state_dict()
→ reload_model_params()
→ 输出精度兼容警告
该分支兼容普通优化器检查点,但只能从当前模型参数重建 main_param。
分支三:检查点只包含部分 main_param
该情况属于检查点内容不完整。加载规划器在写入模型和优化器状态前直接报错,并列出缺失的模型参数完整名称,不允许回退到 reload_model_params()。
9.4 只恢复模型
当 checkpoint.restore_optimizer: false 时,Trainer 不构造优化器加载状态。模型权重加载和精度转换完成后,回调调用 reload_model_params(),使当前包装器与已加载模型参数一致。该路径不恢复 AdamW 的一阶和二阶矩、Muon 动量或优化器步骤。
10. 可观察性
为了区分用户配置、默认值和运行时结果,运行日志包含:
Trainer config:解析完成后的完整 Trainer 配置,包括model_init_dtype和fp32_main_params默认值或最终值;Model config:模型构建完成后的模型配置;Effective adamw config:AdamW 构造器默认值、配置覆盖值和运行时最终值;Effective muon config:Muon 构造器默认值、配置覆盖值和运行时最终值。
优化器日志由 AdamW/Muon 创建过程生成,与是否增加外层包装器无关。
11. 主要实现位置
| 文件 | 作用 |
|---|---|
hyper_parallel/auto_models/trainer/config.py |
定义 model_init_dtype、OptimizerConfig.fp32_main_params 和跨配置校验。 |
hyper_parallel/auto_models/config/resolver.py |
从优化器 YAML 中独立解析 fp32_main_params。 |
hyper_parallel/auto_models/trainer/model_init_dtype.py |
转换模型精度,并校验参数身份、分布式布局和最终精度。 |
hyper_parallel/auto_models/trainer/base.py |
串联模型精度转换、优化器构建和 main_param 包装。 |
hyper_parallel/auto_models/components/distributed/config.py |
定义 FSDP 混合精度配置。 |
hyper_parallel/auto_models/components/distributed/fsdp2.py |
将配置精度和 main_grad 策略传给 FSDP。 |
hyper_parallel/auto_models/components/optim/optimizer/mixed_precision_optimizer.py |
实现参数分组、梯度准备、更新回写、清理和状态保存恢复。 |
hyper_parallel/core/optimizer/optimizer.py |
维护优化器参数与模型参数映射,并保持状态使用模型参数完整名称。 |
hyper_parallel/core/optimizer/muon.py |
参数替换后重建 Muon 运行时缓存。 |
hyper_parallel/auto_models/components/checkpoint/dcp_checkpointer.py |
初始化优化器状态并选择检查点加载分支。 |
hyper_parallel/auto_models/trainer/callbacks/checkpoint_callback.py |
组织模型精度转换、优化器恢复和只恢复模型时的参数刷新。 |
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 352
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/352
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with TrainerConfig.post_init and the optimizer construction path, then trace DistributedSetup into FSDP2Manager and Float16OptimizerWithFloat16Params. Review the existing Muon, AdamW, and distributed-checkpoint tests before adding coverage for FP32 main parameters, main_grad handling, configuration validation, and save/restore continuation. Done means the listed Muon + AdamW scenarios and compatibility checks pass without changing optimizer update formulas.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- distributed-systems, machine-learning
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Active
- Clarity
- Needs clarification
- Newbie friendliness
- 25/100