parallel_mlp.py — Qwen2 并行 MLP¶
文件路径¶
verl/models/qwen2/megatron/layers/parallel_mlp.py
文件概述¶
Qwen2 的 SwiGLU MLP 实现,与 LLaMA 版本代码完全相同。
关键代码讲解¶
class ParallelQwen2MLP(nn.Module):
def __init__(self, config, megatron_config=None):
tp_size = mpu.get_tensor_model_parallel_world_size()
# gate + up 合并
self.gate_up_proj = MergedColumnParallelLinear(
input_size=self.hidden_size,
gate_ouput_size=self.intermediate_size,
up_output_size=self.intermediate_size,
bias=False, gather_output=False, ...
)
self.gate_size = self.intermediate_size // tp_size
# down_proj
self.down_proj = tensor_parallel.RowParallelLinear(
input_size=self.intermediate_size,
output_size=self.hidden_size,
bias=False, input_is_parallel=True, ...
)
self.act_fn = ACT2FN[config.hidden_act]
def forward(self, x):
gate_up = self.gate_up_proj(x)[0]
gate, up = gate_up.split(self.gate_size, dim=-1)
return self.down_proj(self.act_fn(gate) * up)[0]
小结¶
与 LLaMA 的 ParallelLlamaMLP 完全对称,使用 SwiGLU 激活和 TP 并行的列切分/行切分模式。