大模型推理引擎vLLM(30):由一个GLM5 bug,整理MLP中的SwiGLU、算子融合、量化相关问题

目录

[1 bug 背景](#1 bug 背景)

[2 SwiGLU解释](#2 SwiGLU解释)

[3 原来的MLP流程](#3 原来的MLP流程)

[3.1 输入](#3.1 输入)

[3.2 第 1 步:gate_up_proj](#3.2 第 1 步:gate_up_proj)

[3.3 第 2 步:SiluAndMul](#3.3 第 2 步:SiluAndMul)

[3.4 第 3 步:down_proj](#3.4 第 3 步:down_proj)

[3.5 串起来整体流程](#3.5 串起来整体流程)

[4 算子融合之后的 MLP 流程](#4 算子融合之后的 MLP 流程)

[4.1 融合在融什么](#4.1 融合在融什么)

[4.2 融合后逐步过程](#4.2 融合后逐步过程)

[4.3 串起来:融合后的整体流程](#4.3 串起来:融合后的整体流程)

[4.4 和原来比,差在哪](#4.4 和原来比,差在哪)


门控激活结构,逐个元素相乘

GLU(Gated Linear Unit):用一路当"门",控制另一路过多少

Swi:门控那一路用 SiLU / Swish(不是 sigmoid)

量化,矩阵乘,反量化

1 bug 背景

hipblaslt_w8a8_gemm → assert a.shape-1 == b.shape-1

这其实是因为vllm021上面加了一个算子融合,然后导致我之前的代码没法跑了,在解决这个bug的时候,顺便梳理了 一下mlp这块的东西。

2 SwiGLU解释

SwiGLU 是一种 MLP/FFN 里的激活结构,名字可以拆开记:

  • GLU(Gated Linear Unit):用一路当"门",控制另一路过多少
  • Swi:门控那一路用 SiLU / Swish(不是 sigmoid)

公式就是:

y = silu(W_gate · x) ⊙ (W_up · x) #逐元素相乘,不是矩阵乘法。

out = W_down · y

是逐元素相乘。

和老式 silu(W · x) 比:多了一路线性(gate),用激活后的 gate 去"开关" up 的信息。很多现代 LLM(Llama、DeepSeek、GLM 等)的 FFN 都用这类结构;你们代码里的 SiluAndMul / gate_up_proj 就是在实现它。

3 原来的MLP流程

3.1 输入

  • x:形状大概是 [token数 m, hidden_size]

3.2 第 1 步:gate_up_proj

  • 一个合并的线性层,一次算出 gate 和 up 两路,拼在一起。
  • 输出 gate_up[m, 2H]H = intermediate_size
  • 前半是 gate,后半是 up。

3.3 第 2 步:SiluAndMul

  • gate_up 拆成:
    • gate = gate_up[..., :H]
    • up = gate_up[..., H:]
  • 做:y = silu(gate) * up
  • 得到 y[m, H](还是 bf16/fp16)

这一步把中间宽度从 2H 收成 H,后面 down 才能接上。


3.4 第 3 步:down_proj

  • 线性层:H → hidden_size
  • 因为是 W8A8/W4A8 的 dense 层,里面通常还会:
    1. y 按 token 量化成 int8(得到 y_q + scale)
    2. 用 int8 GEMM:y_q × weight → 再乘 scale,得到 bf16/fp16 输出
  • 输出:[m, hidden_size]
  • 若 TP>1,后面还可能 all-reduce。

3.5 串起来整体流程

xm, hidden ← bf16 隐状态

gate_up_proj

│ ① 把激活 x(bf16)量化成 int8 + scale

│ ② 和 int8 权重做 GEMM(累加多为更高精度,如 int32)

│ ③ 再按 scale 反量化 → 输出仍是 bf16

gate_upm, 2H ← bf16(前 H 为 gate,后 H 为 up)

SiluAndMul ← y = silu(gate) * up(仍在 bf16 上做)

ym, H ← bf16

down_proj

│ ① 把激活 y(bf16)量化成 int8 + scale

│ ② 和 int8 权重做 GEMM

│ ③ 再按 scale 反量化 → 输出仍是 bf16

outm, hidden ← bf16

4 算子融合之后的 MLP 流程

4.1 融合在融什么

对照第 3 节「原来」的路径,down_proj 内部本来要先把 bf16 的 y 再量化成 int8,才能做 W8A8 GEMM。

于是中间会出现:

SiluAndMul → 写出 ym,H(bf16)→ down_proj 再读入 y → quant(y) → GEMM

021 增加的融合算子 fuse_silu_mul_quant(代码里常包成 FusedSiluAndMulAndQuant),把下面三步合成一步:

  1. 对 gate 做 SiLU
  2. 与 up 逐元素相乘
  3. per-token int8 量化

也就是:原来的 SiluAndMul + down_proj 前那一次激活量化。

融合后直接得到:

  • xq:int8,形状 [m, H]
  • xs:scale,形状大致 [m, 1]

不再先落一份 bf16 的 y,再在 down_proj 里重新 quant。

环境变量上通常由 VLLM_HCU_USE_FUSED_SILU_MUL_QUANT(以及 VLLM_HCU_USE_CUSTOM_OPS)控制,默认往往是开的。

4.2 融合后逐步过程

输入不变: x[m, hidden],仍是 bf16 隐状态。

第 1 步:gate_up_proj(与原来类似)

  • 内部仍可:量化 x → int8 GEMM → 反量化
  • 输出:gate_up[m, 2H],bf16

第 2 步:fuse_silu_mul_quant(替换原来的 SiluAndMul

  • 输入:gate_up[m, 2H](bf16)
  • 内部:silu(gate) * up,并立刻量化
  • 输出:(xq, xs),其中 xq 已是 [m, H] 的 int8

注意:这里没有再产出给外面用的 bf16 y

第 3 步:down_proj(接口变了)

设计意图是:

  • 激活侧 不再 对 bf16 输入做 quant(y)
  • 直接使用外面传来的预量化结果 (xq, xs)
  • 只做:xq × int8 权重 → 再按 scale 反量化 → bf16 输出

代码形态大致是:

gate_up, _ = self.gate_up_proj(x)

xq, xs = self.act_fn(gate_up, quant_dtype=...) # FusedSiluAndMulAndQuant

out, _ = self.down_proj(gate_up, x_and_scale_quanted=(xq, xs))

这里有个容易误解的点:down_proj 的第一个参数仍写着 gate_up

按融合路径的设计,gate_up 只是占位(Linear 接口需要一个 input_);真正参与 GEMM 的应是 x_and_scale_quanted=(xq, xs)

xq 的最后一维是 H,正好对上 down_proj 权重的 K 维。

4.3 串起来:融合后的整体流程

xm, hidden ← bf16 隐状态

gate_up_proj

│ ① quant(x) → int8 + scale

│ ② int8 × int8 GEMM

│ ③ dequant → bf16

gate_upm, 2H ← bf16

fuse_silu_mul_quant ← Silu + Mul + 激活量化(三合一)

├─► xqm, H ← int8

└─► xsm, 1 ← scale

down_proj(应走预量化输入)

│ ① 直接用 (xq, xs),不再 quant(bf16 y)

│ ② xq × int8 权重 GEMM

│ ③ dequant → bf16

outm, hidden ← bf16

4.4 和原来比,差在哪

| 原来 | 融合后 |
|--------------------|-------------------------|------------------------------------------|
| 激活 | SiluAndMul → bf16 y | fuse_silu_mul_quant → int8 xq + xs |
| down_proj 输入 | bf16 y,内部再 quant | 预量化 (xq, xs),内部跳过 quant |
| 中间是否写回 bf16 y | 是 | 否(这是融合省的一步) |
| 对 LinearMethod 的要求 | 自己 quant 激活即可 | 必须会接 x_and_scale_quanted |

相关推荐
花无缺pize1 天前
vLLM框架:LLM推理的高效机制
服务器·人工智能·vllm
Briwisdom2 天前
MoE 推理优化实战——从“瓶颈罗列“到“性能调优“
gemm·vllm·moe·decode·prefill
Briwisdom3 天前
LLM 推理引擎三强争霸——vLLM vs SGLang vs TensorRT-LLM
tensorrt·vllm·推理引擎·sglang
西柚小萌新3 天前
【大模型:部署】--使用VLLM架构部署本地大模型
vllm
不断学习加努力4 天前
使用rviz2进行可视化时,cpu资源占用过高的bug
bug
happyness444 天前
如何利用 AI 自动编写单元测试(Unit Test)来捕捉隐藏的边缘情况 Bug?
人工智能·单元测试·bug
白驹_过隙4 天前
【大模型OCR落地终极排坑:OvisOCR2+vLLM从报错到批量稳定部署全过程】
人工智能·ocr·vllm
深念Y4 天前
Windows幽灵端口占用:HNS如何无声偷走你的端口
windows·python·bug·环境·端口·特权
一个王同学5 天前
从零到一 | CV转多模态大模型 | week19 | 基于 FastAPI 和 vLLM 的多模态大模型部署
人工智能·深度学习·计算机视觉·fastapi·改行学it·vllm