已关闭
[RFC]: 为 FSDP 添加 MXFP8-32x32 量化支持 #175
kyle_zhang创建于  6月14日关闭于  7月11日
kyle_zhang
kyle_zhang
6月14日 创建

状态(Status): Draft
作者(Authors): @kyle_zhangchi
创建日期(Created): 2026-06-14
更新日期(Updated): 2026-06-14
相关 Issue/PR: https://gitcode.com/Ascend/MindSpeed/pull/3534

1. 概述

1.1 简介

本提案为 FSDP (Fully Sharded Data Parallel) 框架添加 MXFP8-32x32 量化支持,通过引入新的 block size 配置 [1, 32, 32],在保持模型精度的同时,进一步优化 all-gather 通信开销,提升大规模分布式训练性能。该方案扩展了现有的 MXFP8 量化能力,为用户提供更多量化粒度选择,以适应不同模型结构和硬件特性的优化需求。

1.2 动机

当前痛点:

  • 现有 MXFP8 量化仅支持 [1, 1, 32] block size 配置,在某些模型结构下无法充分挖掘量化潜力
  • 在大规模分布式训练场景下,all-gather 通信开销仍然较大,影响训练效率
  • 用户缺乏针对不同模型特性的量化粒度选择,无法灵活权衡精度与性能

使用场景:

  • 大规模分布式训练场景,需要降低通信开销
  • 对量化精度有较高要求的模型训练
  • 需要根据模型结构特性选择最优量化配置的场景

用户价值:

  • 提供更灵活的量化配置选择,适应不同模型特性
  • 通过优化 all-gather 操作,降低通信开销,提升训练性能
  • 保持模型精度的同时,获得更好的训练效率

不做此提案的影响:

  • 用户无法使用更优化的量化配置,训练性能受限
  • 在特定模型结构下,无法充分挖掘量化优化潜力

1.3 目标

目标:

  • 为 FSDP 添加 MXFP8-32x32 量化支持,支持 block size [1, 32, 32] 配置
  • 优化 all-gather 通信逻辑,减少不必要的张量传输
  • 提供新的量化配方 (recipe) mxfp8-32x32,方便用户使用
  • 保持与现有 MXFP8 实现的兼容性

非目标 (边界说明):

  • 不涉及其他量化格式 (如 FP4, INT8) 的支持
  • 不改变现有的 FSDP 核心架构和分片策略
  • 不涉及非 NPU 硬件平台的适配
  • 不改变量化的动态策略 (仍采用 DYNAMIC scaling strategy)

2. 用例分析

2.1 大规模分布式训练场景

功能点:

  • 支持 MXFP8-32x32 量化配置初始化
  • 在 forward 和 backward 过程中正确处理量化权重
  • 优化 all-gather 通信,减少 scale 张量的传输次数

关键性能指标:

  • All-gather 通信时间降低 20%-30% (相比原 MXFP8 方案)
  • 训练吞吐量提升 10%-15% (在大规模场景下)
  • 内存占用保持不变或略有降低

安全隐私及 DFX 要求:

  • 兼容性: 需与现有 MXFP8 配置共存,不影响已有模型
  • 可维护性: 代码结构清晰,量化逻辑模块化,便于后续扩展
  • 可测试性: 提供完整的单元测试和集成测试用例
  • 可靠性: 需验证量化精度损失在可接受范围内 (与原方案对比)

使用限制:

  • 仅支持 NPU 硬件平台
  • 需要硬件支持 npu_dynamic_block_mx_quant 算子
  • Block size 仅支持 [1, 1, 32][1, 32, 32] 两种配置

2.2 MoE (Mixture of Experts) 模型训练场景

功能点:

  • 支持 MoE 模型的 gate_up_proj 和 down_proj 权重量化
  • 正确处理 grouped matmul 的量化计算
  • 支持 expert parallel 场景下的量化

关键性能指标:

  • MoE 模型训练性能提升 15%-20%
  • Expert 间通信开销降低

使用限制:

  • 需配合 MoE 相关算子使用
  • 需正确设置 intermediate_size 和 hidden_dim 参数

3. 方案设计

3.1 总体方案

整体设计思路:

本方案通过扩展现有的 MXFP8 量化框架,引入新的 block size 配置 [1, 32, 32],实现 MXFP8-32x32 量化。核心设计包括:

  1. 配置层扩展:ScalingGranularity 类中支持新的 block size 验证,添加 mxfp8-32x32 配方注册
  2. 量化逻辑适配:PreQuantWeightMXLinearMXFP8GMM 等核心类中,根据配置选择不同的量化路径
  3. All-gather 优化: 针对 MXFP8-32x32 特性,优化 fsdp_pre_all_gatherfsdp_post_all_gather 逻辑,减少 scale 张量的传输

技术方案:

  • 量化算法: 使用 npu_dynamic_block_mx_quant 算子,支持 block-level 的动态量化
  • 数据布局: Block size [1, 32, 32] 表示在 input_dim 和 output_dim 维度上按 32 分块,token_dim 不分块
  • 通信优化: MXFP8-32x32 模式下,forward 和 backward 共享同一权重,仅需传输一次 weight 和两次 scale,减少通信量

架构图:

┌─────────────────────────────────────────────────────────┐
│                    Quantization Config                   │
│  ┌──────────────┐        ┌────────────────┐            │
│  │   mxfp8      │        │  mxfp8-32x32   │            │
│  │ [1, 1, 32]   │        │  [1, 32, 32]   │            │
│  └──────────────┘        └────────────────┘            │
└─────────────────────────────────────────────────────────┘
                          │
                          ▼
┌─────────────────────────────────────────────────────────┐
│              PreQuantWeight (Tensor Subclass)            │
│  ┌──────────────────────────────────────────────┐      │
│  │  fsdp_pre_all_gather:                        │      │
│  │  - mxfp8: (weight_fwd, scale_fwd,            │      │
│  │            weight_bwd, scale_bwd)            │      │
│  │  - mxfp8-32x32: (weight_fwd, scale_fwd,      │      │
│  │                  scale_bwd)                  │      │
│  └──────────────────────────────────────────────┘      │
└─────────────────────────────────────────────────────────┘
                          │
                          ▼
┌─────────────────────────────────────────────────────────┐
│                   Quantized Modules                      │
│  ┌─────────────┐         ┌──────────────┐              │
│  │  MXLinear   │         │  MXFP8GMM    │              │
│  └─────────────┘         └──────────────┘              │
└─────────────────────────────────────────────────────────┘

限制和约束:

  • 硬件平台: 仅支持华为 NPU
  • 软件依赖: 需要 torch_npu 支持 npu_dynamic_block_mx_quant 算子
  • 编程框架: 基于 PyTorch FSDP2 实现
  • 性能前提: 在多卡分布式场景下才能体现通信优化效果

3.2 技术选型

考虑过的方案:

  1. 方案一: 扩展现有 MXFP8,支持多种 block size

    • 优势: 复用现有代码,改动小,维护成本低
    • 劣势: 需要在多处添加条件判断,代码耦合度较高
    • 选择理由: 最终采用此方案,通过配置参数区分不同量化路径,保持代码统一性
  2. 方案二: 独立实现 MXFP8-32x32 模块

    • 优势: 代码独立,便于维护和测试
    • 劣势: 重复代码多,维护成本高,难以共享优化逻辑
    • 放弃理由: 与现有架构不一致,增加代码冗余
  3. 方案三: 使用静态量化替代动态量化

    • 优势: 可进一步降低计算开销
    • 劣势: 需要额外的校准步骤,使用复杂度高,精度损失可能更大
    • 放弃理由: 不符合当前动态量化的设计理念,用户体验差

3.3 功能与性能设计

3.3.1 配置管理功能

实现方案:

  • 扩展 ScalingGranularity 数据类,在 __post_init__ 中验证 block size 合法性
  • 新增 register_mxfp8_32x32_recipe() 函数,注册 mxfp8-32x32 配方
  • 配方参数: scaling_strategy=DYNAMIC, scaling_granularity=MX-[1, 32, 32], dtype 统一为 E4M3

核心流程:

用户指定 recipe_name="mxfp8-32x32"
        ↓
QuantRecipe 解析配置
        ↓
ScalingGranularity 验证 block_size=[1, 32, 32]
        ↓
创建 QuantizeConfig 实例
        ↓
模型转换器应用配置

数据模型变更:

  • ScalingGranularity.block_size 支持 [1, 32, 32] 取值
  • 新增 recipe_name="mxfp8-32x32" 配方
3.3.2 All-gather 优化功能

实现方案:

PreQuantWeight.fsdp_pre_all_gather 中:

  • MXFP8 模式: 返回 (weight_fwd, scale_fwd, weight_bwd, scale_bwd),forward 和 backward 使用不同权重
  • MXFP8-32x32 模式: 返回 (weight_fwd, scale_fwd, scale_bwd),forward 和 backward 共享同一权重,减少传输量

PreQuantWeight.fsdp_post_all_gather 中:

  • 根据 all_gather_outputs 长度判断量化模式,正确恢复权重和 scale

性能影响:

  • 通信量减少: MXFP8-32x32 模式下,减少一次 weight 张量的 all-gather 操作
  • 内存占用: 保持不变,权重仍需完整存储
  • 计算开销: 略异不大,主要优化在通信层面

影响范围:

  • mindspeed/fsdp/quantization/core/pre_quant_weight.py
  • mindspeed/fsdp/quantization/module/linear_mxfp8.py
  • mindspeed/fsdp/quantization/module/gmm_mxfp8.py
3.3.3 量化计算功能

实现方案:

引入 process_quant_result_fn 函数,根据 recipe_name 选择量化路径:

  • mxfp8: 调用 npu_dynamic_mx_quant_with_dual_axis,生成独立的 forward 和 backward 权重
  • mxfp8-32x32: 调用 npu_dynamic_block_mx_quant,forward 和 backward 共享权重,仅 scale 不同

核心流程:

输入权重 weight
        ↓
判断 recipe_name
        ↓
┌───────────────────┬────────────────────┐
│     mxfp8         │   mxfp8-32x32      │
│  dual_axis_quant  │  block_mx_quant    │
│  (weight_fwd,     │  (weight_fwd,      │
│   weight_bwd,     │   scale_fwd,       │
│   scale_fwd,      │   scale_bwd)       │
│   scale_bwd)      │                    │
└───────────────────┴────────────────────┘
        ↓
返回量化结果

3.4 安全隐私与 DFX 设计

安全隐私:

  • 不涉及用户数据隐私,仅处理模型权重
  • 量化过程在本地完成,不引入额外的数据传输风险

兼容性:

  • 向后兼容: 完全兼容现有 MXFP8 配置,不影响已有模型
  • 配置兼容: 新增配方不影响已有配方的使用
  • API 兼容: 不改变现有 API 签名,仅扩展配置参数

可维护性:

  • 代码结构: 通过 process_quant_result_fn 统一量化逻辑,减少重复代码
  • 配置管理: 使用 dataclass 和注册机制,便于扩展新的量化配方
  • 文档完善: 更新 docstring,说明支持的量化模式和配置

可测试性:

  • 单元测试: 需覆盖配置验证、量化计算、all-gather 逻辑
  • 集成测试: 需验证端到端训练流程的正确性
  • 性能测试: 需对比 MXFP8 和 MXFP8-32x32 的性能差异

可靠性:

  • 精度验证: 需对比量化前后的模型精度,确保损失在可接受范围
  • 异常处理: 对不支持的配置抛出明确的 ValueError
  • 边界检查: 验证 block size 的合法性,防止非法配置

3.5 编程与调用设计

3.5.1 编程模型基本设计

开发环境设计:

  • 硬件环境: 华为 NPU (Ascend 系列)
  • 软件环境: PyTorch 2.x, torch_npu, mindspeed
  • 开发工具: 标准的 Python 开发工具链,支持 NPU 调试
  • 编程框架: 基于 PyTorch FSDP2 的分布式训练框架

开发约束:

  • 硬件限制: 仅支持 NPU,不支持 CPU/GPU
  • 编程语言: Python 3.10+
  • 依赖版本: 需要特定版本的 torch_npu 支持 npu_dynamic_block_mx_quant 算子

可验收设计:

  • 功能验收: 提供示例脚本,验证量化配置的正确加载和训练流程
  • 性能验收: 提供性能对比脚本,验证通信开销的降低
  • 精度验收: 提供精度对比脚本,验证模型精度损失在可接受范围
3.5.2 接口定义与设计
3.5.2.1 ScalingGranularity 数据类

接口描述: 定义量化粒度配置,支持 MX 格式的 block size 设置

接口原型:

@dataclass
class ScalingGranularity:
    stype: ScalingGranularityEnum
    block_size: Optional[List[int]] = None

输入/输出参数:

参数名称 输入/输出 类型 描述 取值范围
stype 输入 ScalingGranularityEnum 缩放粒度类型 MX, PER_TENSOR, PER_CHANNEL
block_size 输入 Optional[List[int]] 缩放块大小 [token_dim, input_dim, output_dim] None, [1, 1, 32], [1, 32, 32]

异常处理:

  • stype=MXblock_size 不在支持列表中时,抛出 ValueError

约束说明:

  • block_size 仅在 stype=MX 时需要指定
  • [1, 32, 32] 为新增配置,需硬件支持

变更说明: 新增 [1, 32, 32] block size 支持

调用参考代码:

from mindspeed.fsdp.quantization.config import ScalingGranularity, ScalingGranularityEnum

# 创建 MXFP8-32x32 配置
granularity = ScalingGranularity(
    stype=ScalingGranularityEnum.MX,
    block_size=[1, 32, 32]
)
3.5.2.2 register_mxfp8_32x32_recipe 函数

接口描述: 注册 MXFP8-32x32 量化配方到全局配方字典

接口原型:

@recipe_register("mxfp8-32x32")
def register_mxfp8_32x32_recipe() -> QuantRecipe:
    ...

输入/输出参数: 无参数

返回参数:

参数名称 类型 描述 取值范围
QuantRecipe 实例 QuantRecipe MXFP8-32x32 量化配方 N/A

异常处理:

约束说明: 需在模型转换前调用,确保配方已注册

变更说明: 新增函数

调用参考代码:

from mindspeed.fsdp.quantization.config import get_recipe

# 获取 MXFP8-32x32 配方
recipe = get_recipe("mxfp8-32x32")
3.5.2.3 process_quant_result_fn 函数

接口描述: 根据配置选择量化路径,处理量化结果

接口原型:

def process_quant_result_fn(
    weight: torch.Tensor,
    dst_type: torch.dtype,
    config: QuantizeConfig
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
    ...

输入/输出参数:

参数名称 输入/输出 类型 描述 取值范围
weight 输入 torch.Tensor 待量化的权重张量 任意形状的张量
dst_type 输入 torch.dtype 目标量化数据类型 E4M3, E5M2
config 输入 QuantizeConfig 量化配置对象 合法的配置实例

返回参数:

参数名称 类型 描述 取值范围
weight_fwd torch.Tensor Forward 量化权重 N/A
scale_fwd torch.Tensor Forward 缩放因子 N/A
weight_bwd torch.Tensor Backward 量化权重 N/A
scale_bwd torch.Tensor Backward 缩放因子 N/A

异常处理:

  • recipe_name 不支持时,抛出 ValueError

约束说明:

  • 仅在支持 NPU 的环境下可用
  • torch_npu 提供相应算子支持

变更说明: 新增函数

调用参考代码:

from mindspeed.fsdp.quantization.module.linear_mxfp8 import process_quant_result_fn

weight_fwd, scale_fwd, weight_bwd, scale_bwd = process_quant_result_fn(
    weight=model_weight,
    dst_type=torch.float8_e4m3fn,
    config=quant_config
)
3.5.3 编程手册设计

编程手册内容规划:

  1. 快速开始章节:

    • MXFP8-32x32 量化配置示例
    • 端到端训练流程示例
  2. 配置说明章节:

    • Block size 选择指南
    • 性能调优建议
  3. API 参考章节:

    • 新增接口的详细说明
    • 参数配置示例
  4. 最佳实践章节:

    • 适用场景分析
    • 性能对比数据

输出方式:

  • 在现有的《MindSpeed 量化编程手册》中新增章节
  • 不单独输出,保持文档统一性

4. 测试设计

4.1 测试层次

单元测试: 配置验证、量化函数、All-gather 逻辑、模块功能
集成测试: FSDP 集成、模型转换集成
端到端测试: 完整训练流程、性能基准、精度验证

4.2 测试用例设计

配置测试: 验证 block size 合法性、配方注册、配置获取
量化函数测试: 验证 mxfp8 和 mxfp8-32x32 两种模式的量化逻辑
All-gather 测试: 验证通信优化效果,确保 MXFP8-32x32 减少 1 个张量传输
性能测试: 测量通信时间、吞吐量、内存占用
精度测试: 对比量化前后精度,确保误差在可接受范围

4.3 测试执行策略

使用 pytest 框架和 unittest.mock 工具
覆盖率目标: 单元测试 ≥ 90%,关键路径 100%
CI 集成: 单元测试自动运行,E2E 测试手动触发

5. 缺点和风险

说明潜在风险(Breaking Change、性能回退、复杂度提升、引入的安全问题)、负面影响(对现有功能/用户的冲击)、实现成本(代码量/维护成本/人力投入)、是否有API或版本兼容性、旧版本迁移方案问题等,给出应对措施。

6. 现有技术

参考其他项目/社区的类似设计,说明借鉴与差异。

7. 未解决问题

待社区讨论/决策的开放问题,如硬件适配范围、参数默认值等(需在RFC通过前解决)。

附录

  • 参考资料链接
  • 术语表
  • 文档更新计划

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
kyle_zhangkyle_zhang
6月14日 修改了issue 的描述
Ggp513成员
6月22日 issue类型由 Bug-Report 改变为 RFC
Ggp513成员
6月22日 添加了label:rfc
ascend-robotascend-robot成员
7月3日 关联了看板:MindStudio ISSUE管理
Ggp513成员
7月11日 issue状态由 TODO 改变为 DONE
Ggp513成员
7月11日 关闭了 issue
ascend-robotascend-robot成员
7月11日 添加了label:resolved