跳转至

parallel_linear.py — LLaMA 并行线性层

文件路径

verl/models/llama/megatron/layers/parallel_linear.py

文件概述

定义了三种特殊的并行线性层:QKV 合并投影、gate+up 合并投影、以及 Value 模型输出层。

关键代码讲解

1. QKVParallelLinear

将 Q、K、V 三个投影合并为一个 ColumnParallelLinear,减少通信次数。

class QKVParallelLinear(tensor_parallel.ColumnParallelLinear):
    def __init__(self, input_size, num_heads, num_key_value_heads, head_dim, **kwargs):
        self.q_output_size = num_heads * head_dim
        self.kv_output_size = num_key_value_heads * head_dim
        # 总输出 = Q + K + V
        output_size = (num_heads + 2 * num_key_value_heads) * head_dim
        super().__init__(input_size=input_size, output_size=output_size, ...)

TP 切分方式:ColumnParallelLinear 按 output_size 维度切分权重。每个 TP rank 持有 output_size // tp_size 的权重。由于 Q、K、V 是拼接在一起的,切分后每个 rank 自然得到了对应头数的 Q、K、V 分片。

2. MergedColumnParallelLinear

将 gate_proj 和 up_proj 合并为一个线性层:

class MergedColumnParallelLinear(tensor_parallel.ColumnParallelLinear):
    def __init__(self, input_size, gate_ouput_size, up_output_size, **kwargs):
        output_size = gate_ouput_size + up_output_size
        super().__init__(input_size=input_size, output_size=output_size, ...)

合并后在 MLP 中通过 split 分回 gate 和 up 两部分。

3. LinearForLastLayer

Value 模型的最后一层,支持序列并行:

class LinearForLastLayer(torch.nn.Linear):
    def __init__(self, input_size, output_size, *, config, bias=True):
        super().__init__(in_features=input_size, out_features=output_size, bias=bias)
        self.sequence_parallel = config.sequence_parallel
        if self.sequence_parallel:
            self.weight.sequence_parallel = True  # 标记为 SP 参数

    def forward(self, input_, weight=None, runtime_gather_output=None):
        logits = super().forward(input_)
        logits = logits.float()
        if self.sequence_parallel:
            # 从 SP 分片 gather 完整输出
            logits = tensor_parallel.gather_from_sequence_parallel_region(logits)
        return logits, None

这个类不做 TP 切分(输出维度只有 1),而是处理 SP 的 gather。

核心类列表

类名 父类 用途
QKVParallelLinear ColumnParallelLinear Q+K+V 合并,按 TP 列切分
MergedColumnParallelLinear ColumnParallelLinear gate+up 合并,按 TP 列切分
LinearForLastLayer nn.Linear Value 模型输出,SP gather

与其他模块的关系

  • QKVParallelLinear 被 parallel_attention.py 使用
  • MergedColumnParallelLinear 被 parallel_mlp.py 使用
  • LinearForLastLayer 被 modeling_llama_megatron.py 的 Value 模型变体使用

小结

这三个线性层封装了 LLaMA 中常见的权重合并模式。QKV 合并和 gate+up 合并是常见的优化技巧,可以减少通信次数。LinearForLastLayer 解决了 Value 模型在序列并行下的输出问题。