tensor_parallel.py — 张量并行¶
文件路径: verl/utils/megatron/tensor_parallel.py
文件概述¶
张量并行 (Tensor Parallelism) 工具,提供并行线性层的默认配置、参数分片信息查询,以及词表并行 (Vocab Parallel) 的 entropy 和 log-prob 计算。
核心功能¶
1. 并行配置¶
def get_default_kwargs_for_column_parallel_linear():
"""列并行线性层的默认参数"""
def get_default_kwargs_for_row_parallel_linear():
"""行并行线性层的默认参数"""
2. 词表并行 Entropy¶
class _VocabParallelEntropy(torch.autograd.Function):
@staticmethod
def forward(ctx, vocab_parallel_logits):
"""计算词表分片在各 TP rank 上的 entropy"""
# 1. 全局 max(all_reduce MAX)
# 2. 数值稳定的 softmax
# 3. 计算 entropy = log(sum_exp) - weighted_sum
当词表按 TP 分片时,计算 entropy 需要跨 rank 通信,因为每个 rank 只有词表的一部分。
3. 词表并行 Log-Prob¶
def vocab_parallel_log_probs_from_logits(logits, labels):
"""在 TP 分片的 logits 上计算 log probabilities"""
return -tensor_parallel.vocab_parallel_cross_entropy(vocab_parallel_logits=logits, target=labels)
核心函数/类列表¶
| 函数/类 | 说明 |
|---|---|
get_default_kwargs_for_column_parallel_linear() |
列并行配置 |
get_default_kwargs_for_row_parallel_linear() |
行并行配置 |
_VocabParallelEntropy |
并行 entropy 计算 |
vocab_parallel_log_probs_from_logits() |
并行 log-prob |
is_tensor_parallel_param() |
检查参数是否是 TP 参数 |
与其他模块的关系¶
- 被 Megatron Actor worker 计算 log-prob 时使用
- 依赖
megatron.core.parallel_state获取 TP 进程组
小结¶
在张量并行下正确计算词表相关量(entropy、log-prob)的核心工具。