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 模型在序列并行下的输出问题。