已关闭
[RFC]:TENPU FusedAdam 的 DTensor 分片更新与 FSDP 状态适配 #43
baichuan51141创建于 13 天前关闭于 13 天前
13 天前 添加了label:rfc
13 天前 关联了pull request:feat: support DTensor shards in FusedAdam for MindSpeed FSDP
13 天前 关闭了 issue
13 天前 issue状态由 TODO 改变为 DONE
13 天前 添加了label:resolved
RFC:TENPU FusedAdam 的 DTensor 分片更新与 FSDP 状态适配
基本信息
状态(Status): Draft
作者(Authors): @GuoHaifeng1999
创建日期(Created): 2026-09-14
更新日期(Updated): 2026-09-14
1. 概述
1.1 简介
本提案在 TransformerEngineNPU 的已有
FusedAdam中实现 MindSpeed/Megatron FSDP 所需的 DTensor 支持。沿用 NVIDIA Transformer Engine(下称 NVTE)的基本结构:参数与优化器状态保留分布式表示,执行优化器内核前提取本地 Tensor,保存状态时保留 mesh、placements、全局 shape 和 stride。本次修复已提交到
5c85c0df3e9c1887a4aa5a000f96b2f01d9fb031。仅修改fused_adam.py和既有test_fused_optimizer.py,未修改优化器选择入口、未新增优化器类,也未新增测试文件或测试函数。Core 的分片、HCCL 和检查点入口见 Core FSDP RFC。核心价值:
fsdp_dtensor流程使用get_unscaled_state(..., skip_unscale=False)参数1.2 动机
使用场景
optimizer.step()报 DTensor 不支持。decoupled_grad。当前痛点
_check_unsupported_param_type()对所有 DTensor 抛出NotImplementedError。grad与decoupled_grad必须在同一分片语义下处理。abs().max()不适用于空 Tensor。initialize_state要求第二个参数,而 Core 0.18 存在单参数调用。必要性
用户价值
transformer_engine.pytorch.optimizers.FusedAdam入口可直接消费 FSDP 参数。1.3 目标
目标范围
_local_tensor()和_wrap_state(),集中处理本地视图与分布式包装。None梯度,保持 group step 语义。state_dict()/load_state_dict()的浮点状态语义与分布式元数据。initialize_state的兼容性例外。非目标(边界说明)
bias_correction=False能力。initialize_state的可选参数是明确保留的差异。2. 用例分析
2.1 功能点描述
to_local()供已有更新内核使用param.decoupled_gradskip_unscale并记录签名例外2.2 关键性能指标
以下为设计目标及验收口径;当前没有 NPU 性能数据。
2.3 安全隐私要求
2.4 DFX 要求
兼容性
ccb9ae5,本次实现固定为5c85c0d。42b84005,不将该快照视为所有 TE 版本的统一签名。init_state_fn存在opt.initialize_state(p)调用。可维护性
self.state和 scale 字典的 key。param_groups,不复制整个优化器实现。可测试性
torchrun下执行。可靠性
Partial参数布局需要先完成归约,优化器不隐式补通信。3. 方案设计
3.1 总体方案
设计思路
采用“分布式状态保存—本地算术执行”的结构。外部参数仍为 DTensor,状态字典中的 moments/master 也按同一分布式布局包装;数值更新只消费 local Tensor。对普通 Tensor,helper 直接返回输入,使其继续使用原有实现。
DTensor 的本地 Tensor 不是完整全局参数,也不能任意更改其分片位置。由 Core 提供正确布局和已归约梯度,TENPU 在该布局上做逐元素 Adam 更新,因此无需在优化器内部进行参数收集。低精度 scale 作用于各 rank 本地状态,不在优化器中做跨 rank 最大值归约。
架构图
classDiagram class CoreFSDP { +optimizer_named_parameters +update_main_grads() } class DTensor { +device_mesh +placements +shape +stride() +to_local() } class FusedAdam { +param_groups +state +_scales +initialize_state(param, store_param_remainders) +get_unscaled_state(param, name, skip_unscale) +set_scaled_state(param, name, value) +step() +state_dict() +load_state_dict(state_dict) -_local_tensor(tensor) -_wrap_state(param, value) } class LocalUpdate { +params +grads +exp_avgs +exp_avg_sqs +fused_adamw() } class CoreCheckpoint { +fsdp_dtensor +save() +load() } CoreFSDP --> DTensor : 构造参数和梯度 FusedAdam --> DTensor : 保存参数与状态布局 FusedAdam --> LocalUpdate : 仅传本地 Tensor CoreFSDP --> FusedAdam : 调用优化器 CoreCheckpoint --> FusedAdam : 保存和恢复状态核心逻辑流程
flowchart TD A[进入 FusedAdam step] --> B[已有 group 校验及 step 推进] B --> C[读取 grad 或 decoupled_grad] C --> D{梯度是否为 None} D -->|是| E[跳过该参数] D -->|否| F{参数是否为 DTensor} F -->|是| G[验证梯度布局并取 local grad] F -->|否| H[沿用原梯度] G --> I[初始化或修正状态] H --> I I --> J{DTensor local shard 是否为空} J -->|是| E J -->|否| K{满足 direct 条件} K -->|是| L[本地参数和 FP32 moments 直接更新] K -->|否| M[本地解缩放及 master 工作区更新] L --> N[完成当前分组] M --> O[状态重新缩放并包装 DTensor] O --> P[本地参数写回] P --> N3.2 技术选型
方案对比
选型理由
3.3 功能与性能设计
3.3.1 类型支持范围与前置条件
capturable=Trueamsgrad=Truebias_correction=False支持范围表示代码具备相应路径;一维 Shard/Replicate 已进入既有测试参数化,多维 mesh、TP/DP 复合布局仍需集成验收。构造器及 group 校验的原有约束继续生效,例如同一 group 的某些 FP16/BF16 混用限制。
3.3.2 本地 Tensor 与状态包装
@staticmethod def _local_tensor(tensor): return tensor.to_local() if isinstance(tensor, DTensor) else tensor @staticmethod def _wrap_state(param, value): if not isinstance(param, DTensor): return value return DTensor.from_local( value, device_mesh=param.device_mesh, placements=param.placements, shape=param.shape, stride=param.stride(), run_check=False, )to_local()提供本地数据视图,不执行全局收集。from_local()使用参数的全局 shape 和 stride,而不是根据本地 shard 大小猜测全局布局,这对非均匀分片尤其重要。run_check=False避免在包装状态时增加检查通信,也意味着调用方必须提供正确布局。helper 不负责修复错误分片或验证全局 rank 间一致性;本次没有承诺任意 local Tensor 都可以安全包装成该参数的状态。3.3.3 参数、梯度及状态的数据结构
param_groupsself.state[param]exp_avg、exp_avg_sq、可选master_param_wrap_state()包装_scales[param][name]_step_tensorsgroup["step"]一个参数的典型状态关系为:
状态 Tensor 的身份并非承诺固定:当前
set_scaled_state()的部分分支会重新创建状态和 DTensor 包装,继承原 TENPU 的赋值式更新。调用方应按参数 key 获取最新状态,不应缓存旧 state Tensor 引用并假设始终原地更新。3.3.4 梯度来源与布局校验
use_decoupled_grad=False时读取param.grad;为真时读取param.decoupled_grad。梯度为None时直接跳过该参数。对 DTensor 参数,当前检查为:if not isinstance(grad, DTensor) or ( grad.device_mesh != param.device_mesh or grad.placements != param.placements or grad.shape != param.shape ): raise RuntimeError("DTensor gradient must match the parameter shard layout.")通过检查后提取 local grad,再沿用 sparse 检查与已有数值路径。此处不校验所有可能的 local stride/底层 storage 关系,也不会自动 redistribute;布局必须由 Core 正确构造。
Partial表示尚未完成必要归约,不可当作最终梯度/参数布局进行普通逐元素更新。参数端明确拒绝 Partial,匹配布局检查进一步阻止梯度和参数采用不同 placements。3.3.5 状态初始化与低精度缩放
_init_or_fix_state()对缺失 moments 使用参数的 local Tensor 构造零状态,再调用set_scaled_state()包装。master 适用时,从本地参数转换和克隆,避免对整个 DTensor 执行不必要的分布式操作。skip_unscale=True返回本地 BF16FP16 状态沿用既有缩放公式,目标动态范围为
finfo(float16).max / 2:当
a为零或非有限值时沿用原有单位 scale 处理;本次只新增空状态返回单位 scale,避免对空 Tensor 求 max。该逻辑不额外承担非有限梯度检测或动态 loss scaling。3.3.6 Direct 与低精度更新路径
flowchart TD A[本地参数及梯度] --> B{direct 条件} B -->|满足| C[本地 FP32 moments] C --> D[原 torch fused AdamW 内核] B -->|不满足| E[解缩放 moments] E --> F[选择 master 或参数 FP32 工作区] F --> G[同一 fused AdamW 内核] G --> H[按配置缩放并存回状态] H --> I[包装 DTensor 并写回 local 参数]Direct 条件包括未启用 master weights、参数为支持的浮点 dtype、梯度 dtype 等于参数 dtype、两类 moments 均为 FP32。修改后参数和 moments 列表都通过
_local_tensor(),梯度已在布局校验后本地化。低精度路径中,FP32 moments 可直接取本地存储;其他 moments 调用
get_unscaled_state()。FP32 master 可直接更新本地存储,FP16 master 解缩放后作为工作参数;没有 master 时从本地参数取得 FP32 工作区。计算后只对需要重新缩放的状态调用 setter,普通浮点参数通过 localcopy_()写回。此路径复用已有
_apply_fused_adamw_groups()与torch._fused_adamw_()。它不增加 NVIDIA CUDA 扩展,也不添加第二套 Adam 实现;在 NPU 环境由已有 torch_npu 运行时提供对应执行能力。3.3.7 Adam/AdamW 数值语义
设当前 rank 已归约的本地梯度为
g,参数或 master 工作区为w,组级步数为t:adam_w_mode=False的 L2 模式先把weight_decay * w加入工作梯度,再令 fused AdamW 内核的 weight decay 为零。原梯度在需要时先克隆,避免为 L2 更新污染外部梯度缓冲区。以上算术仍由原内核执行,DTensor 适配只改变访问的数据范围。Core 必须先完成 DP 梯度归约;TENPU 不再除以 DP 大小,也不再次聚合不同 rank 的梯度。否则会重复平均或使用未归约梯度。
3.3.8 空分片与组级 Step
flowchart TD A[进入参数组] --> B[沿用 group step 推进] B --> C{当前参数有梯度} C -->|否| D[不初始化当前参数状态并跳过] C -->|是| E[布局校验和状态初始化] E --> F{DTensor local numel 为零} F -->|是| G[保留必要空状态并跳过内核] F -->|否| H[正常分组更新] D --> I[继续下一参数] G --> I H --> ICore 0.18 可能将空 shard 的梯度设置为
None,此时走已有跳过路径;若调用方提供了空 DTensor 梯度,则初始化状态后跳过内核。两种情况下都不得根据本地是否更新独立重置 group step,否则各 rank 在后续非空梯度恢复时会出现偏差。本次不引入 per-parameter 新计数器,继续使用原
group["step"]及_step_tensors。低精度状态缩放增加空 Tensor 分支,同时也消除了普通空状态可能触发的 reduce-max 错误。3.3.9 State Dict 保存与恢复
sequenceDiagram participant C as Core 检查点调用方 participant O as FusedAdam participant S as DTensor 状态 participant L as 本地状态 Tensor C->>O: state_dict() O->>S: 按原参数读取 moments 和 master S->>L: to_local 并按规则解缩放 L->>O: 本地 FP32 状态 O->>S: 用参数 mesh 和全局元数据重新包装 O-->>C: 含 DTensor 的状态字典 Note over C,O: 实际 DCP I/O 由 Core 负责 C->>O: load_state_dict(saved_state) O->>O: 保留原保存值并调用父类加载 O->>O: 恢复 group step 与相关缓存 O->>L: 本地化保存值并转换至目标存储 dtype O->>S: set_scaled_state 重新包装状态 O-->>C: 可继续更新的 optimizer保存时先获取父类状态字典,再复制每个已打包的参数状态字典,避免直接替换运行中的 state 条目。moments 和 master 通过
get_unscaled_state()导出,并通过_wrap_state()恢复 DTensor 元数据。浮点 moments/master 对外采用 FP32;普通 Tensor 的 INT16 remainder 特例不属于本次 DTensor 路径。加载时先按保存参数 id 建立到当前参数的映射并保留保存 Tensor 的副本,再调用父类加载。这样可以在父类发生 dtype 转换后,用保留值重建正确的浮点状态。既有逻辑重建 scale、quantizer、step Tensor 等相关缓存,并通过 setter 恢复目标精度;setter 新增 local 提取及 DTensor 包装。
state_dict()可返回正确结构并不代表 DCP 已完成磁盘保存、跨 rank 读写或重分片。本次既有测试验证的是内存态保存加载与下一步连续性,实际 I/O 必须由 Core 集成用例验收。3.3.10 TE 接口对齐与明确例外
__init__zero_gradset_to_none=Nonestepclosure=None, grad_scaler=Noneget_unscaled_stateparam, state_name, skip_unscale=Falseset_scaled_stateparam, state_name, unscaled_statestate_dict/load_state_dictstate_dictinitialize_stateparam, store_param_remainders必传param, store_param_remainders=None,返回注解为 None例外原因:已核对的 Core 0.18
init_state_fn会在 precision-aware 模式调用opt.initialize_state(p);TENPU 既有用例也存在单参数调用。直接复制 NVTE 的必传参数会使这些调用报TypeError,与保持现有功能兼容的要求冲突。当前提交保留单参数调用,并接受 TE 风格的显式第二参数。若显式值与该参数实际 remainder 配置不一致,沿用既有错误检查。此选择是调用兼容性扩展,不应写成“公开签名完全一致”。后续若要求严格一致,须先同步调用方并完成普通路径回归,见第 7.1 节。
3.3.11 普通 Tensor 路径隔离与性能影响
_local_tensor_wrap_statecopy_()copy_()新增开销主要是 Python 类型判断和必要的 DTensor 包装;低精度与检查点路径仍存在本地 Tensor 分配/复制。不能据此承诺零开销或 state Tensor 身份不变。Direct 路径没有增加全局参数收集,整体分片通信仍由 Core 调度。
3.4 安全隐私与 DFX 设计
3.4.1 兼容性设计
transformer_engine.pytorch.optimizers.FusedAdaminitialize_state兼容例外3.4.2 可维护性设计
3.4.3 可测试性设计
flowchart TD A[既有 test_fused_optimizer] --> B[普通 Tensor] A --> C[Shard] A --> D[Replicate] B --> E[手工 AdamW 参考] C --> E D --> E B --> F[低精度状态保存恢复] C --> F D --> F F --> G[恢复后下一步参数比较] C --> H[非均匀和空分片] A --> I[普通 SGD 与 Adam 回归]3.4.4 可靠性设计
self.state[param]key 不变,防止 optimizer scheduler 和调用方状态查找错位。3.5 编程与调用设计
3.5.1 编程模型基本设计
开发环境
2.9.0+cpu;以源码提取方式隔离 NPU 导入依赖。开发约束
master_weight_dtype受 TENPU 既有 FP32/FP16 限制。可验收设计
3.5.2 接口定义与设计
3.5.2.1 FusedAdam 构造与更新
接口描述:使用 TE 导入入口的 NPU Adam/AdamW 优化器。
接口原型:
class FusedAdam(torch.optim.Optimizer): def __init__( self, params: Iterable[torch.nn.Parameter | dict], lr: float = 1e-3, betas: tuple[float, float] = (0.9, 0.999), eps: float = 1e-8, weight_decay: float = 0.0, amsgrad: bool = False, *, bias_correction=True, adam_w_mode=True, capturable=False, master_weights=False, master_weight_dtype=torch.float32, exp_avg_dtype=torch.float32, exp_avg_sq_dtype=torch.float32, use_decoupled_grad=False, store_param_remainders=False, set_grad_none: Optional[bool] = None, ): ... def zero_grad(self, set_to_none: Optional[bool] = None) -> None: ... def step(self, closure=None, grad_scaler=None): ...[0, 1)异常处理:非法标量值沿用 ValueError;AMSGrad、capturable、关闭 bias correction 等沿用既有拒绝行为;不支持的 DTensor 类型组合和不匹配梯度布局明确报错。
3.5.2.2 状态接口
接口原型:
def initialize_state(self, param, store_param_remainders=None) -> None: ... def get_unscaled_state( self, param: torch.nn.Parameter, state_name: str, skip_unscale: bool = False ) -> torch.Tensor: ... def set_scaled_state(self, param, state_name, unscaled_state): ... def state_dict(self): ... def load_state_dict(self, state_dict): ...exp_avg/exp_avg_sq/master_param异常处理:保留旧版 uint8 checkpoint、非法 remainder 和配置不一致检查。读取尚未初始化的状态仍由现有状态查找行为处理,不隐式创建缺失状态。
调用参考代码:
import torch from transformer_engine.pytorch.optimizers import FusedAdam # param 和 grad 由 Core FSDP 提供,都是同布局的 DTensor。 optimizer = FusedAdam( [param], lr=1e-3, master_weights=True, master_weight_dtype=torch.float32, use_decoupled_grad=True, ) optimizer.initialize_state(param, store_param_remainders=False) param.decoupled_grad = grad optimizer.step() optimizer.zero_grad() # 状态交由 Core 的 fsdp_dtensor / DCP 流程保存。 optimizer_state = optimizer.state_dict()示例强调已有 Core 参数对象的调用方式,不在优化器中自行构造通信组或重分片。
4. 测试设计
4.1 测试代码结构
新增 fixture 为普通 Tensor、Shard、Replicate 提供参数构造方式;继续复用原数值断言与状态测试。未新增
test_*函数。其他普通优化器测试保留为回归集合。4.2 参数化覆盖设计
单进程运行的 Shard 并不能产生真正的跨 rank 非均匀分片。真实空分片和非均匀分片验收必须使用至少两个 rank;本地验证另以真实 DTensor 加模拟进程组分别执行 rank 0/1 的本地逻辑。
4.3 测试执行流水线
4.4 断言维度
布局拒绝分支、BF16
skip_unscale的独立覆盖和多维 mesh 仍属于后续验收关注点,不能把设计断言全部记为当前已自动化通过。4.5 容差体系
rtol=1e-6, atol=1e-7,沿用已有断言rtol=1e-2, atol=1e-2torch.testing.assert_close的 dtype 默认容差4.6 测试通过条件完整定义
当前验证状态:
unsupported gloo device,未形成真实多进程通信证据39 项定向结果与 80 项回归存在重叠,不应相加成为独立覆盖总数。两项 CPU 失败为 FP16/BF16 参数配 FP32 moments 的原有直接内核场景,已在
ccb9ae5重现;这只说明它们不是本次改动引入,不能证明 NPU 对应场景已通过。原 FP8/量化/remainder 路径未在此次 CPU 回归中得到验证。4.7 运行方式
在安装可用 TENPU、torch_npu 的仓库根目录复用已有文件:
# 普通环境下运行现有优化器集合 pytest tests/pytorch/test_fused_optimizer.py -v # 两个 NPU rank 验证已有 DTensor 参数化场景 torchrun --standalone --nproc_per_node=2 -m pytest \ tests/pytorch/test_fused_optimizer.py -v \ -k 'adamw_matches_manual_adamw_reference or low_precision_state_checkpoint_semantics or master_weights_checkpoint_semantics'fixture 在已有 process group 存在时复用;否则使用 HCCL 初始化,并依据
LOCAL_RANK选择 NPU。普通非分布式执行时使用单 rank 初始化;多卡执行使用torchrun环境。以上为目标环境运行方式,不是已在本机执行通过的命令记录。4.8 集成测试
fsdp_dtensor保存恢复5. 缺点和风险
5.1 潜在风险
5.2 实现成本
transformer_engine/pytorch/optimizers/fused_adam.pytests/pytorch/test_fused_optimizer.py+161/-455c85c0d相对ccb9ae55.3 兼容性问题
5.4 迁移方案
018/feat/fsdp-dtensor-fused-adam的5c85c0d,按既有安装方式部署 TENPU。354bb817,使用fsdp_dtensor的已有系统用例。迁移检查清单:
6. 现有技术
6.1 NVIDIA Transformer Engine FusedAdam
设计特点:参考快照在状态初始化时按 local Tensor 分配,并将状态包装为 DTensor;更新时向 fused CUDA 内核提供本地参数、梯度和状态。
借鉴之处:分布式状态保存、本地更新、checkpoint 元数据恢复以及
skip_unscale参数。差异之处:TENPU 复用 NPU 更新内核,保留原状态 setter 行为,功能集合小于 NVTE,且
initialize_state保留单参数调用兼容例外。NVTE 支持 FusedAdam 不代表 NVIDIA CUDA 扩展能够直接在 NPU 运行。6.2 Megatron-Core 0.18 Optimizer 与 FSDP
设计特点:优化器工厂优先尝试 TE FusedAdam,FSDP 提供分布式参数与梯度,precision-aware 配置决定主参数和独立梯度处理。
借鉴之处:不修改优化器选择流程,直接满足已有 TE 调用入口的数据约定。
差异之处:TENPU 不拥有分片布局生成、梯度通信和 DCP 调度,只消费 Core 的结果。
6.3 MindSpeed master Apex Optimizer Patch
设计特点:master 基础适配将
apex.optimizers.FusedAdam映射到 MindSpeed AdamW 实现,旧 FSDP 使用旧数据结构和入口。借鉴之处:沿用已有 NPU Adam/AdamW 的数值目标及普通训练兼容要求。
差异之处:0.18 实际选择 TE 时,仅 patch Apex 名称未必生效;在 MindSpeed 全局替换优化器还需处理提前导入的绑定和 precision-aware 参数。因此本次直接在 TENPU 完成 DTensor 支持。
7. 未解决问题
7.1
initialize_state严格签名收敛问题描述:当前
store_param_remainders=None与 NVTE 必传参数不同。影响范围:严格
inspect.signature比较不相等;删除默认值会破坏 Core 0.18 和已有 TENPU 的单参数调用。解决计划:先固定目标 NVTE 版本,审计并适配所有调用方显式传参,完成普通/precision-aware 回归后,再评估删除兼容默认值。不得仅修改 TENPU 签名而忽略上游调用。
7.2 NPU 多卡与真实检查点验证
问题描述:当前通过的是 CPU 参考与本地 DTensor 逻辑,未覆盖 NPU 内核、HCCL 和实际 DCP 读写。
影响范围:训练时序、混合 dtype 内核、跨进程状态恢复和设备内存行为。
解决计划:运行第 4 章已有测试和 Core 系统用例,记录完整环境及失败归因。
7.3 量化 DTensor 与 Parameter Remainder
问题描述:当前明确拒绝 DTensor 量化模型、FP8 moments 与 INT16 remainder。
影响范围:FP8 主权重及低精度训练状态压缩组合。
解决计划:只有出现明确训练需求后,再围绕 local QuantizedTensor、scale/amax、存储别名与 checkpoint 表示单独设计,不在当前最小适配中放宽 guard。
7.4 复合 Mesh、状态身份与更完整断言
问题描述:当前新增覆盖以一维 Shard/Replicate 为主;setter 仍可能替换状态 Tensor,布局校验也不是全局一致性证明。
影响范围:TP/DP 复合布局、外部状态引用和严格 TE 行为对齐。
解决计划:复用既有测试补充多维 mesh、布局拒绝、BF16 skip 与状态引用检查;若调用方依赖原地状态身份,再评估与 NVTE setter 进一步收敛。
附录
A. 参考资料链接
TENPU 运行时源码应按提交读取
transformer_engine/pytorch/optimizers/fused_adam.py,避免当前仓库检出的其他分支覆盖审查结论。B. 术语表
param.grad挂载的梯度