目录
[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 层,里面通常还会:
- 把
y按 token 量化成 int8(得到y_q+ scale) - 用 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),把下面三步合成一步:
- 对 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 串起来:融合后的整体流程
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 |