已关闭
[RFC]:TENPU FusedAdam 的 DTensor 分片更新与 FSDP 状态适配 #43
baichuan51141创建于  13 天前关闭于  13 天前
baichuan51141成员
13 天前 创建

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。

核心价值:

  • 解除普通浮点 DTensor 参数被 FusedAdam 一律拒绝的问题
  • 在同一优化器中复用普通 Tensor 的 Adam/AdamW 数值更新路径
  • 支持浮点 moments、master weights、独立梯度和空分片处理
  • 保留优化器状态的分布式元数据,供 Core 的 fsdp_dtensor 流程使用
  • 补齐 TE 的 get_unscaled_state(..., skip_unscale=False) 参数
  • 明确接口兼容例外、量化边界和真实 NPU 验收状态

1.2 动机

使用场景

  1. MindSpeed 0.18 启用 Megatron FSDP,Core 向优化器提供 DTensor 参数。
  2. 参数校验已经通过,但首次初始化状态或 optimizer.step() 报 DTensor 不支持。
  3. 使用 precision-aware optimizer,将高精度梯度挂载到 decoupled_grad。
  4. 使用 FP16/BF16 参数与 FP32 或 FP16 master weights。
  5. 保存并恢复 optimizer moments 和主参数,恢复后继续训练。
  6. 参数在 DP rank 间分片不均,部分 rank 持有空的本地 shard。

当前痛点

  1. 显式拒绝:旧 _check_unsupported_param_type() 对所有 DTensor 抛出 NotImplementedError。
  2. 内核类型:仅删除检查会让 DTensor 进入已有 fused AdamW 内核,不能保证分布式 dispatch 正确。
  3. 状态布局:直接保存 local Tensor 会丢失分布式位置;把 DTensor 和本地 scale 混算又可能触发混合类型错误。
  4. 梯度来源:普通 grad 与 decoupled_grad 必须在同一分片语义下处理。
  5. 空分片:低精度状态缩放中的 abs().max() 不适用于空 Tensor。
  6. 接口冲突:NVTE 的 initialize_state 要求第二个参数,而 Core 0.18 存在单参数调用。

必要性

  • 优化器状态与参数必须指向同一 rank 的逻辑分片,才能正确更新与恢复。
  • DTensor 支持不仅是类型放行,还涉及状态初始化、数值计算、写回和序列化。
  • 将修复放在 TENPU 可以让已有 TE 导入入口继续工作,避免在 MindSpeed 再维护一套 Adam。
  • 必须区分 API 参数对齐、功能支持范围和验证覆盖,不能以其中一项替代其他项。

用户价值

  • 原有 transformer_engine.pytorch.optimizers.FusedAdam 入口可直接消费 FSDP 参数。
  • 普通 Tensor 与 DTensor 复用既有内核和低精度状态处理逻辑。
  • 检查点保留分片元数据,减少恢复时状态与参数错位的风险。
  • 面向实际 FSDP 浮点训练需求完成最小改动,量化组合另行演进。

1.3 目标

目标范围

  1. 提供内部 _local_tensor() 和 _wrap_state(),集中处理本地视图与分布式包装。
  2. 支持 FP32、FP16、BF16 普通浮点参数的 DTensor 表示。
  3. 支持 FP32、FP16、BF16 moments,以及已有 FP32/FP16 master weight 路径。
  4. 校验 DTensor 参数与梯度的 mesh、placements 和全局 shape。
  5. 将 direct 与低精度更新路径中的内核输入转换为本地 Tensor。
  6. 处理空本地 shard 与 None 梯度,保持 group step 语义。
  7. 保留 state_dict()/load_state_dict() 的浮点状态语义与分布式元数据。
  8. 对齐 TE 已有调用参数,记录 initialize_state 的兼容性例外。
  9. 复用现有测试函数扩展分片、精度与恢复场景。

非目标(边界说明)

  • 不提供通用 DTensor optimizer dispatch,也不增加新的 all-gather、all-reduce 或 reduce-scatter。
  • 不在 TENPU 中重新实现 Core 的 FSDP、梯度同步和 DCP 调度。
  • 不支持本次 DTensor 路径中的量化模型参数、FP8/uint8 moments 或 INT16 parameter remainder。
  • 不改变普通 Tensor 的既有量化模型、FP8 moments 和 remainder 实现。
  • 不新增 capturable、AMSGrad 或 bias_correction=False 能力。
  • 不保证 NVTE 所有高级功能都已在 TENPU 实现。
  • 不声明现有接口签名已经完全一致;initialize_state 的可选参数是明确保留的差异。

2. 用例分析

2.1 功能点描述

功能点 描述 优先级
DTensor 类型接入 对支持范围内的 DTensor 放行 P0
本地数据视图 to_local() 供已有更新内核使用 P0
状态包装 使用参数的完整分布式元数据重建状态 P0
梯度布局校验 校验类型、mesh、placements 和 shape P0
Direct 更新 FP32 moments 等满足条件时复用直接 fused 路径 P0
低精度更新 本地解缩放、FP32 工作区、写回存储 dtype P0
Master weights 在既有条件下创建和更新主参数 P0
独立梯度 使用 param.decoupled_grad P0
空分片 避免空 reduce-max 和空内核更新 P0
保存恢复 对外状态保持 DTensor 布局 P0
TE 参数兼容 补齐 skip_unscale 并记录签名例外 P0
普通路径回归 复用 Tensor、SGD、Adam 既有检查 P0

2.2 关键性能指标

以下为设计目标及验收口径;当前没有 NPU 性能数据。

指标 目标值 说明
优化器新增 collective 0 分片同步由 Core 完成
Direct 路径参数全量复制 0 次新增全局复制 本地视图进入原内核
DTensor 元数据保留 mesh/placements/shape/stride 一致 不用 local shape 替代 global shape
数值更新 满足既有浮点容差 使用已有手工 AdamW 参考
恢复连续性 下一步参数满足已有容差 保存后恢复再更新
空 shard 错误 0 次 不对空状态执行 max,不提交空更新
普通 Tensor 回归 无本改动新增失败 与相同环境基线比较
NPU 时间与显存 记录相对原路径的变化 低精度状态仍需本地工作 Tensor

2.3 安全隐私要求

  • 优化器处理模型参数及训练状态,不额外输出参数、梯度或 moments 数值。
  • 本实现不新增网络连接和跨 rank 通信。
  • 保存恢复沿用调用方的状态字典和存储权限,不新增独立检查点目录约定。
  • 错误信息描述类型、精度和布局约束,不包含训练样本。

2.4 DFX 要求

兼容性

  • TENPU 基线固定为 ccb9ae5,本次实现固定为 5c85c0d。
  • NVTE 比较快照为本地 42b84005,不将该快照视为所有 TE 版本的统一签名。
  • Core 0.18 的 precision-aware init_state_fn 存在 opt.initialize_state(p) 调用。
  • 普通 Tensor 路径保留原内核、状态 dtype 规则和已有不支持项。

可维护性

  • DTensor 转换封装为两个私有静态 helper。
  • 优化器仍以原始参数对象为 self.state 和 scale 字典的 key。
  • 不临时替换 param_groups,不复制整个优化器实现。
  • 公开方法中的兼容差异在本 RFC 与源码注释中同时说明。

可测试性

  • 只扩展既有数值与检查点测试函数。
  • fixture 在普通 Tensor、Shard 和 Replicate 间切换。
  • 多 rank 场景可用既有文件在 torchrun 下执行。
  • 本地模拟进程组验证本地 shard 计算,不用来证明 collective 或 DCP I/O 正确。

可靠性

  • 在不支持的量化/remainder DTensor 组合上明确报错。
  • Partial 参数布局需要先完成归约,优化器不隐式补通信。
  • 梯度布局不匹配时拒绝更新,不自动猜测或重分片。
  • 空 shard 的处理不改变已有 group step 推进规则。

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 --> N

3.2 技术选型

方案对比

方案 描述 优点 缺点 是否采用
A:在 TENPU 内本地化计算 状态仍为 DTensor,内核只见 local Tensor 与 NVTE 结构接近,复用现有路径 需覆盖状态和检查点所有边界 是
B:仅删除类型拒绝 将 DTensor 直接传原内核 代码行数少 无法保证 dispatch、scale 混算及状态布局 否
C:MindSpeed 全局 patch Adam 替换 TE/Apex/optimizer 导入绑定 可以绕开当前 TENPU 错误 影响其他优化器模式,增加跨层逻辑 否
D:收集完整参数后更新 优化器内部 all-gather 容易按完整参数理解计算 破坏分片内存目标,增加通信 否
E:临时替换参数组和状态 key 更新时改成 local Parameter 复用普通优化器外壳 参数身份、调度器和状态 hook 容易错位 否

选型理由

  1. Adam 的逐元素更新适合在已经完成归约的本地 shard 上执行。
  2. Core 已通过 TE 导入路径选择优化器,TENPU 修复可直接被调用方使用。
  3. NVTE 已采用本地计算与 DTensor 状态包装的方式,可对齐结构而不复制 CUDA 扩展。
  4. 仅增加必要转换和校验,不改变普通 Tensor 的内核与参数分组策略。

3.3 功能与性能设计

3.3.1 类型支持范围与前置条件

对象/配置 本次 DTensor 路径 普通 Tensor 路径
FP32/FP16/BF16 参数 支持本地 shard 更新 沿用已有实现
FP32/FP16/BF16 moments 支持,保留既有缩放规则 沿用已有实现
FP32/FP16 master weight 支持既有适用条件 沿用已有实现
量化模型参数 明确不支持 不修改原量化支持范围
uint8/FP8 moments 明确不支持 原有能力及依赖检查保留
INT16 parameter remainder 明确不支持 原有 BF16 remainder 路径保留
Shard/Replicate 布局 本地更新模型适用 不适用
Partial 参数布局 拒绝 不适用
capturable=True 既有不支持项 既有不支持项
amsgrad=True 既有不支持项 既有不支持项
bias_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_groups 原参数对象和超参数 不替换为临时 local Parameter
self.state[param] exp_avg、exp_avg_sq、可选 master_param 继续用原参数作为 key
DTensor state local storage 与全局布局 由 _wrap_state() 包装
_scales[param][name] 状态缩放因子 本地标量 Tensor,不额外分布式归约
_step_tensors group 关联的内核步数 Tensor 延用缓存和推进机制
group["step"] 组级步数 保存加载时转为整数
direct groups 本地参数、梯度、两类 moments 列表 直接交给已有内核
low precision groups 原参数身份及本地 FP32 工作区 计算后按原参数写回状态
FP8/remainder 缓存 原高级路径元数据 本次不扩展到 DTensor

一个参数的典型状态关系为:

原参数 p: DTensor(local=p_r, mesh=M, placements=L, global_shape=S)
self.state[p]["exp_avg"]:    DTensor(local=m_r, 同 M/L/S/stride)
self.state[p]["exp_avg_sq"]: DTensor(local=v_r, 同 M/L/S/stride)
self.state[p]["master_param"]: 可选 DTensor(local=w32_r 或缩放主参数)
self._scales[p][state_name]: 本地 scale

状态 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 执行不必要的分布式操作。

存储 dtype 写入规则 读取规则
FP32 转为 FP32 后保存 返回本地 FP32 状态副本
BF16 转为 BF16,scale 为 1 默认转 FP32;skip_unscale=True 返回本地 BF16
FP16 使用本地最大绝对值计算 scale 后保存 转 FP32 并乘本地 scale
uint8/FP8 不属于 DTensor 支持范围 普通 Tensor 原路径保留

FP16 状态沿用既有缩放公式,目标动态范围为 finfo(float16).max / 2:

a = max(abs(local_unscaled_state))
s = max(a / storage_amax, finfo(float32).tiny)
stored = cast_fp16(local_unscaled_state / s)
restored = cast_fp32(stored) * s

当 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,普通浮点参数通过 local copy_() 写回。

此路径复用已有 _apply_fused_adamw_groups() 与 torch._fused_adamw_()。它不增加 NVIDIA CUDA 扩展,也不添加第二套 Adam 实现;在 NPU 环境由已有 torch_npu 运行时提供对应执行能力。

3.3.7 Adam/AdamW 数值语义

设当前 rank 已归约的本地梯度为 g,参数或 master 工作区为 w,组级步数为 t:

m_t = beta1 * m_(t-1) + (1 - beta1) * g
v_t = beta2 * v_(t-1) + (1 - beta2) * g * g
m_hat = m_t / (1 - beta1^t)
v_hat = v_t / (1 - beta2^t)
AdamW: w_t = (1 - lr * weight_decay) * w_(t-1)
             - lr * m_hat / (sqrt(v_hat) + eps)

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 --> I

Core 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 接口对齐与明确例外

方法 NVTE 参考快照 TENPU 当前实现 结论
__init__ 既有 TE 参数顺序、默认值、keyword-only 分界 保持对应结构 本次不新增构造参数
zero_grad set_to_none=None 相同调用参数 沿用 TENPU 已有细节
step closure=None, grad_scaler=None 相同调用参数 能力边界仍受 TENPU 限制
get_unscaled_state param, state_name, skip_unscale=False 本次补齐第三参数 对齐该调用签名
set_scaled_state param, state_name, unscaled_state 同样的无注解参数形式 本次对齐签名形式
state_dict / load_state_dict 无参数 / state_dict 相同调用参数 分布式状态语义接入
initialize_state param, 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 路径隔离与性能影响

位置 DTensor 行为 普通 Tensor 行为
_local_tensor 返回 local Tensor 返回输入对象
_wrap_state 使用原参数元数据包装 返回输入状态
参数类型检查 放行浮点布局,拒绝限定组合 原量化参数检查保留
更新内核 本地 Tensor 列表 原 Tensor 列表
参数写回 local copy_() 原参数 copy_()
低精度状态 本地解缩放与包装 原缩放路径

新增开销主要是 Python 类型判断和必要的 DTensor 包装;低精度与检查点路径仍存在本地 Tensor 分配/复制。不能据此承诺零开销或 state Tensor 身份不变。Direct 路径没有增加全局参数收集,整体分片通信仍由 Core 调度。

3.4 安全隐私与 DFX 设计

3.4.1 兼容性设计

兼容性维度 设计方案
TE 导入入口 保留原 transformer_engine.pytorch.optimizers.FusedAdam
参数身份 原参数继续作为状态及 scale key
数值内核 复用原 fused AdamW 与低精度工作区
检查点 使用参数元数据重新包装状态
普通高级功能 不扩展也不移除原 FP8/remainder 路径
公开接口 明确记录 initialize_state 兼容例外

3.4.2 可维护性设计

  • helper 只负责本地访问和状态包装,不掺入优化器选择策略。
  • 显式区分 NVTE 结构借鉴与 NPU 能力实现。
  • 文档引用固定提交,避免当前工作区其他分支掩盖实际修复。
  • 后续扩展量化 DTensor 时单独审查 quantizer、scale、存储身份和检查点约束。

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 可靠性设计

  • 在梯度布局不匹配时明确失败,不静默转换布局。
  • local Tensor 只用于算术,保存时恢复全局元数据。
  • 原始 self.state[param] key 不变,防止 optimizer scheduler 和调用方状态查找错位。
  • 保留错误处理边界:本优化器不修复上游未同步的梯度,也不接管非有限值检测。

3.5 编程与调用设计

3.5.1 编程模型基本设计

开发环境
  • 目标环境:Linux Ascend NPU,安装版本匹配的 PyTorch、torch_npu、CANN 与 TENPU。
  • 框架调用方:MindSpeed 0.18 Megatron FSDP 分支。
  • 测试依赖:pytest、PyTorch distributed/DTensor,真实多卡使用 HCCL。
  • 本地检查环境:PyTorch 2.9.0+cpu;以源码提取方式隔离 NPU 导入依赖。
开发约束
  • 参数及梯度须已由 Core 完成正确分片与归约。
  • 只对支持的浮点状态使用 DTensor,本次不启用 DTensor FP8/remainder。
  • master_weight_dtype 受 TENPU 既有 FP32/FP16 限制。
  • 不将 CPU 模拟环境当作 TENPU 完整运行环境。
  • 不新增测试文件或测试函数。
可验收设计
  • 接口验收:检查 TE 参数形式与唯一已记录的初始化签名例外。
  • 数值验收:与已有手工 AdamW 参考比较本地更新。
  • 状态验收:moments/master 的 dtype、shape 和分布式元数据正确。
  • 恢复验收:恢复目标参数及 optimizer 后执行相同步骤并比较。
  • 回归验收:普通 Tensor 和原有高级路径在对应硬件环境中运行已有用例。

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): ...
参数名称 输入/输出 类型 描述 取值范围
params 输入 参数或参数组迭代器 普通参数或支持范围内 DTensor FP32/FP16/BF16
lr 输入 float 学习率 非负
betas 输入 tuple 一阶和二阶衰减因子 每项 [0, 1)
eps 输入 float 分母稳定项 非负
weight_decay 输入 float AdamW 衰减或 L2 系数 非负
adam_w_mode 输入 bool 选择解耦衰减或 L2 True/False
master_weights 输入 bool 是否使用既有 master 路径 默认 False
exp_avg_dtype/exp_avg_sq_dtype 输入 torch.dtype moments 存储 dtype DTensor 限浮点支持集合
use_decoupled_grad 输入 bool 是否读取独立梯度 默认 False
closure 输入 callable 或 None 计算损失闭包 沿用已有行为
step 返回值 输出 loss 或 None 闭包计算结果 沿用已有行为

异常处理:非法标量值沿用 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): ...
参数名称 输入/输出 类型 描述 取值范围
param 输入 原参数对象 状态字典的 key 必须使用优化器持有的参数
store_param_remainders 输入 bool 或 None 显式校验 remainder 配置 DTensor 不支持 remainder
state_name 输入 str 目标状态 exp_avg/exp_avg_sq/master_param
skip_unscale 输入 bool BF16 状态保留原 dtype 默认 False
unscaled_state 输入 Tensor 或 DTensor 待保存的未缩放值 本地化后进入原精度转换
get 返回值 输出 local Tensor 未缩放状态 通常 FP32,BF16 skip 例外
state_dict 返回值 输出 dict 分布式优化器状态 包含 param_groups 与 state

异常处理:保留旧版 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 测试代码结构

tests/pytorch/
└── test_fused_optimizer.py
    ├── optimizer_param                              # 新增 fixture,非测试函数
    ├── _local / _param_grad_like                     # 辅助函数适配
    ├── test_adamw_matches_manual_adamw_reference      # 扩展既有函数
    ├── test_adam_low_precision_state_checkpoint_semantics
    └── test_adam_master_weights_checkpoint_semantics

新增 fixture 为普通 Tensor、Shard、Replicate 提供参数构造方式;继续复用原数值断言与状态测试。未新增 test_* 函数。其他普通优化器测试保留为回归集合。

4.2 参数化覆盖设计

既有测试 新增/沿用维度 取值
手工 AdamW 参考 参数表示 Tensor、Shard、Replicate
手工 AdamW 参考 全局元素数 1、7;多 rank 下覆盖空/非均匀 shard
手工 AdamW 参考 梯度来源 普通 grad、decoupled_grad
手工 AdamW 参考 更新次数 5 次
低精度状态 checkpoint moments dtype FP16/FP16、BF16/BF16、FP16/FP32、FP32/BF16
低精度状态 checkpoint 参数 dtype FP16
Master checkpoint 参数 dtype FP16、BF16
Master checkpoint master dtype FP32、FP16
两类 checkpoint 参数表示与恢复 Tensor/Shard/Replicate,恢复后下一步比较
原有回归 普通 optimizer 行为 SGD、Adam/L2、zero_grad、step 和选项检查

单进程运行的 Shard 并不能产生真正的跨 rank 非均匀分片。真实空分片和非均匀分片验收必须使用至少两个 rank;本地验证另以真实 DTensor 加模拟进程组分别执行 rank 0/1 的本地逻辑。

4.3 测试执行流水线

  1. 核对构造器与状态接口,明确单参数初始化兼容例外。
  2. 运行普通 Tensor 已有数值与状态用例。
  3. 用 fixture 构造 Shard/Replicate 参数和同布局梯度。
  4. 执行更新,与手工参考比较 local 参数。
  5. 导出 state dict,检查 moments 精度和分布式元数据。
  6. 将模型参数同步到目标对象,加载 optimizer state 后执行下一步比较。
  7. 在真实 NPU 多进程环境运行同一文件,再接入 Core 系统脚本。
  8. 最后验收真实 DCP I/O 与其他高级路径回归。

4.4 断言维度

断言对象 验证内容
参数数值 多次更新后 local 参数与已有参考一致
梯度来源 普通与独立梯度都能更新,zero_grad 行为保留
状态 dtype 运行时存储 dtype 与配置一致
未缩放值 默认读取为 FP32,BF16 skip 分支按约定
Checkpoint dtype 浮点 moments/master 对外为 FP32
DTensor 布局 既有低精度 checkpoint 用例检查 mesh、placements、shape、stride
恢复连续性 目标参数与源参数执行同一步后比较
空 shard 不触发 max 或内核空输入错误
普通路径 与基线对比,定位环境失败与新增回归

布局拒绝分支、BF16 skip_unscale 的独立覆盖和多维 mesh 仍属于后续验收关注点,不能把设计断言全部记为当前已自动化通过。

4.5 容差体系

场景 容差
手工 FP32 AdamW 参考 rtol=1e-6, atol=1e-7,沿用已有断言
参数与 master 一致性 既有 rtol=1e-2, atol=1e-2
恢复后下一步 既有 torch.testing.assert_close 的 dtype 默认容差
状态 dtype、布局、参数选项 精确相等
NPU 普通模型差分 使用原文件对应场景容差
性能 单独记录,不使用数值容差替代

4.6 测试通过条件完整定义

  1. 支持范围内 DTensor 不再触发旧的统一拒绝错误。
  2. 本地更新数值与已有参考一致,参数组计数不因空 shard 分化。
  3. 状态 dtype、scale 和 DTensor 元数据满足配置要求。
  4. 保存加载后能够执行下一步,并与连续训练对照一致。
  5. 普通 Tensor 路径没有本改动新增回归。
  6. 不支持组合明确报错,不能静默降级为全量参数更新。
  7. 真实 NPU 测试与 Core 联合保存恢复通过后,才能认定端到端验收完成。

当前验证状态:

验证层次 结果 证据边界
Python 语法与 diff 检查 通过 静态检查
CPU 定向检查 模拟 rank 1 的 39 项通过 包含普通/Shard/Replicate 的相关参数组合
CPU 扩展回归 80 项通过,2 项失败,27 项未选入 源码提取后运行已有用例,排除量化/remainder 等依赖项
两项 CPU 失败基线对照 未修改 main 同样失败 CPU fused AdamW 不接受对应混合 dtype 组合
Windows Gloo 初始化 失败 unsupported gloo device,未形成真实多进程通信证据
DTensor 本地验证 使用真实 DTensor 与模拟进程组 分别验证两个 rank 的本地逻辑,无真实 collective
NPU/HCCL、多卡 DCP I/O 未执行 仍需硬件集成验收

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 集成测试

场景 描述 验证点
Core 标准 FSDP 上游管理主参数并安装 DTensor 梯度 原导入入口完成 step
Precision-aware TENPU 管理 master,使用独立梯度 dtype、master 与参数回写
非均匀分片 全局元素数不可整除 DP 大小 local shape、global metadata
空分片 rank 数大于某参数可分片元素数 None/空梯度和 group step
fsdp_dtensor 保存恢复 Core 调用真实 DCP 模型、moments、master、step 连续性
普通量化/remainder 使用原有对应硬件用例 证明本次改动未影响旧功能

5. 缺点和风险

5.1 潜在风险

风险类型 描述 影响 应对措施
元数据风险 错用 local shape 导致分片定位错误 高 使用原参数全局 shape/stride
状态引用风险 setter 重建状态,旧引用失效 中 调用方按参数 key 获取最新状态
布局风险 复合 mesh 或 local storage 不匹配 高 Core 集成验收,必要时扩展既有断言
签名风险 单参数调用与 NVTE 严格签名冲突 高 明确保留例外,协调调用方后再收敛
精度风险 本地低精度 scale 或 master 写回错误 高 多 dtype、恢复连续性检查
环境风险 CPU 模拟掩盖 NPU 内核差异 高 多卡 NPU 单独验收
量化回归风险 原高级路径未在本地实际运行 高 使用已有 NPU 量化测试集合
检查点风险 内存加载通过但真实 DCP 失败 高 Core 端到端保存恢复

5.2 实现成本

成本类型 估算 说明
业务源码 1 个文件 transformer_engine/pytorch/optimizers/fused_adam.py
测试适配 1 个文件 tests/pytorch/test_fused_optimizer.py
提交差异 +161/-45 5c85c0d 相对 ccb9ae5
新增优化器类 0 复用 FusedAdam
新增测试文件/函数 0 新增 fixture/helper 并扩展已有测试
新增通信 0 不改变 Core 分片同步
后续工作 硬件验收与签名收敛 不计为已完成实现

5.3 兼容性问题

兼容性维度 问题 解决方案
TE 签名 初始化第二参数在两端必选性不同 当前保留可选参数并记录差异
TE 功能 相同参数不代表 NPU 支持全部 NVTE 功能 保留明确不支持检查
普通 Tensor helper 可能触及共享状态路径 普通输入返回原对象,执行基线回归
Checkpoint 低精度运行态和导出精度不同 延用 FP32 未缩放导出语义
DTensor 量化 local Tensor 可能是 QuantizedTensor 当前明确拒绝,不进入原量化内核
老 TENPU 环境 实际导入仍是旧安装包 核对模块文件路径与部署 commit

5.4 迁移方案

  1. 获取 018/feat/fsdp-dtensor-fused-adam 的 5c85c0d,按既有安装方式部署 TENPU。
  2. 确认 Python 实际导入该版本的 FusedAdam,而不是其他环境中的旧包。
  3. 配套 MindSpeed 354bb817,使用 fsdp_dtensor 的已有系统用例。
  4. 先验证默认浮点状态,再扩展低精度 moments、master 和独立梯度。
  5. 在原测试文件中执行普通 Tensor 与多卡 DTensor 回归。
  6. 完成 Core 真实保存恢复后再扩大生产训练范围。

迁移检查清单:


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. 术语表

术语 全称 说明
TENPU Transformer Engine NPU TE 接口的 NPU 实现
NVTE NVIDIA Transformer Engine 本提案的结构与 API 参考
DTensor Distributed Tensor 带 mesh 与分片位置的 Tensor
DeviceMesh - 定义设备/rank 网格
Shard - 沿某维分片的 placement
Replicate - 在 mesh 维上复制的 placement
Partial - 仍需归约的分布式表示
Moments - Adam 的一阶、二阶累积状态
Master weight - 用于更高精度更新的主参数
Decoupled gradient - 独立于 param.grad 挂载的梯度
Parameter remainder - 用 INT16 保存 BF16 主参数补充信息的既有机制
likedislike
ascend-robotascend-robot成员
13 天前 添加了label:rfc
Bbaichuan51141成员
13 天前 关联了pull request:feat: support DTensor shards in FusedAdam for MindSpeed FSDP
ascend-robotascend-robot成员
13 天前 关闭了 issue
ascend-robotascend-robot成员
13 天前 issue状态由 TODO 改变为 DONE
ascend-robotascend-robot成员
13 天前 添加了label:resolved