跳转至

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 并行的列切分/行切分模式。