issue待分派,添加triage-review标签


torchao_npu/
├── README.md
├── LICENSE
├── setup.py
├── benchmarks/ # 基准测试(不变)
├── tests/ # 单元测试(不变)
├── examples/ # 使用示例(不变)
├── torchao_npu/ # 主包
│ ├── init.py # 包入口,注册各模块
│ ├── interfaces/ # 外部框架对接(如 torchtitan)
│ │ ├── init.py
│ │ └── torchtitan.py
│ ├── core/ # 核心三层抽象
│ │ ├── config/ # —— 策略配置层 ——
│ │ │ ├── init.py
│ │ │ ├── configs.py # ParamSwapConfig,扩展 QATConfig
│ │ │ ├── quant_configs.py # MXQuantizeConfig / BlockQuantizeConfig
│ │ │ └── strategies.py # 昇腾亲和量化策略(训练/MOE/QAT,对应架构中的策略)
│ │ ├── wrapper_tensor/ # —— 张量抽象层 ——
│ │ │ ├── init.py
│ │ │ ├── base_wrapper_tensor.py # BaseTrainingWeightWrapperTensor
│ │ │ ├── mx_wrapper_tensor.py # MXTrainingWeightWrapperTensor
│ │ │ ├── block_wrapper_tensor.py # BlockTrainingWeightWrapperTensor
│ │ │ └── float8_wrapper_tensor.py# Float8TrainingWeightWrapperTensor
│ │ └── ops/ # —— 算子实现层 ——
│ │ ├── init.py
│ │ ├── mx_ops.py # MXFP8 mm / _grouped_mm / linear
│ │ ├── block_ops.py # Block FP8 mm / _grouped_mm
│ │ └── float8_ops.py # FP8 row-wise 算子
│ └── patches/ # 算子劫持与分发(属于张量管理的运行时补丁)
│ ├── init.py
│ ├── mx_linear.py # 拦截/替换 torchao 的 CUDA linear
│ ├── mxfp8_grouped_mm.py # 替换 grouped_mm 至 NPU 实现
│ ├── mx_ac_config.py # activation checkpoint 配置适配
│ ├── mx_capability_check.py # NPU 能力探测
│ └── activation_checkpoint_state.py # AC 状态下的量化感知
└── (其余顶层目录保持不变)


状态(Status): Draft / Reviewing / Approved / Rejected / Superseded
作者(Authors): @wanglijun55
创建日期(Created): 2026-08-18
更新日期(Updated): 2026-08-18
相关 Issue/PR: 无
背景与动机
1. torchao 简介
torchao(PyTorch Architecture Optimization)是 PyTorch 官方维护的量化与稀疏化工具库(仓库地址:
pytorch/ao)。它构建在torch.matmul/torch._grouped_mm等核心算子之上,通过quantize_API 对nn.Module进行量化变换,提供从推理量化(PTQ)到量化感知训练(QAT)的完整能力,是 PyTorch 生态中模型低精度化的主要入口。torchao 现有架构在非 CUDA 平台上存在以下问题和不足:
float8/mx系列的mm、_grouped_mm、linear)默认调用 CUDA kernel,权重/激活直接落盘到 CUDA 设备,在 NPU 上会因算子缺失或设备不匹配而报错quant_primitives中的缩放/反缩放、转置、cast 等实现依赖 CUDA 张量语义,缺乏对 NPUtorch.npu原生量化算子的对接路径AOBaseConfig/QATConfig的prepare/convert流程在quantize_module_handler注册机制下进行模块替换,难以为 NPU 定制「参数级包装 + 算子拦截」的训练路径而不侵入上游2. 为什么要适配昇腾 NPU
torchao 是 PyTorch 量化的官方主线:随着 PyTorch 2.x 持续把量化能力收敛到 torchao,NPU 平台需要及早对接,避免上层训练框架(如 torchtitan-npu)被迫维护私有量化分叉。
NPU 已具备原生低精度算子能力:
torch_npu与 CANN 在 MXFP8 / Block FP8 上提供了原生算子(如npu_quantize、npu_mxfp8_mm、npu_grouped_mm等),具备在训练前向/反向中真正执行低精度计算的条件,而非仅做伪量化。需要一座桥接层把这些算子暴露给 torchao 的量化流程。解耦设计降低适配成本:torchao 的
register_quantize_module_handler注册机制与AOBaseConfig配置体系允许 out-of-tree 扩展,新增 NPU 后端无需修改 torchao 核心代码,极大降低昇腾平台的适配和维护成本。面向大模型训练:在 DeepSeek-V3 / Flash 等超大 MoE 模型训练中,MXFP8 / Block FP8 是降显存、提吞吐的关键手段。已有实验数据表明,在 NPU 上对 Flash 模型启用
all_mxfp8训练相比 BF16 基线有约 6% 的吞吐提升(all_block_fp8约 3%),亟需将其沉淀为独立、可演进的项目。建议方案
1. 总体方案:独立代码仓 torchao_npu
方案对比
torchao_npu,实现 NPU 量化算子对接、wrapper tensor、patch 机制,并通过 torchao 配置注册集成pytorch/ao仓库中增加 NPU 后端代码torch_npu,上游难以合入带硬件绑定的实现综合考量:建议选择方案一,以独立代码仓方式提供基于 NPU 的量化实现。主要原因:
torch_npu提供的原生算子(带硬件绑定),此类代码不符合上游「无硬件依赖」的合入准则;pytorch/ao由 Meta 维护,外部贡献的 review 和合入周期不可控;register_quantize_module_handler、AOBaseConfig)本身就是为了支持这种独立开发模式。代码仓命名与结构
代码仓名称:
torchao_npu建议的仓库结构(基于现有代码迁移):
2. torchao 整体架构与交互流程
2.1 torchao 架构总览与 NPU 接入点
2.2 torchao_npu 与 torchao 的交互流程
关键集成点说明:
ParamSwapConfig:继承 torchao 的QATConfig,复用其prepare/convert两阶段机制。与上游「模块级替换(把nn.Linear换成FakeQuantizedLinear)」不同,ParamSwap 走「参数级包装」:把匹配的nn.Parameter.data包成BaseTrainingWeightWrapperTensor子类。__torch_function__拦截:wrapper tensor 子类重写__torch_function__,在torch.mm/torch._grouped_mm/F.linear等算子被调用时,把计算重定向到ops/mx_ops.py等模块中调用torch.npu原生低精度算子的实现。patches/运行时补丁:对 torchao 内部仍走 CUDA 的代码路径(如mx_linear、mxfp8_grouped_mm),通过 monkey-patch 在运行时替换为 NPU 实现,避免 fork torchao 源码。3. torchao_npu 具体适配方案
3.1 需要实现的模块和功能
基于对
CANN/torchtitan-npu中experiments/ao_npu/torchao_npu的分析,独立仓需要包含以下模块:3.1.1 ParamSwapConfig — 参数级量化配置(
configs.py)对应 torchao 参考:
torchao.quantization.qat.QATConfig功能:扩展
QATConfig,在prepare阶段把匹配的nn.Parameter包装为 wrapper tensor,在convert阶段还原。class ParamSwapConfig(QATConfig): """ 通过参数包装实现低精度训练的配置,配合 quantize_ 使用。 复用 QATConfig 的 prepare/convert 机制,但走「参数级包装」而非 「模块级替换」。匹配的 nn.Parameter.data 被包成 BaseTrainingWeightWrapperTensor 子类,通过 __torch_function__ 拦截 torch.mm / torch._grouped_mm 等,施加子类专属的精度变换 (带 STE 的伪量化,或真实低精度 matmul)。 支持:FP8 row-wise、NPU MX block-wise、NPU Block FP8。 """ def __init__( self, base_config: AOBaseConfig | None = None, activation_config: FakeQuantizeConfigBase | None = None, weight_config: FakeQuantizeConfigBase | None = None, *, step: QATStep = QATStep.PREPARE, params_filter_fn: ParameterFilterFn = _is_parameter, ): ... @register_quantize_module_handler(ParamSwapConfig) def _param_swap_config_transform(module, config): ...3.1.2 wrapper_tensors — 训练时权重包装(
wrapper_tensors/)对应 torchao 参考:torchao 的 tensor subclass 体系
功能:定义拦截计算算子的 tensor 子类,是 NPU 低精度算子注入的核心。
# base_wrapper_tensor.py class BaseTrainingWeightWrapperTensor(torch.Tensor): """基类:持有原始权重 + 量化参数,重写 __torch_function__ 把 mm/_grouped_mm/linear 路由到 ops/*。""" # mx_wrapper_tensor.py class MXTrainingWeightWrapperTensor(BaseTrainingWeightWrapperTensor): """MXFP8 权重包装,路由到 ops/mx_ops。""" # block_wrapper_tensor.py class BlockTrainingWeightWrapperTensor(BaseTrainingWeightWrapperTensor): """Block FP8 权重包装,路由到 ops/block_ops。""" # float8_wrapper_tensor.py class Float8TrainingWeightWrapperTensor(BaseTrainingWeightWrapperTensor): """FP8 row-wise 权重包装,路由到 ops/float8_ops。"""3.1.3 ops — NPU 低精度算子封装(
ops/)对应 torchao 参考:
torchao.quantization.quant_primitives+ CUDA kernel功能:把
torch.npu的原生低精度算子封装为 mm / _grouped_mm / linear 接口,供 wrapper tensor 调用。# ops/mx_ops.py def npu_mxfp8_mm(...) -> torch.Tensor: ... # 调 torch.npu MXFP8 matmul def npu_mxfp8_grouped_mm(...) -> torch.Tensor: ... def npu_mxfp8_linear(...) -> torch.Tensor: ... # ops/block_ops.py def npu_block_fp8_mm(...) -> torch.Tensor: ... def npu_block_fp8_grouped_mm(...) -> torch.Tensor: ... # ops/float8_ops.py def npu_float8_row_mm(...) -> torch.Tensor: ...3.1.4 quantization — 量化配置与变换(
quantization/)对应 torchao 参考:
torchao.quantization.quant_config+transform_module功能:定义 NPU 专属量化配置与参数变换函数。
# quant_configs.py class MXQuantizeConfig(FakeQuantizeConfigBase): ... # NPU MX block-wise class BlockQuantizeConfig(FakeQuantizeConfigBase): ... # NPU Block FP8 # transform.py _PARAM_SWAP_QUANTIZE_CONFIG_HANDLER = { Float8FakeQuantizeConfig: ..., MXQuantizeConfig: ..., BlockQuantizeConfig: ..., } def _replace_params_with_custom_fn_if_matches_filter(module, fn, filter_fn, extra_args): ... def unwrap_param(param, config): ... # quant_primitives/mx.py, block_fp8.py def quantize_mx(...) / dequantize_mx(...) ... def quantize_block_fp8(...) / dequantize_block_fp8(...) ...3.1.5 patches — 运行时补丁(
patches/)功能:对 torchao / torchtitan 中仍走 CUDA 的路径做运行时替换,避免 fork 上游源码。
# patches/mx_linear.py — 替换 torchao 的 CUDA mx linear # patches/mxfp8_grouped_mm.py — 替换 grouped_mm 至 NPU 实现 # patches/mx_ac_config.py — activation checkpoint 配置适配 # patches/mx_capability_check.py — NPU 能力探测(避免 CUDA capability 检查失败) # patches/activation_checkpoint_state.py — AC 状态下的量化感知3.2 模块关系与集成图
3.3 与上游 torchao 的文件/概念映射关系
QATConfigconfigs.ParamSwapConfigquantize_module_handler注册configs._param_swap_config_transformFakeQuantizeConfigBasequantization.quant_configs.MXQuantizeConfig/BlockQuantizeConfigwrapper_tensors.*TrainingWeightWrapperTensorquant_primitives+ CUDA kernelops/{mx,block,float8}_ops.py+quant_primitives/transform_module._replace_params...quantization.transformpatches/mx_linear.py/mxfp8_grouped_mm.py3.4 支持的量化 Recipe
all_mxfp8mixall_block_fp83.5 依赖与性能基线
依赖:
性能基线(Flash 模型,单节点训练):
all_mxfp8all_block_fp83.6 使用示例
示例一:MXFP8 训练(prepare → train → convert)
import torch import torchao_npu # 触发 wrapper/ops/patches 注册 from torchao.quantization import quantize_ from torchao_npu.configs import ParamSwapConfig from torchao_npu.quantization.quant_configs import MXQuantizeConfig # 1. prepare:把匹配的权重包成 MXTrainingWeightWrapperTensor quantize_( model, ParamSwapConfig(weight_config=MXQuantizeConfig()), ) # 2. 正常训练前向/反向,mm/_grouped_mm 自动走 torch.npu MXFP8 算子 for x, y in dataloader: loss = model(x).loss loss.backward() # 3. convert:还原为普通权重 quantize_( model, ParamSwapConfig(weight_config=MXQuantizeConfig(), step=QATStep.CONVERT), )示例二:MoE 混合精度(mix recipe)
from torchao_npu.configs import ParamSwapConfig from torchao_npu.quantization.quant_configs import MXQuantizeConfig, BlockQuantizeConfig # Attention/Shared Expert → MXFP8;Routed Expert → Block FP8 quantize_(model, ParamSwapConfig( weight_config=..., # 由 recipe_converter 按 module filter 分配 params_filter_fn=moe_filter, ))迁移与建仓计划
CANN/torchtitan-npu的torchtitan_npu/experiments/ao_npu/torchao_npu整体迁移至新仓Ascend/torchao_npu顶层包;benchmarks/、tests/一并迁移。torchtitan_npu内部模块的引用收敛到interfaces/torchtitan.py一个文件,便于两个仓独立演进。参考资料:
https://github.com/pytorch/aohttps://gitcode.com/cann/torchtitan-npu/tree/master/torchtitan_npu/experiments/ao_npuhttps://gitcode.com/Ascend/pytorch/issues/3087https://gitcode.com/Ascend/pytorch