作者​:昇腾实战派
知识地图​: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 必须为 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 计算 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:

  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) == 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/ 目录。

Logo

鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。

更多推荐