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 并行封装。