推理计算流梳理

推理计算流梳理

基本计算

Norm&Residual

采用 Pre-Norm,RMSNorm。 \[ U_\ell=\operatorname{RMSNorm}_{\ell,\mathrm{attn}}(X_\ell),\qquad R_\ell=X_\ell+\operatorname{Attention}_\ell(U_\ell) \]

\[ V_\ell=\operatorname{RMSNorm}_{\ell,\mathrm{ffn}}(R_\ell),\qquad X_{\ell+1}=R_\ell+\operatorname{FFN}_\ell(V_\ell) \]

\(\gamma_j\) 是可学习的量。 \[ \operatorname{RMSNorm}(x)_j=\gamma_j\frac{x_j}{\sqrt{\frac{1}{4096}\sum_{k=0}^{4095}x_k^2+10^{-5}}} \]

RoPE

对于一个向量,通常会将其中一部分标量两两绑定,对两两绑定的标量做 RoPE。

一个例子是对于 dim=192,约定将前 64 的位置按照 (i, i+32) 为一组绑定,对这些位置应用 RoPE 旋转。 \[ R(\theta)=\begin{bmatrix}\cos\theta&-\sin\theta\\\sin\theta&\cos\theta\end{bmatrix},\qquad\begin{bmatrix}u'\\v'\end{bmatrix}=R(\theta)\begin{bmatrix}u\\v\end{bmatrix} \]

通信原语

原语 输入与输出关系 典型使用位置
All-Reduce 各 rank 的相同形状张量逐元素归约,完整结果分发给各 rank 汇总 Row Parallel 的部分和
All-Gather 拼接各 rank 的分片,完整结果分发给各 rank 重建 token 或特征维度
Reduce-Scatter 先归约,再让每个 rank 只保留结果的一片 为下一阶段建立 token 分片布局
All-to-All 每个 rank 将不同数据发送给不同目标 rank MoE token dispatch/combine

张量并行和通信

设一批 token 的 hidden 为 \(X\in \mathbb{R}^{m\times d}\),线性层为 \(W\in \mathbb{R}^{d\times f}\),输出为 \(Y=XW\)。

通常用一个行向量代表一个 token 的信息,每一列是不同的特征量

列并行

拆权重的列(输出维),如果下一步需要完整的 \(Y\),可以做 All-Gather。 \[ W=[W_0\;W_1\;\cdots\;W_{p-1}],\qquad Y_i=XW_i \]

行并行

拆输入的列和输出的行,每个输出 \(X_iW_i\) 有完整形状但是只有和的一部分,后续一般要做和 Reduce 有关的操作以便处理。 \[ X=[X_0\;X_1\;\cdots\;X_{p-1}],\qquad W=\begin{bmatrix}W_0\\W_1\\\vdots\\W_{p-1}\end{bmatrix} \]

\[ Y=\sum_{i=0}^{p-1}X_iW_i \]

实际 MLP 计算策略

MiMo-V2.5 使用这样的门控 MLP。 \[ Y=\left(\operatorname{SiLU}(XW_{\mathrm{gate}})\odot XW_{\mathrm{up}}\right)W_{\mathrm{down}} \] 对于 \(W_\text{gate}, W_\text{up}\),采用列并行,\(\text{SiLU}\) 和 pointwise 的计算无需通信。

此时每个 rank 持有一份分片,恰好是行并行要求的 \(X\) 分片方式,对 \(W_\text{down}\) 做行并行,可以得到最终计算结果。

Attention 的切分

Attention 通常按照 heads 切分。

对于 GQA 的 attention,通常 Q 比 KV 会多,如果 world_size=8,kv_heads=4,通常需要将 KV 复制一份,按列硬拆会在 SPDA 过程引入额外的通信(需要获取别的 rank 上的 KV)。

正常 Attention 过程中,先复制分发 hidden stats,各个 Rank 计算自身的 QKV,独立完成自己负责的 attention,在输出投影 \(W_O\) 后可以得到部分和,此时做一次 All-Reduce。

专家并行

MoE Router 过程

使用一层 \(\text{sigmoid}\) 激活带偏置的神经网络对 Expert 做打分,取 Top8,并分配权重。注意偏置不参与权重分配打分。 \[ z=xW_r^{\mathsf T},\quad s=\operatorname{sigmoid}(z),\quad I=\operatorname{Top8}(s+b)\qquad w_i=\frac{s_i}{\sum_{j\in I}s_j}\;(i\in I) \]

EP 路径

实际的实现中,Attention 的 All-Reduce 变成了 Reduce-Scatter,每个 rank 都持有一定量 token 和 \(W_r\),计算之后,做 All2All dispatch 通信。

dispatch 后每个 rank 各自做 MLP 计算(无通信)。MLP 计算需要单个 rank 激活多个专家,可以使用 grouped GEMM 统一调度计算。

MLP 结束后,某种实现是在来源 rank 做加权求和得到最终状态。

MiMo-V2.5 大致计算路径

  • token_id 发到每个 rank,查表(不在表中的置零),All-Reduce。
  • Input Pre-Norm; attention; post-attention Norm; SiLU MLP