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

昇腾平台融合算子 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_leftTrue │ └─ 变种 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_scaleNone,activation_scaleNone,biasNone,quant_scaleNone,quant_offsetNone,group_indexNone,activate_leftFalse,quant_mode0,swiglu_mode0,clamp_limit7.0,glu_alpha1.702,glu_bias1.0,dst_typetorch.int8,round_moderint,)-(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 时必须为 Nonequant_scaleTensor[groupNum, H]None平滑量化系数float32quant_offsetTensor-None量化偏移float32。group_index 非 None 时必须为 Nonegroup_indexTensor[groupNum]None分组索引count 模式int64activate_leftbool-FalseTrue: swish(A) * BFalse: A * swish(B)quant_modeint-00静态量化1动态量化swiglu_modeint-00标准 SwiGLU1变种 SwiGLUclamp_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_e5m2round_modestr-“rint”舍入模式rint银行家舍入/ floor向下取整2.3 返回值输出类型形状说明outputTensor[TokensNum, H]量化输出dtype 由 dst_type 决定scaleTensor[TokensNum]量化 scalefloat323. 计算公式3.1 反量化DequantINT32 输入x_float x * weight_scale * activation_scale biasBF16 输入x_float x # 无需反量化直接使用3.2 SwiGLU 激活将 x_float 沿最后一维拆分为 A x_float[:, 0:H] 和 B x_float[:, H:2H]。标准 SwiGLUswiglu_mode0左激活activate_leftTrueoutput swish(A) * B右激活activate_leftFalseoutput A * swish(B)其中 swish(z) z * sigmoid(z)sigmoid(z) 1 / (1 exp(-z))变种 SwiGLUswiglu_mode1按奇偶交错拆分x_glu clamp(x_even, maxclamp_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_scale3.4 量化Quant静态量化quant_mode0output clamp(round(output / quant_scale quant_offset), -max_val, max_val) scale quant_scale # 静态量化时 scale 为输入参数动态量化quant_mode1scale max(|output|) / max_val # 逐行求最大绝对值 output clamp(round(output / scale), -max_val, max_val)max_val 取值INT8: 127.0FP8 E4M3FN: 448.0FP8 E5M2: 57344.04. 约束条件4.1 输入类型约束输入类型weight_scaleactivation_scalebias说明int32必选必选可选需要反量化bfloat16必须为 None必须为 None必须为 None无需反量化4.2 分组量化约束group_index仅支持动态量化quant_mode1group_index非 None 时bias 和 quant_offset 必须为 Nonegroup_index求和不超过 TokensNumgroup_index为 count 模式每个元素表示该分组的 token 数量4.3 形状约束x 必须为 2D 张量最后一维为偶数2Hweight_scale 形状[groupNum, 2H]单组时 groupNum1activation_scale 形状[TokensNum, 1]quant_scale 形状[groupNum, H]group_index 形状[groupNum]4.4 其他约束clamp_limit、glu_alpha、glu_bias 仅在 swiglu_mode1 时生效输出 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 # 动态量化 kernel5.2 Kernel 设计静态量化 Kernel单阶段处理反量化 → SwiGLU → 平滑量化 → 静态量化数据在寄存器中流转无中间缓冲区所有计算在寄存器中完成减少显存访问支持 quant_offset静态量化特有的偏移参数动态量化 Kernel两阶段处理第一阶段反量化 → SwiGLU → 平滑量化 → 求行级 ReduceMax第二阶段使用 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_RINT5.4 分组 Scale 展开入口函数中通过_expand_group_scale将分组 scale 展开为逐行 scale根据group_index计算row_to_group映射使用 advanced indexing 展开scale[row_to_group]单组groupNum1时 squeeze 为 1D5.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.03.5 → 4.0实现逻辑floor_xfloor(x)fracx-floor_x is_half(frac0.5)is_even(int(floor_x)1)0resultwhere(is_halfis_even,floor_x,floor_x1.0)resultwhere(is_half,result,where(frac0.5,floor_x1.0,floor_x))6.2 floor向下取整直接使用tl.floor(x)实现。7. 精度说明7.1 INT8 输出精度由于 Triton 和 NPU 的 SwiGLU 中间浮点计算存在 ULPUnit 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 A2ShapeTriton (ms)NPU (ms)加速比(64, 512)0.0120.0060.50(1024, 2048)0.1250.0310.25(4096, 8192)1.6180.5000.31BF16 动态量化单组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-knot TestFP8Output9.2 性能测试cdtests python test_benchmark_dequant_swiglu_quant.py性能测试结果保存到../perf_time/和../perf_throughput/目录。

最新新闻

日新闻

周新闻

月新闻