Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
TX HY3 针对Router GEMM算子BF16*FP32计算,将权重FP32进行拆分了2个BF16(高16位与低16位),将原来的FP32计算转换为BF16计算,充分利用硬件的BF16算力,提升算子计算性能
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
NPU 910D平台下(32核), BF16算力432T FP32算力 27T
算力bound场景,理论性能,可提升8x 搬运bound场景,无收益
离线权重拆分,将 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)
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)