flowchart LR
classDef input fill:#495057,stroke:#212529,color:#ffffff,stroke-width:1px;
classDef rank0 fill:#1971c2,stroke:#0b4a85,color:#ffffff,stroke-width:1px;
classDef rank1 fill:#0b7285,stroke:#064653,color:#ffffff,stroke-width:1px;
classDef comm fill:#5f3dc4,stroke:#3b2585,color:#ffffff,stroke-width:2px;
classDef model fill:#2b8a3e,stroke:#176329,color:#ffffff,stroke-width:1px;
I["token ids"]:::input
E0["Rank 0 Embedding<br/>使用 W0<br/>vocab [0,V/2)"]:::rank0
E1["Rank 1 Embedding<br/>使用 W1<br/>vocab [V/2,V)"]:::rank1
C["all-reduce<br/>或 reduce-scatter"]:::comm
T["Transformer blocks"]:::model
L0["Rank 0 LM Head<br/>复用 W0<br/>local logits [0,V/2)"]:::rank0
L1["Rank 1 LM Head<br/>复用 W1<br/>local logits [V/2,V)"]:::rank1
CE["Vocab Parallel<br/>Cross Entropy"]:::comm
I --> E0
I --> E1
E0 --> C
E1 --> C
C --> T
T --> L0
T --> L1
L0 --> CE
L1 --> CE
E0 -. "同一份 W0 shard" .-> L0
E1 -. "同一份 W1 shard" .-> L1
5 Vocab Parallelism
5.1 为什么还需要 Vocab Parallelism
TP 已经切分了 Attention 和 MLP 中的大矩阵,SP 又切分了两个 TP 区域之间的 activation。接下来检查模型边界处还没有被这两种策略覆盖的部分。
此时,仍未切分的主要大张量是 embedding 和 LM head:
embedding weight: [V,H]
LM head weight: [V,H]它们在每个 rank 上都是完整副本,因为 TP 切的是 Transformer block 内的 hidden/intermediate dimension,SP 切的是 sequence dimension,两者都没有触及 vocab dimension。
MiniMind 将 embedding 与 LM head 绑定为同一个参数,因此不会实际保存两份 weight;但这份共享的 [V,H] 参数及其 gradient、optimizer states 仍在所有 rank 上完整复制。LM head 产生的 logits [S,B,V] 也不会随 TP size 缩小。词表 \(V\) 越大,这部分显存就越难忽略。
既然这些 tensor 都沿 vocab dimension 增长,是否可以让不同 rank 各自负责一段词表?Vocab Parallelism(VP)正是沿这一维继续切分:
embedding weight per rank: [V/T,H]
LM head weight per rank: [V/T,H]
local logits: [S,B,V/T]它要求:
V % T == 0对应的配置是:
vocab_parallel: bool = False5.2 VP 的模型边界
Embedding 负责把 token id 映射到 hidden states,LM head 再把 hidden states 映射回 logits。绑定权重启用时,它们共享同一个 [V,H] 参数,因此较合理的设计是让 LM head 也沿相同的 vocab range 切分。
Rank 0 在模型两端都使用 \(W_0\),Rank 1 都使用 \(W_1\),从单个 rank 的角度看,现在的行为与不切分 Embedding/LM head 完全一致,两个 layer 在底层共享同一份 Parameter。若不切分 LM head,则每次更新后需要在所有 rank 上同步 LM head 的完整参数,显然不如直接切分 LM head。
5.2.1 Vocab Parallel Embedding
Rank \(t\) 只负责区间:
\[ \left[t\frac{V}{T},(t+1)\frac{V}{T}\right). \]
先看一个具体例子。设 \(V=8\)、TP size \(=2\),输入的两个 token id 是 [1,6]:
| token 1 | token 6 | |
|---|---|---|
Rank 0,负责 [0,4) |
\(E_1\) | \(0\) |
Rank 1,负责 [4,8) |
\(0\) | \(E_6\) |
| SUM | \(E_1\) | \(E_6\) |
Rank 0 只认识 token 1,Rank 1 只认识 token 6;不属于本地 vocab shard 的位置输出零。两个 local embedding 相加后,正好恢复 Dense Embedding 的结果:
\[ [E_1,0]+[0,E_6]=[E_1,E_6]. \]
本地 lookup 的关键逻辑是:
mask = (token < self.vocab_start) | (token >= self.vocab_end)
local_token = token - vocab_start
local_token[mask] = 0
output = local_embedding(local_token)
output[mask] = 0对于任意 token,只有一个 rank 产生非零 embedding。因此,根据 SP 是否开启,可以用不同的方法得到当前 rank 需要的 embedding:
SP off: all-reduce local embeddings
SP on: reduce-scatter local embeddings5.2.2 Vocab Parallel LM Head
LM head 要在输出维度 V 上切成 V/T,本质上是 Column Parallel Linear:
input: [S,B,H]
LM head weight shape: [V/T,H]
output: [S,B,V/T]MiniMind 绑定 embedding 和 LM head 权重。开启 VP 后,两者必须引用同一个 vocab shard,这个很简单,因为都在同一个 rank 上,设计时 VocabParallelEmbedding 和 ColumnParallelLinear 都是模仿普通的 nn.Embedding 和 nn.Linear,因此绑定权重和一般的模块完全一致:
if self.config.tie_word_embeddings:
self.model.embed_tokens.weight = self.lm_head.weight5.3 Vocab Parallel Cross Entropy
5.3.1 Forward:不恢复完整 Logits
当 LM head 被切分后,每个 rank 仅持有形状为 [B, S, V/T] 的 local logits。最直接的实现方式是通过全收集(all-gather)恢复完整张量:
local logits [B, S, V/T]
-> all-gather
-> full logits [B, S, V]
-> standard cross entropy但这种方式会在每个 rank 上重新构造 [B, S, V],导致此前通过词表并行(VP)省下的完整 logits 显存重新被占用,且通信量也会随词表大小 \(V\) 线性增长。因此我们需要思考:能否不恢复完整 logits,直接利用各 rank 的 vocab shards 算出与 Dense Cross Entropy 完全一致的结果?
设某一 token 位置上的完整 logits 为 \(x\),目标类别为 \(k\)。Cross Entropy 可以展开为:
\[\operatorname{CE}(x,k) = -x_k + m + \log\sum_{i=1}^{V}e^{x_i-m}\]
其中 \(m = \max_i x_i\) 用于保证数值稳定性。
虽然每个 rank 仅持有局部词表的 logits,但只需通过三次形状为 [B, S] 的通信,即可收集到所需的变量:
- 目标类别的 target logit (\(x_k\)):正确词的 logit \(x_k\) 只会出现在负责该词表 shard 的单个 rank 上。每个 rank 先检查每个 token 的目标类别是否落在自己的 vocab shard 内;若是,则将对应的 logit 写入本地 target logit 张量,否则填零。随后通过
all-reduce(SUM 算子)求和,所有 rank 便都能获得准确的 \(x_k\)。 - 全局最大值 (\(m\)):每个 rank 先独立计算本地 logits 沿着 vocab 维度的最大值,再通过
all-reduce(MAX 算子)规约,得到全局最大值 \(m\)。 - 全局 sum-exp 项:每个 rank 利用全局最大值 \(m\) 计算本地的 \(\sum \exp(x_{\text{local}} - m)\),再通过
all-reduce(SUM 算子)求和,即可得到全局的归一化分母。
借助这三次仅传输 [B, S] 张量的通信,我们以极低的显存和通信开销精确计算出了 Cross Entropy,避免了传输 [B, S, V] 完整 logits 的高昂代价。
5.3.2 Backward:本地无分支更新
回顾 Cross Entropy 的定义:
\[\operatorname{CE}(x,k) = -x_k + \log\sum_{i=1}^{V}e^{x_i}\]
损失函数对第 \(j\) 个 logit \(x_j\) 的偏导数为:
\[\frac{\partial \operatorname{CE}}{\partial x_j} = \frac{e^{x_j}}{\sum_{i=1}^{V}e^{x_i}} - \mathbb{I}_{\{k=j\}} = p_j - y_j\]
即 Cross Entropy 对 logits 的梯度为:
\[\frac{\partial \operatorname{CE}}{\partial x_i} = p_i - y_i\]
在反向传播过程中:
- 概率 \(p_i\) 可以通过本地 rank 的 logits 与 forward 阶段通信获取的全局 sum-exp 直接计算得出。
- one-hot \(y_i\) 可以通过掩码逻辑判断 target 是否属于当前 rank:若在本地,则在 target 位置减 1;若不在,则减 0。
因此,backward 阶段的梯度计算完全不需要额外的跨卡通信,各 rank 均可在本地独立完成。
Vocab Parallel Cross Entropy 的 forward/backward 关键逻辑
# forward 阶段
target_mask = (target < vocab_start) | (target >= vocab_end) # [B, S],True 表示不在当前 rank
masked_target = target.clone() - vocab_start
masked_target[target_mask] = 0 # [B, S],属于本地的映射为 local id,不属于的置 0 防越界
# forward 中其他计算 ...
prob = exp_logits / exp_logits_sum # [B, S, v]
ctx.save_for_backward(prob, masked_target, target_mask)
# backward 阶段
prob, masked_target, target_mask = ctx.saved_tensors
prob_update = 1.0 - target_mask.float() # 在当前 rank 为 1.0,不在为 0.0
grad_input = prob.view(-1, v) # [B * S, v]
arange_1d = torch.arange(B * S, device=grad_input.device)
grad_input[arange_1d, masked_target.view(-1)] -= prob_update.view(-1)
# 逻辑说明:
# 1. 若 target 在本地:masked_target 为有效的 local id,prob_update 为 1.0,对应位置执行 p - 1.0。
# 2. 若 target 不在本地:masked_target 虽被置 0,但 prob_update 为 0.0,0 号位置执行 p - 0.0,保持不变。5.3.3 ignore_index
对于设置了 ignore_index 的位置,其对应的 loss 和 logits gradient 均必须为零。最终 loss 的归一化计算公式为:
\[L = \frac{\sum_{n \in \text{valid}} \ell_n}{\max(1, N_{\text{valid}})}\]
由于 TP 各个 rank 处理的是完全同一批序列 token,各 rank 接收到的 ignore_index 掩码与有效 token 数量 \(N_{\text{valid}}\) 完全一致,因此 loss 缩放与掩码操作均可在本地直接计算,无需额外的同步通信。
当前实现位置:
VocabParallelEmbedding:
model/tensor_parallel_layers.py
VocabParallelCrossEntropy:
model/tensor_parallel_layers.py5.4 收益验证
前面的推导说明,VP 直接作用于随词表大小 \(V\) 增长的 embedding、LM head、logits 和交叉熵中间状态。为了观察这部分收益如何随 \(V\) 变化,固定 8 层 Transformer 和 2048 sequence length,只扫描 vocab size:
Vocab size benchmark 命令
for v in 6400 12800 25600 51200; do
CUDA_DEVICE_MAX_CONNECTIONS=1 python -m evaluation.parallel benchmark \
--kind memory \
--tp_size 2 \
--pp_size 1 \
-- \
--modes ddp tp tp_vp \
--batch_policy fixed_per_rank \
--num_hidden_layers 8 \
--seq_len 2048 \
--hidden_size 768 \
--num_attention_heads 8 \
--num_key_value_heads 4 \
--vocab_size "$v" \
--micro_batch_size 4 \
--num_microbatches 1 \
--dtype bfloat16 \
--warmup_iters 10 \
--benchmark_iters 5 \
--output_csv "results/tp/vp_vocab_${v}.csv" \
--output_plot "results/tp/vp_vocab_${v}.png"
done峰值 allocated memory 如下:
可以看到 vocab size 每次翻倍时,VP 节省的绝对显存也几乎翻倍:约 227、455、910 和 1827 MiB。Transformer block 的其他显存保持不变时,VP 的相对收益也随 \(V\) 增大,从 8.2% 上升到 30.6%。
对应的 end-to-end step time 为:
VP 确实引入了额外通信:Vocab Parallel Embedding 需要合并各 rank 的局部输出,Parallel Cross Entropy 也要同步 target logit、全局最大值和 sum-exp。但是这些 collective 传输的是 [B,S,H] 的 embedding 输出或 Cross Entropy 中 [B,S] 的统计量,通信 tensor 的 shape 并不随 vocab size 增长。
与此同时,每个 rank 的 LM head 输出宽度从 \(V\) 降为 \(V/T\),本地 logits 以及交叉熵中的指数、归约和梯度计算也只覆盖自己的 vocab shard,这部分计算和显存访问会随 \(V\) 增长,因此切分后减少的计算也随 \(V\) 增长。在较小词表上,新增通信带来约 2%~3% 的时间开销。词表增大后,本地计算减少带来的收益逐渐抵消通信,实测 step time 甚至略有下降。这里的几百分点仍可能受到 kernel 和短时 benchmark 波动影响,但整体趋势说明 VP 并不是单纯用通信换显存,它同时切小了词表维度上的计算。