ML Revisit: Infra


2026-09-05

训练就基本是跟着 (Hashimoto & Liang, 2026; Tazi et al., 2025) 过了一遍。

Collective Operations

每行一张 GPU · root = R0颜色跟随数据块,下标表示分片逐元素求和的图例Σ₀ = A₀ + B₀ + C₀ + D₀,表示同一分片位置的逐元素求和。A₀+B₀+C₀+D₀=Σ₀
Broadcast:整份复制R0 的完整张量 A、B、C、D 被复制到所有 rank,每个 rank 都得到完整的一份。输入输出R0ABCDR1R2R3ABCDABCDABCDABCD
Scatter:切开分发R0 持有 A、B、C、D 四块数据。分发后 R0 得到 A,R1 得到 B,R2 得到 C,R3 得到 D。输入输出R0ABCDR1R2R3ABCD
Gather:收集拼接四个 rank 分别提供 A、B、C、D,在 R0 按 rank 顺序拼接为四块数据。只在 R0 指定输出。输入输出R0AR1BR2CR3DABCD
Reduce:逐项相加四个 rank 的同形状数据 A、B、C、D 逐元素相加,R0 得到同形状的结果 Σ = A + B + C + D。只在 R0 指定输出。输入输出+R0AR1BR2CR3DΣ
All-gather:每处都拼接四个 rank 分别提供 A、B、C、D,每个 rank 都得到按 rank 顺序拼接的完整数据 A、B、C、D。输入输出R0AR1BR2CR3DABCDABCDABCDABCD
Reduce-scatter:相加后分片各 rank 先按相同分片位置逐元素相加。Σᵢ = Aᵢ + Bᵢ + Cᵢ + Dᵢ。R0 只得到 Σ₀,R1 只得到 Σ₁,R2 只得到 Σ₂,R3 只得到 Σ₃。输入输出+R0A₀A₁A₂A₃R1B₀B₁B₂B₃R2C₀C₁C₂C₃R3D₀D₁D₂D₃Σ₀Σ₁Σ₂Σ₃
All-reduce:每处都得总和各 rank 按相同分片位置逐元素相加,每个 rank 都得到完整结果 Σ₀、Σ₁、Σ₂、Σ₃。输入输出+R0A₀A₁A₂A₃R1B₀B₁B₂B₃R2C₀C₁C₂C₃R3D₀D₁D₂D₃Σ₀Σ₁Σ₂Σ₃Σ₀Σ₁Σ₂Σ₃Σ₀Σ₁Σ₂Σ₃Σ₀Σ₁Σ₂Σ₃
All-to-all:按列交换每个 rank 将第 i 个分片发给 Ri。Ri 按来源 rank 顺序收到 Aᵢ、Bᵢ、Cᵢ、Dᵢ。输入的每一列成为输出的一行,没有求和或复制。可用图下方的单选按钮高亮一个目标 rank。输入输出R0A₀A₁A₂A₃R1B₀B₁B₂B₃R2C₀C₁C₂C₃R3D₀D₁D₂D₃A₀B₀C₀D₀A₁B₁C₁D₁A₂B₂C₂D₂A₃B₃C₃D₃
追踪去向
All-reduce = Reduce-scatter + All-gather
All-reduce 的两步分解左侧每个 rank 有四个输入分片。Reduce-scatter 后,中间的每个 rank 持有一个求和分片 Σᵢ。再经 All-gather,右侧每个 rank 都持有完整结果 Σ₀、Σ₁、Σ₂、Σ₃,与 All-reduce 的输出相同。这是结果等价的分解,不限定实际通信算法。输入每处一片每处一整份R0A₀A₁A₂A₃R1B₀B₁B₂B₃R2C₀C₁C₂C₃R3D₀D₁D₂D₃R0Σ₀R1Σ₁R2Σ₂R3Σ₃R0Σ₀Σ₁Σ₂Σ₃R1Σ₀Σ₁Σ₂Σ₃R2Σ₀Σ₁Σ₂Σ₃R3Σ₀Σ₁Σ₂Σ₃Reduce-scatterAll-gather
更详细的代码可以看 CS336 · Lecture 7。归约以 SUM 为例;箭头表示数据关系。虚线框表示无指定输入或输出,并不表示清空原数据。

Tensor Parallel

对于 Column-wise TP:

\[ X \begin{bmatrix} W_1&W_2&\cdots&W_n \end{bmatrix} = \begin{bmatrix} XW_1 & XW_2 & \cdots & XW_n \end{bmatrix} \]

反向传播的时候,令 \(Y_i=XW_i\)

\[ \frac{\partial\mathcal{L}}{\partial W_i} = X^\top\frac{\partial\mathcal{L}}{\partial Y_i}, \quad \frac{\partial\mathcal{L}}{\partial X} = \sum_{i=1}^n\frac{\partial L}{\partial Y_i}W_i^\top \]

对于 Row-wise TP:

\[ \begin{bmatrix} X_1&X_2&\cdots&X_n \end{bmatrix} \begin{bmatrix} W_1\\W_2\\ \vdots\\W_n \end{bmatrix} = \sum_{i=1}^nX_iW_i \]

反向传播的时候,令求和结果为 \(Y\)

\[ \frac{\partial\mathcal{L}}{\partial W_i} = X_i^{\top}\frac{\partial\mathcal{L}}{\partial Y}, \quad \frac{\partial\mathcal{L}}{\partial X_i}=\frac{\partial\mathcal{L}}{\partial Y}W_i^{\top} \]

对于一个 MLP

\[ \mathrm{MLP}(x) = \sigma(XW_1)W_2 \]

我们可以将 \(W_1\) 做 Column-wise 拆分,\(W_2\) 做 Row-wise 拆分,这样子做完 \(XW_1\) 之后不需要做 All-Gather,反向传播过完 \(W_2\) 之后也不需要 All-Gather 完整梯度。

对于 attention,考虑到 Megatron (Shoeybi et al., 2020) 里没有考虑 number of attention heads 大于 TP 的情况(代码位置),我们也只讨论这一点

python
if self.num_attention_heads % self.tensor_model_parallel_size != 0:
    raise ValueError(
        f"num_attention_heads ({self.num_attention_heads}) must be a multiple of "
        f"tensor_model_parallel_size ({self.tensor_model_parallel_size})."
    )

对于 MHA 就非常好做了,对于

\[ O = \begin{bmatrix} O_1 & O_2 & \cdots & O_h \end{bmatrix}W_o\\ O_t = \mathrm{Attention}\left(XW_q^{(t)}, XW_k^{(t)}, XW_v^{(t)}\right) \]

我们将 \(W_o\) 进行 Row-wise 切分。然后每张卡单独计算一定量的 heads,这样子天然是一个 Column-wise 切分的状态。

References

Hashimoto, T., & Liang, P. (2026). CS336: Language Modeling from Scratch. Stanford University, Spring 2026. cs336.stanford.edu
Shoeybi, M., Patwary, M., Puri, R., LeGresley, P., Casper, J., & Catanzaro, B. (2020). Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism. arxiv.org
Tazi, N., Mom, F., Zhao, H., Nguyen, P., Mekkouri, M., von Werra, L., Wolf, T., & HuggingFace, O. (2025). The ultra-scale playbook: Training LLMs on GPU clusters. Hugging Face. huggingface.co

Cite this post

@misc{pu2026mlrevisitinfra,
  author = {Pu, Fanyi},
  title  = {ML Revisit: Infra},
  year   = {2026},
  month  = {9},
  url    = {https://pufanyi.com/blog/ml-revisit-infra}
}