目录

1 bug 背景

2 SwiGLU解释

3 原来的MLP流程

3.1 输入

3.2 第 1 步:gate_up_proj

3.3 第 2 步:SiluAndMul

3.4 第 3 步:down_proj

3.5 串起来整体流程

4 算子融合之后的 MLP 流程

4.1 融合在融什么

4.2 融合后逐步过程

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

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 串起来整体流程

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),把下面三步合成一步:

  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 串起来:融合后的整体流程

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 和原来比,差在哪

原来 融合后

激活

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

Logo

免费领 150 小时云算力,进群参与显卡、AI PC 幸运抽奖

更多推荐