昇腾平台融合算子 dequant_swiglu_quant 的设计与实现

📅 2026/8/1 11:22:37 👁️ 阅读次数 📝 编程学习
昇腾平台融合算子 dequant_swiglu_quant 的设计与实现

作者​:昇腾实战派
知识地图​:https://blog.csdn.net/Lumos_Lovegood/article/details/161601003

背景概述

在深度学习推理场景中,模型量化与激活函数的组合操作频繁出现,通常需要依次执行反量化(Dequant)、激活函数(如 SwiGLU)和量化(Quant)三个步骤。传统分步执行方式会产生大量中间张量的显存读写,导致推理延迟增加。为解决这一问题,本文设计并实现了一个融合算子dequant_swiglu_quant,将上述三个操作合并为一次 kernel 调用,显著减少显存访问开销,提升推理性能。该算子基于 Triton-Ascend DSL 开发,运行于 Ascend NPU 平台。

1. 算子功能概述

dequant_swiglu_quant是一个融合算子,将反量化(Dequant)、SwiGLU 激活、量化(Quant)三个操作融合为一次 kernel 调用,减少中间结果的显存读写开销,提升推理性能。

该算子对标torch_npu.npu_dequant_swiglu_quantNPU 原生算子,使用 Triton-Ascend DSL 实现,在 Ascend NPU 上运行。

1.1 计算流程

输入 x [TokensNum, 2H] │ ├─ Dequant(反量化) │ ├─ x = x * weight_scale (权重反量化,INT32 输入时) │ ├─ x = x * activation_scale (激活反量化,INT32 输入时) │ └─ x = x + bias (可选偏置) │ ├─ SwiGLU(激活) │ ├─ 将 x 沿最后一维拆分为 A[:, 0:H] 和 B[:, H:2H] │ ├─ 标准 SwiGLU: swish(A) * B (activate_left=True) │ └─ 变种 SwiGLU: clamp + swish(z, α) * (z_linear + bias) │ ├─ Smooth Quant(平滑量化,可选) │ └─ out = out * quant_scale │ └─ Quant(量化) ├─ 静态量化: out = clamp(round(out / quant_scale + quant_offset), -max, max) └─ 动态量化: scale = max(|out|); out = clamp(round(out / scale), -max, max) │ 输出 output [TokensNum, H], scale [TokensNum]

1.2 分组量化

支持 count 模式的分组量化,通过group_index参数指定每个分组的 token 数量。每组使用不同的 scale 参数(weight_scale、activation_scale、quant_scale)。

示例:x.shape = [128, 2H],group_index = [2, 1, 3],表示 3 个分组:

  • group0 = x[0:2, :],使用 scale[0, :]
  • group1 = x[2:3, :],使用 scale[1, :]
  • group2 = x[3:6, :],使用 scale[2, :]

2. 算子接口

2.1 函数签名

defdequant_swiglu_quant(x,*,weight_scale=None,activation_scale=None,bias=None,quant_scale=None,quant_offset=None,group_index=None,activate_left=False,quant_mode=0,swiglu_mode=0,clamp_limit=7.0,glu_alpha=1.702,glu_bias=1.0,dst_type=torch.int8,round_mode="rint",)->(Tensor,Tensor)

2.2 参数说明

必选参数

参数类型形状说明
xTensor[TokensNum, 2H]输入张量,支持 int32 / bfloat16,最后一维必须为偶数

可选参数

参数类型形状默认值说明
weight_scaleTensor[groupNum, 2H]None权重反量化系数,float32。int32 输入时必选
activation_scaleTensor[TokensNum, 1]None激活反量化系数,float32。int32 输入时必选
biasTensor-None偏置,int32。group_index 非 None 时必须为 None
quant_scaleTensor[groupNum, H]None平滑量化系数,float32
quant_offsetTensor-None量化偏移,float32。group_index 非 None 时必须为 None
group_indexTensor[groupNum]None分组索引(count 模式),int64
activate_leftbool-FalseTrue: swish(A) * B;False: A * swish(B)
quant_modeint-00=静态量化,1=动态量化
swiglu_modeint-00=标准 SwiGLU,1=变种 SwiGLU
clamp_limitfloat-7.0变种 SwiGLU 的 clamp 限制
glu_alphafloat-1.702变种 SwiGLU 的 alpha 参数
glu_biasfloat-1.0变种 SwiGLU 的 bias 参数
dst_typetorch.dtype-torch.int8输出类型:int8 / float8_e4m3fn / float8_e5m2
round_modestr-“rint”舍入模式:rint(银行家舍入)/ floor(向下取整)

2.3 返回值

输出类型形状说明
outputTensor[TokensNum, H]量化输出,dtype 由 dst_type 决定
scaleTensor[TokensNum]量化 scale,float32

3. 计算公式

3.1 反量化(Dequant)

INT32 输入

x_float = x * weight_scale * activation_scale + bias

BF16 输入

x_float = x # 无需反量化,直接使用

3.2 SwiGLU 激活

将 x_float 沿最后一维拆分为 A = x_float[:, 0:H] 和 B = x_float[:, H:2H]。

标准 SwiGLU(swiglu_mode=0)

左激活(activate_left=True):

output = swish(A) * B

右激活(activate_left=False):

output = A * swish(B)

其中 swish(z) = z * sigmoid(z),sigmoid(z) = 1 / (1 + exp(-z))

变种 SwiGLU(swiglu_mode=1)

按奇偶交错拆分:

x_glu = clamp(x_even, max=clamp_limit) x_linear = clamp(x_odd, -clamp_limit, clamp_limit) output = swish(x_glu, α) * (x_linear + glu_bias)

其中 swish(z, α) = z * sigmoid(α * z)

3.3 平滑量化(Smooth Quant,可选)

output = output * quant_scale

3.4 量化(Quant)

静态量化(quant_mode=0)

output = clamp(round(output / quant_scale + quant_offset), -max_val, max_val) scale = quant_scale # 静态量化时 scale 为输入参数

动态量化(quant_mode=1)

scale = max(|output|) / max_val # 逐行求最大绝对值 output = clamp(round(output / scale), -max_val, max_val)

max_val 取值:

  • INT8: 127.0
  • FP8 E4M3FN: 448.0
  • FP8 E5M2: 57344.0

4. 约束条件

4.1 输入类型约束

输入类型weight_scaleactivation_scalebias说明
int32必选必选可选需要反量化
bfloat16必须为 None必须为 None必须为 None无需反量化

4.2 分组量化约束

  • group_index仅支持动态量化(quant_mode=1)
  • group_index非 None 时,bias 和 quant_offset 必须为 None
  • group_index求和不超过 TokensNum
  • group_index为 count 模式,每个元素表示该分组的 token 数量

4.3 形状约束

  • x 必须为 2D 张量,最后一维为偶数(2H)
  • weight_scale 形状:[groupNum, 2H](单组时 groupNum=1)
  • activation_scale 形状:[TokensNum, 1]
  • quant_scale 形状:[groupNum, H]
  • group_index 形状:[groupNum]

4.4 其他约束

  • clamp_limit、glu_alpha、glu_bias 仅在 swiglu_mode=1 时生效
  • 输出 out 和 scale 超过 group_index 总和的部分为未定义数据

5. 实现架构

5.1 文件结构

src/ ├── dequant_swiglu_quant.py # 算子入口,参数验证、分组 scale 展开、kernel 调度 ├── dequant_swiglu_quant_static_base.py # 静态量化 kernel └── dequant_swiglu_quant_dynamic_base.py # 动态量化 kernel

5.2 Kernel 设计

静态量化 Kernel

  • 单阶段处理:反量化 → SwiGLU → 平滑量化 → 静态量化,数据在寄存器中流转
  • 无中间缓冲区:所有计算在寄存器中完成,减少显存访问
  • 支持 quant_offset:静态量化特有的偏移参数

动态量化 Kernel

  • 两阶段处理
    1. 第一阶段:反量化 → SwiGLU → 平滑量化 → 求行级 ReduceMax
    2. 第二阶段:使用 ReduceMax 结果计算 scale → 量化输出
  • 需要中间缓冲区swiglu_tmp暂存 SwiGLU 结果,供第二阶段使用

5.3 辅助 Kernel

函数功能说明
sigmoid_kernel计算 sigmoid1.0 / (1.0 + exp(-x))
swish_kernel计算 swishx * sigmoid(x)
rint_kernel银行家舍入round half to even,匹配 NPU 的 CAST_RINT

5.4 分组 Scale 展开

入口函数中通过_expand_group_scale将分组 scale 展开为逐行 scale:

  1. 根据group_index计算row_to_group映射
  2. 使用 advanced indexing 展开:scale[row_to_group]
  3. 单组(groupNum=1)时 squeeze 为 1D

5.5 BLOCK_SIZE 配置

BLOCK_M 和 BLOCK_N 通过 triton.autotune 自动寻优,不在此处固定配置。
优化目标:确保不超出 NPU UB 容量限制(约 196 KB)。

6. 舍入模式

6.1 rint(银行家舍入,round half to even)

默认舍入模式,匹配 NPU 的 CAST_RINT 操作:

  • 非 x.5 值:标准四舍五入
  • x.5 值:舍入到最近的偶数(如 2.5 → 2.0,3.5 → 4.0)

实现逻辑:

floor_x=floor(x)frac=x-floor_x is_half=(frac==0.5)is_even=(int(floor_x)&1)==0result=where(is_half&is_even,floor_x,floor_x+1.0)result=where(is_half,result,where(frac>=0.5,floor_x+1.0,floor_x))

6.2 floor(向下取整)

直接使用tl.floor(x)实现。

7. 精度说明

7.1 INT8 输出精度

由于 Triton 和 NPU 的 SwiGLU 中间浮点计算存在 ULP(Unit in the Last Place)级别的差异,经 x.5 边界舍入放大后,可能导致极少数 INT8 输出元素差 ±1。

这是浮点运算的固有特性,不是实现 bug。在精度测试中,允许极少量 INT8 ±1 差异(比例 ≤ 1e-5)。

7.2 Scale 精度

动态量化时,scale 输出与 NPU 参考实现完全一致(float32 精度范围内)。

静态量化时,NPU 的 scale 输出语义不明确,精度测试中不检查 scale。

8. 性能特征

8.1 融合优势

相比分步执行(反量化 → SwiGLU → 量化),融合算子:

  • 减少中间结果的显存读写(2 次完整读写 → 0 次)
  • 减少 kernel launch 开销(3 次 → 1 次)
  • 提高数据局部性,更好地利用 NPU UB 缓存

8.2 静态 vs 动态量化

特性静态量化动态量化
Kernel 阶段单阶段两阶段
中间缓冲区不需要需要 swiglu_tmp
量化 scale输入参数运行时计算
延迟更低稍高
精度依赖 quant_scale 质量自适应,精度更稳定

8.3 典型性能数据

INT32 动态量化单组(NPU: Atlas 800I A2):

ShapeTriton (ms)NPU (ms)加速比
(64, 512)0.0120.0060.50
(1024, 2048)0.1250.0310.25
(4096, 8192)1.6180.5000.31

BF16 动态量化单组:

ShapeTriton (ms)NPU (ms)加速比
(64, 512)0.0100.0070.70
(1024, 2048)0.0930.0210.23
(4096, 8192)1.1690.2880.25

注:当前 Triton 实现与 NPU 原生算子仍有性能差距,后续可通过优化 BLOCK_SIZE、向量化策略等提升性能。

9. 测试

9.1 精度测试

cdtests pytest test_accuracy_dequant_swiglu_quant.py-v-k"not TestFP8Output"

9.2 性能测试

cdtests python test_benchmark_dequant_swiglu_quant.py

性能测试结果保存到../perf_time/../perf_throughput/目录。