状态(Status): Draft 作者(Authors): @kyle_zhangchi 创建日期(Created): 2026-06-14 更新日期(Updated): 2026-06-14 相关 Issue/PR: https://gitcode.com/Ascend/MindSpeed/pull/3534
本提案为 FSDP (Fully Sharded Data Parallel) 框架添加 MXFP8-32x32 量化支持,通过引入新的 block size 配置 [1, 32, 32],在保持模型精度的同时,进一步优化 all-gather 通信开销,提升大规模分布式训练性能。该方案扩展了现有的 MXFP8 量化能力,为用户提供更多量化粒度选择,以适应不同模型结构和硬件特性的优化需求。
[1, 32, 32]
当前痛点:
[1, 1, 32]
使用场景:
用户价值:
不做此提案的影响:
目标:
mxfp8-32x32
非目标 (边界说明):
功能点:
关键性能指标:
安全隐私及 DFX 要求:
使用限制:
npu_dynamic_block_mx_quant
整体设计思路:
本方案通过扩展现有的 MXFP8 量化框架,引入新的 block size 配置 [1, 32, 32],实现 MXFP8-32x32 量化。核心设计包括:
ScalingGranularity
PreQuantWeight
MXLinear
MXFP8GMM
fsdp_pre_all_gather
fsdp_post_all_gather
技术方案:
架构图:
┌─────────────────────────────────────────────────────────┐ │ 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 │ │ │ └─────────────┘ └──────────────┘ │ └─────────────────────────────────────────────────────────┘
限制和约束:
torch_npu
考虑过的方案:
方案一: 扩展现有 MXFP8,支持多种 block size
方案二: 独立实现 MXFP8-32x32 模块
方案三: 使用静态量化替代动态量化
实现方案:
__post_init__
register_mxfp8_32x32_recipe()
scaling_strategy=DYNAMIC
scaling_granularity=MX-[1, 32, 32]
核心流程:
用户指定 recipe_name="mxfp8-32x32" ↓ QuantRecipe 解析配置 ↓ ScalingGranularity 验证 block_size=[1, 32, 32] ↓ 创建 QuantizeConfig 实例 ↓ 模型转换器应用配置
数据模型变更:
ScalingGranularity.block_size
recipe_name="mxfp8-32x32"
在 PreQuantWeight.fsdp_pre_all_gather 中:
PreQuantWeight.fsdp_pre_all_gather
(weight_fwd, scale_fwd, weight_bwd, scale_bwd)
(weight_fwd, scale_fwd, scale_bwd)
在 PreQuantWeight.fsdp_post_all_gather 中:
PreQuantWeight.fsdp_post_all_gather
all_gather_outputs
性能影响:
影响范围:
mindspeed/fsdp/quantization/core/pre_quant_weight.py
mindspeed/fsdp/quantization/module/linear_mxfp8.py
mindspeed/fsdp/quantization/module/gmm_mxfp8.py
引入 process_quant_result_fn 函数,根据 recipe_name 选择量化路径:
process_quant_result_fn
recipe_name
npu_dynamic_mx_quant_with_dual_axis
输入权重 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) │ │ └───────────────────┴────────────────────┘ ↓ 返回量化结果
安全隐私:
兼容性:
可维护性:
可测试性:
可靠性:
开发环境设计:
开发约束:
可验收设计:
接口描述: 定义量化粒度配置,支持 MX 格式的 block size 设置
接口原型:
@dataclass class ScalingGranularity: stype: ScalingGranularityEnum block_size: Optional[List[int]] = None
输入/输出参数:
异常处理:
stype=MX
block_size
ValueError
约束说明:
变更说明: 新增 [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] )
接口描述: 注册 MXFP8-32x32 量化配方到全局配方字典
@recipe_register("mxfp8-32x32") def register_mxfp8_32x32_recipe() -> QuantRecipe: ...
输入/输出参数: 无参数
返回参数:
异常处理: 无
约束说明: 需在模型转换前调用,确保配方已注册
变更说明: 新增函数
from mindspeed.fsdp.quantization.config import get_recipe # 获取 MXFP8-32x32 配方 recipe = get_recipe("mxfp8-32x32")
接口描述: 根据配置选择量化路径,处理量化结果
def process_quant_result_fn( weight: torch.Tensor, dst_type: torch.dtype, config: QuantizeConfig ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: ...
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 )
编程手册内容规划:
快速开始章节:
配置说明章节:
API 参考章节:
最佳实践章节:
输出方式:
单元测试: 配置验证、量化函数、All-gather 逻辑、模块功能 集成测试: FSDP 集成、模型转换集成 端到端测试: 完整训练流程、性能基准、精度验证
配置测试: 验证 block size 合法性、配方注册、配置获取 量化函数测试: 验证 mxfp8 和 mxfp8-32x32 两种模式的量化逻辑 All-gather 测试: 验证通信优化效果,确保 MXFP8-32x32 减少 1 个张量传输 性能测试: 测量通信时间、吞吐量、内存占用 精度测试: 对比量化前后精度,确保误差在可接受范围
使用 pytest 框架和 unittest.mock 工具 覆盖率目标: 单元测试 ≥ 90%,关键路径 100% CI 集成: 单元测试自动运行,E2E 测试手动触发
说明潜在风险(Breaking Change、性能回退、复杂度提升、引入的安全问题)、负面影响(对现有功能/用户的冲击)、实现成本(代码量/维护成本/人力投入)、是否有API或版本兼容性、旧版本迁移方案问题等,给出应对措施。
参考其他项目/社区的类似设计,说明借鉴与差异。
待社区讨论/决策的开放问题,如硬件适配范围、参数默认值等(需在RFC通过前解决)。
欢迎加入社区,感谢您对社区的贡献 🎉!
状态(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 动机
当前痛点:
[1, 1, 32]block size 配置,在某些模型结构下无法充分挖掘量化潜力使用场景:
用户价值:
不做此提案的影响:
1.3 目标
目标:
[1, 32, 32]配置mxfp8-32x32,方便用户使用非目标 (边界说明):
2. 用例分析
2.1 大规模分布式训练场景
功能点:
关键性能指标:
安全隐私及 DFX 要求:
使用限制:
npu_dynamic_block_mx_quant算子[1, 1, 32]和[1, 32, 32]两种配置2.2 MoE (Mixture of Experts) 模型训练场景
功能点:
关键性能指标:
使用限制:
3. 方案设计
3.1 总体方案
整体设计思路:
本方案通过扩展现有的 MXFP8 量化框架,引入新的 block size 配置
[1, 32, 32],实现 MXFP8-32x32 量化。核心设计包括:ScalingGranularity类中支持新的 block size 验证,添加mxfp8-32x32配方注册PreQuantWeight、MXLinear、MXFP8GMM等核心类中,根据配置选择不同的量化路径fsdp_pre_all_gather和fsdp_post_all_gather逻辑,减少 scale 张量的传输技术方案:
npu_dynamic_block_mx_quant算子,支持 block-level 的动态量化[1, 32, 32]表示在 input_dim 和 output_dim 维度上按 32 分块,token_dim 不分块架构图:
限制和约束:
torch_npu支持npu_dynamic_block_mx_quant算子3.2 技术选型
考虑过的方案:
方案一: 扩展现有 MXFP8,支持多种 block size
方案二: 独立实现 MXFP8-32x32 模块
方案三: 使用静态量化替代动态量化
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核心流程:
数据模型变更:
ScalingGranularity.block_size支持[1, 32, 32]取值recipe_name="mxfp8-32x32"配方3.3.2 All-gather 优化功能
实现方案:
在
PreQuantWeight.fsdp_pre_all_gather中:(weight_fwd, scale_fwd, weight_bwd, scale_bwd),forward 和 backward 使用不同权重(weight_fwd, scale_fwd, scale_bwd),forward 和 backward 共享同一权重,减少传输量在
PreQuantWeight.fsdp_post_all_gather中:all_gather_outputs长度判断量化模式,正确恢复权重和 scale性能影响:
影响范围:
mindspeed/fsdp/quantization/core/pre_quant_weight.pymindspeed/fsdp/quantization/module/linear_mxfp8.pymindspeed/fsdp/quantization/module/gmm_mxfp8.py3.3.3 量化计算功能
实现方案:
引入
process_quant_result_fn函数,根据recipe_name选择量化路径:npu_dynamic_mx_quant_with_dual_axis,生成独立的 forward 和 backward 权重npu_dynamic_block_mx_quant,forward 和 backward 共享权重,仅 scale 不同核心流程:
3.4 安全隐私与 DFX 设计
安全隐私:
兼容性:
可维护性:
process_quant_result_fn统一量化逻辑,减少重复代码可测试性:
可靠性:
3.5 编程与调用设计
3.5.1 编程模型基本设计
开发环境设计:
开发约束:
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=MX且block_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: ...输入/输出参数: 无参数
返回参数:
异常处理: 无
约束说明: 需在模型转换前调用,确保配方已注册
变更说明: 新增函数
调用参考代码:
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]: ...输入/输出参数:
返回参数:
异常处理:
recipe_name不支持时,抛出ValueError约束说明:
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 编程手册设计
编程手册内容规划:
快速开始章节:
配置说明章节:
API 参考章节:
最佳实践章节:
输出方式:
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通过前解决)。
附录
欢迎加入社区,感谢您对社区的贡献 🎉!