JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

把模型本身切開:張量、管線與序列平行

當連單一層的權重或激活值都放不下時,你不再複製模型,而是開始把它的運算切分到不同裝置上。

張量平行:把一個矩陣乘切到多張裝置

資料平行與分片都讓每一層的運算在單一裝置上保持完整張量平行 打破了這個假設:它把一層內部的權重矩陣切分,使矩陣乘法本身由數張裝置共同計算。對 Transformer 的前饋區塊 `Y = GeLU(X W1) W2`,你把 `W1` 按欄切、把 `W2` 按列切;每張裝置算出部分結果,最後一次 all-reduce 把它們加成正確的輸出。注意力區塊則沿著頭(heads)自然切分。這正是 Megatron-LM 的核心技巧,也是讓一個對單裝置太大的層,能跨在例如伺服器內八張 GPU 上運行的原因。

Y\,B=\begin{bmatrix}Y_1 & Y_2\end{bmatrix}\begin{bmatrix}B_1\\ B_2\end{bmatrix}=\underbrace{Y_1B_1+Y_2B_2}_{\text{all-reduce across devices}}

张量并行把权重矩阵切成若干片,每个设备只持有一片;各自的部分积 Y_iB_i 必须用一次位于关键路径上的 all-reduce 求和。

序列平行:分片剩下的激活值

張量平行切分了矩陣乘,但它們之間的運算——層正規化、dropout、殘差相加——仍是複製的,每張裝置都持有那些區域的完整激活值。在長序列下,這份複製的激活值就成了新瓶頸。序列平行 補上這個缺口,把那些激活值沿序列維度分片到同一個張量平行群組上,把區域邊界的 all-reduce 換成 all-gather 加 reduce-scatter 的組合:搬動同樣的資料量,卻全程保持記憶體分片。實務上它與張量平行幾乎能無成本地組合,對長上下文訓練不可或缺。

在 Transformer 块内部,张量并行切分注意力和前馈的矩阵乘法,但层归一化、dropout 和残差相加仍被复制——序列并行正是沿序列维度把这些剩余激活值切分开。

一个 Transformer 块,显示层归一化、注意力、残差相加和前馈网络。

管線平行:把層切成階段

張量平行切的是層之內管線平行 切的是層之間。你把例如第 1–8 層指派給階段 A、9–16 層給階段 B,依此沿著一串裝置排下去。一個批次向前流經各階段、梯度往回流,像一條生產線。危險在於管線氣泡:當階段 A 在算第一個微批次時,階段 B、C、D 還沒東西可做,只能閒置。解法是把批次切成許多 微批次 並讓它們同時在管線中流動,使管線一旦填滿,每個階段都忙碌。氣泡比例會隨著微批次數相對於階段數的增加而縮小。

# 1F1B schedule: interleave forward and backward to bound memory
for step in pipeline_schedule(num_microbatches, num_stages):
    if step.is_forward:
        act = stage.forward(recv_from_prev())
        send_to_next(act)
    else:                          # backward of an earlier microbatch
        grad = stage.backward(recv_from_next())
        send_to_prev(grad)
1F1B(一前向一反向)排程使每個階段只保留少量在用的激活值。
\text{bubble fraction}=\frac{p-1}{m+p-1}

流水线气泡:随着微批数量 m 相对于阶段数 p 的增大,'填充与排空'的空闲占比随之缩小。

重算:通用的記憶體折扣

無論你用哪種切分,激活值在大規模下都主宰記憶體,而永遠管用的槓桿就是 激活值重算(又稱梯度檢查點)。你不再為反向傳播儲存每個中間激活值,而是只存少數檢查點——通常是每個 Transformer 區塊的輸入——其餘的在反向傳播時即時重算。這筆交易很明確:用大約多一次前向的計算(若每個區塊都做檢查點,約 30% 的減速)換取激活記憶體下降一個數量級。選擇性重算只重算那些便宜但佔空間的運算,能以一小部分代價換得大半的節省。

\underbrace{O(N)}_{\text{store all activations}}\;\xrightarrow{\;\text{recompute}\;}\;\underbrace{O(\sqrt{N})}_{\text{store }\sqrt{N}\text{ checkpoints}}

梯度检查点用计算换显存:只保存 √N 个检查点而非全部 N 个激活值,把激活显存降到 O(√N),代价是多做一次前向传播。