跳转至

parallel_linear.py — Qwen2 并行线性层

文件路径

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

文件概述

定义 Qwen2 的并行线性层。与 LLaMA 版本类似,但只包含两个类(没有 LinearForLastLayer)。

关键代码讲解

QKVParallelLinear

class QKVParallelLinear(tensor_parallel.ColumnParallelLinear):
    def __init__(self, input_size, num_heads, num_key_value_heads, head_dim, **kwargs):
        output_size = (num_heads + 2 * num_key_value_heads) * head_dim
        super().__init__(input_size=input_size, output_size=output_size, ...)

MergedColumnParallelLinear

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, ...)

与 LLaMA 版本的代码完全相同。

与 LLaMA 版本的区别

  • 没有 LinearForLastLayer(Qwen2 的 Value 模型输出层直接使用 nn.Linear)

小结

Qwen2 的并行线性层与 LLaMA 版本功能相同,是 QKV 合并和 gate+up 合并的 TP 并行封装。