跳转至

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)的核心工具。