已关闭
[Requirement|需求建议]: 新增DualMatmulAdd算子torch接口 #4403
dutiegang创建于  7月29日关闭于  8月6日
dutiegang
7月29日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud(背景信息)

TX HY3 针对Router GEMM算子BF16*FP32计算,将权重FP32进行拆分了2个BF16(高16位与低16位),将原来的FP32计算转换为BF16计算,充分利用硬件的BF16算力,提升算子计算性能

Origin(信息来源)

TX HY3网络模型提出来的Router GEMM算子性能优化方案
https://www.163.com/dy/article/L0CE3E5P0518R7MO.html?referFrom=
基于GPU的算子实现
https://github.com/Tencent/hpc-ops/blob/main/src/gemm/sm90/gemm_bf16xfp32.cu

Benefit / Necessity (价值/作用)

NPU 910D平台下(32核), BF16算力432T FP32算力 27T

算力bound场景,理论性能,可提升8x
搬运bound场景,无收益

Design(设计方案)

离线权重拆分,将 FP32 权重 W 分解为高位 BF16 分量与低位残差 BF16 分量:
W_fp32 ≈ W_high + 其中:scale × W_low

W_high = W_fp32.to(bfloat16) — 截取 FP32 权重的高 16 位(主精度部分)
W_low = ((W_fp32 - W_high.float()) / scale).to(bfloat16) — 残差部分,缩放后再次截断为 BF16
scale = 1 / 256 — 固定缩放因子,对齐 BF16 的 8 位尾数(2^8 = 256)
推理阶段计算公式:

Y = X × W_high + scale × (X × W_low)
两路 BF16 GEMM 均运行在 Tensor Core 的 BF16×BF16→F32 MMA 指令上,激活值全程 BF16 无需类型转换,充分发挥硬件算力。

算子接口

def dual_matmul_add(
x: torch.Tensor,
w_high: torch.Tensor,
w_low: torch.Tensor,
*,
w_low_scale: float = 0.00390625,
y_dtype: int = 0,
)

接口调用示例

import cann_ops_nn
x = torch.tensor(m,k)
w = torch.tensor(k,n)
w_high = w.to(torch.bfloat16)
w_low = (w - w_high) / 256
w_low_scale = 1.0 / 256.0
y_dtype = 0
cann_ops_nn.dual_add_matmul(x, w_high, w_low, w_low_scale=w_low_scale, y_dype=y_dtype)

likedislike
yuning_chenyuning_chen成员
7月29日 将 dutiegang 设为负责人
Wwuyufei成员
7月30日 将 wuyufei 设为负责人
Wwuyufei成员
7月30日 修改标题为 “[Requirement|需求建议]: 新增DualMatmulAdd算子torch接口”,原标题为“[Requirement|需求建议]: 新增DualMatmulAdd算子”
Wwuyufei成员
7月31日 修改了issue 的描述
Wwuyufei成员
8月1日 修改了issue 的描述
CANN-robotCANN-robot成员
8月6日 关闭了 issue
CANN-robotCANN-robot成员
8月6日 添加了label:resolved