大模型推理引擎vLLM(30):由一个GLM5 bug,整理MLP中的SwiGLU、算子融合、量化相关问题
目录
门控激活结构,逐个元素相乘
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 层,里面通常还会:
- 把
y按 token 量化成 int8(得到y_q+ scale) - 用 int8 GEMM:
y_q×weight→ 再乘 scale,得到 bf16/fp16 输出
- 把
- 输出:
[m, hidden_size] - 若 TP>1,后面还可能 all-reduce。
3.5 串起来整体流程
x[m, hidden] ← bf16 隐状态
│
▼
gate_up_proj
│ ① 把激活 x(bf16)量化成 int8 + scale
│ ② 和 int8 权重做 GEMM(累加多为更高精度,如 int32)
│ ③ 再按 scale 反量化 → 输出仍是 bf16
▼
gate_up[m, 2H] ← bf16(前 H 为 gate,后 H 为 up)
│
▼
SiluAndMul ← y = silu(gate) * up(仍在 bf16 上做)
│
▼
y[m, H] ← bf16
│
▼
down_proj
│ ① 把激活 y(bf16)量化成 int8 + scale
│ ② 和 int8 权重做 GEMM
│ ③ 再按 scale 反量化 → 输出仍是 bf16
▼
out[m, hidden] ← bf16
4 算子融合之后的 MLP 流程
4.1 融合在融什么
对照第 3 节「原来」的路径,down_proj 内部本来要先把 bf16 的 y 再量化成 int8,才能做 W8A8 GEMM。
于是中间会出现:
SiluAndMul → 写出 y[m,H](bf16)→ down_proj 再读入 y → quant(y) → GEMM
021 增加的融合算子 fuse_silu_mul_quant(代码里常包成 FusedSiluAndMulAndQuant),把下面三步合成一步:
- 对 gate 做 SiLU
- 与 up 逐元素相乘
- 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 串起来:融合后的整体流程
x[m, hidden] ← bf16 隐状态
│
▼
gate_up_proj
│ ① quant(x) → int8 + scale
│ ② int8 × int8 GEMM
│ ③ dequant → bf16
▼
gate_up[m, 2H] ← bf16
│
▼
fuse_silu_mul_quant ← Silu + Mul + 激活量化(三合一)
│
├─► xq[m, H] ← int8
└─► xs[m, 1] ← scale
│
▼
down_proj(应走预量化输入)
│ ① 直接用 (xq, xs),不再 quant(bf16 y)
│ ② xq × int8 权重 GEMM
│ ③ dequant → bf16
▼
out[m, hidden] ← bf16
4.4 和原来比,差在哪
| 原来 | 融合后 | |
|---|---|---|
|
激活 |
|
|
|
|
bf16 |
预量化 |
|
中间是否写回 bf16 |
是 |
否(这是融合省的一步) |
|
对 LinearMethod 的要求 |
自己 quant 激活即可 |
必须会接 |
更多推荐




所有评论(0)