昇腾平台融合算子 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_quant NPU 原生算子,使用 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 函数签名
def dequant_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 参数说明
必选参数
| 参数 | 类型 | 形状 | 说明 |
|---|---|---|---|
| x | Tensor | [TokensNum, 2H] | 输入张量,支持 int32 / bfloat16,最后一维必须为偶数 |
可选参数
| 参数 | 类型 | 形状 | 默认值 | 说明 |
|---|---|---|---|---|
| weight_scale | Tensor | [groupNum, 2H] | None | 权重反量化系数,float32。int32 输入时必选 |
| activation_scale | Tensor | [TokensNum, 1] | None | 激活反量化系数,float32。int32 输入时必选 |
| bias | Tensor | - | None | 偏置,int32。group_index 非 None 时必须为 None |
| quant_scale | Tensor | [groupNum, H] | None | 平滑量化系数,float32 |
| quant_offset | Tensor | - | None | 量化偏移,float32。group_index 非 None 时必须为 None |
| group_index | Tensor | [groupNum] | None | 分组索引(count 模式),int64 |
| activate_left | bool | - | False | True: swish(A) * B;False: A * swish(B) |
| quant_mode | int | - | 0 | 0=静态量化,1=动态量化 |
| swiglu_mode | int | - | 0 | 0=标准 SwiGLU,1=变种 SwiGLU |
| clamp_limit | float | - | 7.0 | 变种 SwiGLU 的 clamp 限制 |
| glu_alpha | float | - | 1.702 | 变种 SwiGLU 的 alpha 参数 |
| glu_bias | float | - | 1.0 | 变种 SwiGLU 的 bias 参数 |
| dst_type | torch.dtype | - | torch.int8 | 输出类型:int8 / float8_e4m3fn / float8_e5m2 |
| round_mode | str | - | “rint” | 舍入模式:rint(银行家舍入)/ floor(向下取整) |
2.3 返回值
| 输出 | 类型 | 形状 | 说明 |
|---|---|---|---|
| output | Tensor | [TokensNum, H] | 量化输出,dtype 由 dst_type 决定 |
| scale | Tensor | [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_scale | activation_scale | bias | 说明 |
|---|---|---|---|---|
| int32 | 必选 | 必选 | 可选 | 需要反量化 |
| bfloat16 | 必须为 None | 必须为 None | 必须为 None | 无需反量化 |
4.2 分组量化约束
group_index仅支持动态量化(quant_mode=1)group_index非 None 时,bias 和 quant_offset 必须为 Nonegroup_index求和不超过 TokensNumgroup_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
- 两阶段处理:
- 第一阶段:反量化 → SwiGLU → 平滑量化 → 求行级 ReduceMax
- 第二阶段:使用 ReduceMax 结果计算 scale → 量化输出
- 需要中间缓冲区:
swiglu_tmp暂存 SwiGLU 结果,供第二阶段使用
5.3 辅助 Kernel
| 函数 | 功能 | 说明 |
|---|---|---|
sigmoid_kernel |
计算 sigmoid | 1.0 / (1.0 + exp(-x)) |
swish_kernel |
计算 swish | x * sigmoid(x) |
rint_kernel |
银行家舍入 | round half to even,匹配 NPU 的 CAST_RINT |
5.4 分组 Scale 展开
入口函数中通过 _expand_group_scale 将分组 scale 展开为逐行 scale:
- 根据
group_index计算row_to_group映射 - 使用 advanced indexing 展开:
scale[row_to_group] - 单组(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) == 0
result = 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):
| Shape | Triton (ms) | NPU (ms) | 加速比 |
|---|---|---|---|
| (64, 512) | 0.012 | 0.006 | 0.50 |
| (1024, 2048) | 0.125 | 0.031 | 0.25 |
| (4096, 8192) | 1.618 | 0.500 | 0.31 |
BF16 动态量化单组:
| Shape | Triton (ms) | NPU (ms) | 加速比 |
|---|---|---|---|
| (64, 512) | 0.010 | 0.007 | 0.70 |
| (1024, 2048) | 0.093 | 0.021 | 0.23 |
| (4096, 8192) | 1.169 | 0.288 | 0.25 |
注:当前 Triton 实现与 NPU 原生算子仍有性能差距,后续可通过优化 BLOCK_SIZE、向量化策略等提升性能。
9. 测试
9.1 精度测试
cd tests
pytest test_accuracy_dequant_swiglu_quant.py -v -k "not TestFP8Output"
9.2 性能测试
cd tests
python test_benchmark_dequant_swiglu_quant.py
性能测试结果保存到 ../perf_time/ 和 ../perf_throughput/ 目录。
鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。
更多推荐


所有评论(0)