已关闭
[RFC]:在Ascend下新建ao代码仓,适配torchao量化库至NPU #4242
wanglijun55创建于  8月20日关闭于  18 天前
wanglijun55成员
8月20日 创建

状态(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 平台上存在以下问题和不足:

问题 说明
CUDA 算子强绑定 torchao 的低精度算子(如 float8 / mx 系列的 mm、_grouped_mm、linear)默认调用 CUDA kernel,权重/激活直接落盘到 CUDA 设备,在 NPU 上会因算子缺失或设备不匹配而报错
量化原语面向 CUDA quant_primitives 中的缩放/反缩放、转置、cast 等实现依赖 CUDA 张量语义,缺乏对 NPU torch.npu 原生量化算子的对接路径
配置与调度耦合 AOBaseConfig / QATConfig 的 prepare/convert 流程在 quantize_module_handler 注册机制下进行模块替换,难以为 NPU 定制「参数级包装 + 算子拦截」的训练路径而不侵入上游
大模型训练量化缺口 torchao 主线聚焦推理与 CUDA 侧训练量化,MoE / 大模型在 NPU 上的 MXFP8 / Block FP8 训练量化长期缺乏官方对接方案

2. 为什么要适配昇腾 NPU

  1. torchao 是 PyTorch 量化的官方主线:随着 PyTorch 2.x 持续把量化能力收敛到 torchao,NPU 平台需要及早对接,避免上层训练框架(如 torchtitan-npu)被迫维护私有量化分叉。

  2. NPU 已具备原生低精度算子能力:torch_npu 与 CANN 在 MXFP8 / Block FP8 上提供了原生算子(如 npu_quantize、npu_mxfp8_mm、npu_grouped_mm 等),具备在训练前向/反向中真正执行低精度计算的条件,而非仅做伪量化。需要一座桥接层把这些算子暴露给 torchao 的量化流程。

  3. 解耦设计降低适配成本:torchao 的 register_quantize_module_handler 注册机制与 AOBaseConfig 配置体系允许 out-of-tree 扩展,新增 NPU 后端无需修改 torchao 核心代码,极大降低昇腾平台的适配和维护成本。

  4. 面向大模型训练:在 DeepSeek-V3 / Flash 等超大 MoE 模型训练中,MXFP8 / Block FP8 是降显存、提吞吐的关键手段。已有实验数据表明,在 NPU 上对 Flash 模型启用 all_mxfp8 训练相比 BF16 基线有约 6% 的吞吐提升(all_block_fp8 约 3%),亟需将其沉淀为独立、可演进的项目。


建议方案

1. 总体方案:独立代码仓 torchao_npu

方案对比

方案一:独立代码仓(推荐) 方案二:向 torchao 社区贡献
描述 创建独立仓 torchao_npu,实现 NPU 量化算子对接、wrapper tensor、patch 机制,并通过 torchao 配置注册集成 直接在 pytorch/ao 仓库中增加 NPU 后端代码
优点 自主可控,快速迭代;不受社区 review 周期限制;可灵活跟随 torchao 与 torch_npu 版本 与社区主线同步,减少维护分叉成本
缺点 需要跟随 torchao API 变更进行适配;需独立维护 社区 review 周期长;NPU 算子依赖 torch_npu,上游难以合入带硬件绑定的实现
风险 torchao API 变更导致的适配成本 贡献周期不可控,昇腾平台短期难以使用

综合考量:建议选择方案一,以独立代码仓方式提供基于 NPU 的量化实现。主要原因:

  1. torchao 的 NPU 适配强依赖 torch_npu 提供的原生算子(带硬件绑定),此类代码不符合上游「无硬件依赖」的合入准则;
  2. pytorch/ao 由 Meta 维护,外部贡献的 review 和合入周期不可控;
  3. torchao 的 out-of-tree 扩展点(register_quantize_module_handler、AOBaseConfig)本身就是为了支持这种独立开发模式。

代码仓命名与结构

代码仓名称:torchao_npu

现状说明:当前 torchao_npu 的代码以 experiments/ao_npu/torchao_npu 的形式寄生在 CANN/torchtitan-npu 仓库中。本 RFC 建议将其整体迁移为独立顶层仓 torchao_npu,与 torchcomms_npu 的建仓模式保持一致。

建议的仓库结构(基于现有代码迁移):

torchao_npu/
├── README.md                                  # 项目说明
├── LICENSE                                    # BSD-3-Clause 许可证
├── setup.py                                   # Python 包安装配置
├── torchao_npu/
│   ├── __init__.py                            # 包入口(触发 wrapper/ops 注册)
│   ├── configs.py                             # ParamSwapConfig(扩展 QATConfig)
│   ├── interfaces/
│   │   └── torchtitan.py                      # 与 torchtitan-npu 训练框架的对接接口
│   ├── ops/                                   # NPU 低精度算子封装
│   │   ├── __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/                               # 对 torchao / torchtitan 的运行时补丁
│   │   ├── __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 状态下的量化感知
│   ├── quantization/
│   │   ├── filters.py                         # ParameterFilterFn 等过滤函数
│   │   ├── quant_configs.py                   # MXQuantizeConfig / BlockQuantizeConfig
│   │   ├── transform.py                       # _replace_params_with_custom_fn_if_matches_filter / unwrap_param
│   │   └── quant_primitives/
│   │       ├── mx.py                          # MXFP8 量化/反量化原语
│   │       └── block_fp8.py                   # Block FP8 量化/反量化原语
│   └── wrapper_tensors/                       # 训练时权重 wrapper tensor 子类
│       ├── __init__.py
│       ├── base_wrapper_tensor.py            # BaseTrainingWeightWrapperTensor
│       ├── mx_wrapper_tensor.py              # MXTrainingWeightWrapperTensor
│       ├── block_wrapper_tensor.py           # BlockTrainingWeightWrapperTensor
│       └── float8_wrapper_tensor.py          # Float8TrainingWeightWrapperTensor
├── benchmarks/
│   └── e2e/
│       ├── dsv4_flash_single_node/           # 推理基准
│       └── dsv4_flash_single_node_train/    # 训练基准(config_registry / recipe_converter / run_train.sh)
├── tests/
│   ├── __init__.py
│   ├── conftest.py
│   ├── reference_moe.py                       # MoE 参考实现
│   ├── testing_utils.py
│   ├── test_everything.sh                     # 一键测试脚本
│   ├── test_param_swap_config.py
│   ├── test_training.py
│   ├── ops/
│   │   ├── test_mx_ops.py
│   │   ├── test_block_ops.py
│   │   └── test_compile_ops.py               # torch.compile 兼容性
│   └── wrapper_tensors/
│       ├── test_wrapper_tensor.py
│       ├── test_wrapper_ops.py
│       └── test_param_swap_transform.py
└── examples/
    └── mxfp8_training.py                      # MXFP8 训练最小示例

2. torchao 整体架构与交互流程

2.1 torchao 架构总览与 NPU 接入点

┌─────────────────────────────────────────────────────────────────┐
│                     Python 用户层                                │
│                                                                  │
│  from torchao.quantization import quantize_                      │
│  quantize_(model, ParamSwapConfig(weight_config=MXQuantizeConfig()))│
└──────────────────────┬──────────────────────────────────────────┘
                       │ register_quantize_module_handler
┌──────────────────────▼──────────────────────────────────────────┐
│                    torchao 核心层(pytorch/ao)                    │
│                                                                  │
│  ┌──────────────┐  ┌───────────────────┐  ┌──────────────────┐  │
│  │ AOBaseConfig │  │ QATConfig          │  │ quantize_module  │  │
│  │ (配置基类)    │──│ (prepare/convert) │──│ _handler 注册表  │  │
│  └──────┬───────┘  └────────┬──────────┘  └────────┬─────────┘  │
│         │ 继承/扩展           │                      │ dispatch    │
│  ┌──────▼──────────────────────────────────────────▼─────────┐  │
│  │  NPU 扩展点(torchao_npu,独立代码仓)                      │  │
│  │                                                            │  │
│  │  • ParamSwapConfig         ← 扩展 QATConfig               │  │
│  │  • MX/Block/Float8 QuantizeConfig                          │  │
│  │  • BaseTrainingWeightWrapperTensor(参数级包装)            │  │
│  │  • ops/{mx,block,float8}_ops(接入 torch.npu 算子)         │  │
│  │  • patches/*(运行时替换 torchao 的 CUDA 路径)            │  │
│  └──────────────────────────┬─────────────────────────────────┘  │
└─────────────────────────────┴───────────────────────────────────┘
                              │ 调用
┌─────────────────────────────▼───────────────────────────────────┐
│                    torch_npu / CANN 算子层                       │
│  npu_quantize / npu_mxfp8_mm / npu_grouped_mm / npu_block_mm …  │
└──────────────────────────────────────────────────────────────────┘

2.2 torchao_npu 与 torchao 的交互流程

用户训练代码                              torchao / torchao_npu
────────────────────                    ────────────────────────────
model = Flash(args)                      │
                                         │
# 1. prepare:参数包装                    │
quantize_(model,                          │
  ParamSwapConfig(                        │
    weight_config=MXQuantizeConfig()))    │
  │                                       │
  ├─ register_quantize_module_handler     │
  │   命中 ParamSwapConfig                │
  │                                       │
  ├─ _replace_params_with_custom_fn_      │
  │   if_matches_filter()                 │
  │   ├─ 对每个匹配 nn.Parameter          │
  │   │   套上 MXTrainingWeightWrapper    │
  │   │   Tensor(__torch_function__)    │
  │   └─ 完成 prepare                     │
                                         │
# 2. 训练前向                              │
out = model(x)                            │
  │                                       │
  ├─ torch.mm / torch._grouped_mm /      │
  │   F.linear 触发                       │
  │   └─ wrapper tensor 的 __torch_       │
  │       function__ 拦截                │
  │       └─ mx_ops / block_ops           │
  │           └─ torch.npu 低精度算子     │
  │                                       │
# 3. convert:训练后还原                  │
quantize_(model, ParamSwapConfig(         │
    ..., step=QATStep.CONVERT))           │
  └─ unwrap_param() 还原为普通 tensor     │

关键集成点说明:

  1. ParamSwapConfig:继承 torchao 的 QATConfig,复用其 prepare/convert 两阶段机制。与上游「模块级替换(把 nn.Linear 换成 FakeQuantizedLinear)」不同,ParamSwap 走「参数级包装」:把匹配的 nn.Parameter.data 包成 BaseTrainingWeightWrapperTensor 子类。

  2. __torch_function__ 拦截:wrapper tensor 子类重写 __torch_function__,在 torch.mm / torch._grouped_mm / F.linear 等算子被调用时,把计算重定向到 ops/mx_ops.py 等模块中调用 torch.npu 原生低精度算子的实现。

  3. 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 模块关系与集成图

┌──────────────────────────────────────────────────────────┐
│            torchao.quantization.quantize_               │
│              (注册表 dispatch)                            │
└──────────────────────┬───────────────────────────────────┘
                       │ 命中 ParamSwapConfig
┌──────────────────────▼───────────────────────────────────┐
│                  configs.ParamSwapConfig                  │
│  (继承 QATConfig, prepare→包装 / convert→还原)            │
└──────────┬───────────────────────────────────────────────┘
           │ prepare 时调用
┌──────────▼───────────────────────────────────────────────┐
│          quantization.transform                          │
│  _replace_params_with_custom_fn_if_matches_filter()      │
│  └─ 按 weight_config 类型查 _PARAM_SWAP_QUANTIZE_CONFIG  │
│     _HANDLER,把 nn.Parameter.data 包成 wrapper tensor   │
└──────────┬───────────────────────────────────────────────┘
           │ 产生
┌──────────▼───────────────────────────────────────────────┐
│        wrapper_tensors.*TrainingWeightWrapperTensor       │
│  (重写 __torch_function__, 拦截 mm/_grouped_mm/linear)   │
└──────────┬───────────────────────────────────────────────┘
           │ 路由到
┌──────────▼───────────────────────────────────────────────┐
│                   ops/{mx,block,float8}_ops              │
│  (封装 torch.npu 原生低精度算子)                          │
└──────────┬───────────────────────────────────────────────┘
           │ 调用
┌──────────▼───────────────────────────────────────────────┐
│            torch_npu / CANN 量化算子层                    │
│  npu_quantize / npu_mxfp8_mm / npu_grouped_mm / ...      │
└──────────────────────────────────────────────────────────┘

并行旁路:
┌──────────────────────────────────────────────────────────┐
│  patches/*  ——运行时 monkey-patch torchao / torchtitan   │
│  的 CUDA 路径,使其落到 ops/* 或 torch.npu                │
└──────────────────────────────────────────────────────────┘

3.3 与上游 torchao 的文件/概念映射关系

torchao 上游 torchao_npu 对应 功能
QATConfig configs.ParamSwapConfig 扩展为参数级包装的 QAT 配置
quantize_module_handler 注册 configs._param_swap_config_transform ParamSwap 的模块变换 handler
FakeQuantizeConfigBase quantization.quant_configs.MXQuantizeConfig / BlockQuantizeConfig NPU 专属量化配置
tensor subclass 体系 wrapper_tensors.*TrainingWeightWrapperTensor 拦截算子的权重包装
quant_primitives + CUDA kernel ops/{mx,block,float8}_ops.py + quant_primitives/ NPU 低精度算子封装
transform_module._replace_params... quantization.transform 参数级包装/还原
(上游走 CUDA 的 linear/grouped_mm) patches/mx_linear.py / mxfp8_grouped_mm.py 运行时替换为 NPU 路径

3.4 支持的量化 Recipe

Recipe 权重量化 激活量化 适用
all_mxfp8 全模块 MXFP8 MXFP8 通用 MoE/稠密训练
mix Attention/Shared Expert: MXFP8;Routed Expert: Block FP8(可选 MXFP4 伪量化) MXFP8 MoE 混合精度
all_block_fp8 全模块 Block FP8 MXFP8 Block FP8 训练

3.5 依赖与性能基线

依赖:

  • CANN >= 9.2.0-20260805
  • torch_npu >= 2.12.0-20260808
  • torchao >= 0.17

性能基线(Flash 模型,单节点训练):

Recipe TFLOPS 较 BF16 基线
BF16(baseline) 248.7 -
all_mxfp8 262.5 ~6%
all_block_fp8 256.9 ~3%

3.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,
))

迁移与建仓计划

  1. 代码迁移:将 CANN/torchtitan-npu 的 torchtitan_npu/experiments/ao_npu/torchao_npu 整体迁移至新仓 Ascend/torchao_npu 顶层包;benchmarks/、tests/ 一并迁移。
  2. 依赖解耦:把对 torchtitan_npu 内部模块的引用收敛到 interfaces/torchtitan.py 一个文件,便于两个仓独立演进。
  3. CI:搭建 torchao_npu 的单元测试 + e2e 基准 CI,对齐 torch_npu / CANN 版本矩阵。
  4. 文档:补齐 README、安装说明、3 种 recipe 的使用文档。

参考资料:

  1. torchao GitHub 仓库:https://github.com/pytorch/ao
  2. torchao_npu 现有代码(迁移源):https://gitcode.com/cann/torchtitan-npu/tree/master/torchtitan_npu/experiments/ao_npu
  3. 同类建仓参考 RFC(torchcomms_npu):https://gitcode.com/Ascend/pytorch/issues/3087
  4. torch_npu 仓库:https://gitcode.com/Ascend/pytorch
likedislike
Wwanglijun55成员
8月20日 添加了label:rfc
Wwanglijun55成员
8月20日 关联了看板:FrameworkPTAdapter 版本issue看板
TorchNPU-BotTorchNPU-Bot成员
8月20日 添加了label:triage-review
TorchNPU-Bot
TorchNPU-Bot成员
8月20日 评论:

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

likedislike
Wwanglijun55成员
8月20日 修改标题为 “[RFC]:在Ascend下新建ao代码仓,适配torchao量化库至NPU”,原标题为“[RFC]: 1”
Wwanglijun55成员
8月20日 修改了issue 的描述
Wwanglijun55成员
8月20日 将 wanglijun55 设为负责人
梁松伟
梁松伟成员
8月22日 评论:

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 状态下的量化感知
└── (其余顶层目录保持不变)

likedislike
Wwanglijun55成员
18 天前 issue状态由 TODO 改变为 DONE
Wwanglijun55成员
18 天前 关闭了 issue
ascend-robotascend-robot成员
18 天前 添加了label:resolved