已开启
[Feature]: torch_npu.npu_mhc_pre_backward新增可选参数以支持开启HF32模式 #4763
hyhhh14创建于  8 天前
hyhhh14
8 天前 创建

提交提案之前,请先检索仓库内是否已有相同的提案,如已有请在同一提案中进行讨论。

💻 需求背景、当前现状、期望实现的功能内容、具体的设计方案、以及测试方案

需求背景

本需求面向 Ascend 950 上的 MHC 训练场景。对 MhcPreBackward 算子进行性能优化,并新增 aclnnMhcPreBackwardV2 接口,用于显式选择 Cube 的 HF32 计算模式。

相应地,torch_npu.npu_mhc_pre_backward 接口需要提供计算模式配置参数。

当前现状

当前 Torch 接口形式不包含 inner_precise 参数:

torch_npu.npu_mhc_pre_backward(
    x,
    phi,
    alpha,
    grad_h_in,
    grad_h_post,
    grad_h_res,
    inv_rms,
    h_mix,
    h_pre,
    h_post,
    gamma=None,
    hc_eps=1e-6,
    grad_x_post=None,
)

接口固定调用 aclnnMhcPreBackward,Cube 使用 FP32 模式,无法选择新增的 aclnnMhcPreBackwardV2 HF32 模式。

期望实现的功能

torch_npu.npu_mhc_pre_backward 新增可选参数:

inner_precise: int = 0

参数含义如下:

  • inner_precise=0:调用 aclnnMhcPreBackward,Cube 使用 FP32 模式。
  • inner_precise=1:调用 aclnnMhcPreBackwardV2,Cube 使用 HF32 模式。
  • 其他取值:返回参数错误。

更新后的接口形式为:

torch_npu.npu_mhc_pre_backward(
    x,
    phi,
    alpha,
    grad_h_in,
    grad_h_post,
    grad_h_res,
    inv_rms,
    h_mix,
    h_pre,
    h_post,
    gamma=None,
    hc_eps=1e-6,
    grad_x_post=None,
    inner_precise=0,
)

inner_precise 默认为 0,因此现有业务代码不需要修改,接口行为和精度模式保持不变。用户只有显式传入 inner_precise=1 时才启用 HF32。

具体设计方案

  1. npu_mhc_pre_backward 的 Torch Schema 和 C++ 接口中增加可选参数 inner_precise,默认值为 0

  2. 在 OpAPI 实现中校验参数范围,仅允许取值 01

  3. 根据 inner_precise 选择底层 aclnn 接口:

    • inner_precise=0:调用 aclnnMhcPreBackward
    • inner_precise=1:调用 aclnnMhcPreBackwardV2,并将 inner_precise 透传为底层计算模式参数。
  4. 调用前检查对应 aclnn 接口是否可用。当运行环境不支持所选择的接口时,返回明确的版本或能力不支持提示。

  5. 输出 Tensor 的 Shape、数据类型和分配方式保持不变,不改变 gammagrad_x_post 等可选输入的现有语义。

  6. 保持反向兼容:未传入 inner_precise 的现有调用仍走 FP32 V1 接口。

测试方案

  1. FP32 兼容性测试

    验证不传 inner_precise 和显式传入 inner_precise=0 时均调用 V1 接口,结果与修改前保持一致。

  2. HF32 功能测试

    显式传入 inner_precise=1,验证能够正常调用 aclnnMhcPreBackwardV2,并与 CPU 参考结果进行精度比较。

  3. 输入组合覆盖

    FP32 和 HF32 模式分别覆盖:

    • BSND 和 TND 输入格式。
    • 传入和不传入 gamma
    • 传入和不传入 grad_x_post
  4. 参数校验测试

    验证 inner_precise 取非 0/1 值时能够返回明确的参数错误。

  5. 兼容性测试

    验证旧版本调用方式无需增加参数即可继续正常执行,输入输出 Shape、数据类型和返回值数量均不发生变化。

  6. 平台测试

    在 Ascend 950 环境运行相关 UT,确认:

    • inner_precise=0 实际调用 aclnnMhcPreBackward
    • inner_precise=1 实际调用 aclnnMhcPreBackwardV2
    • 两种模式均满足对应精度要求。

替代方案

补充说明

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

likedislike
TorchNPU-BotTorchNPU-Bot成员
8 天前 添加了label:triage-review
ascend-robotascend-robot成员
8 天前 添加了label:feature
TorchNPU-Bot
TorchNPU-Bot成员
8 天前 评论:

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

likedislike
Hhyhhh14
8 天前 修改了issue 的描述