跳转至

parallel_mlp.py — LLaMA 并行 MLP

文件路径

verl/models/llama/megatron/layers/parallel_mlp.py

文件概述

实现 LLaMA 的 SwiGLU MLP,使用 Megatron 的 TP 并行线性层。

关键代码讲解

class ParallelLlamaMLP(nn.Module):
    def __init__(self, config, megatron_config=None):
        tp_size = mpu.get_tensor_model_parallel_world_size()

        # gate_proj + up_proj 合并为一个 ColumnParallelLinear
        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,  # 输出保持 TP 分片状态
            **column_kwargs,
        )
        # 每个 TP rank 上的 gate 大小
        self.gate_size = self.intermediate_size // tp_size

        # down_proj 使用 RowParallelLinear(输入已按 TP 切分)
        self.down_proj = tensor_parallel.RowParallelLinear(
            input_size=self.intermediate_size,
            output_size=self.hidden_size,
            bias=False,
            input_is_parallel=True,  # 输入已经是 TP 分片
            **row_kwargs,
        )

        self.act_fn = ACT2FN[config.hidden_act]  # SiLU

    def forward(self, x):
        # 1. gate + up 合并计算
        gate_up = self.gate_up_proj(x)[0]  # (seq, 1, 2*intermediate//tp)
        gate, up = gate_up.split(self.gate_size, dim=-1)

        # 2. SwiGLU: act(gate) * up
        # 3. down_proj 内部做 reduce-scatter
        return self.down_proj(self.act_fn(gate) * up)[0]

SwiGLU MLP 的 TP 并行模式

输入 x: (seq, hidden_size)
    |
  [ColumnParallel: all-gather -> gate_up matmul]
    |
  gate_up: (seq, 2 * intermediate_size // tp)
    |
  [split -> SiLU(gate) * up]
    |
  intermediate: (seq, intermediate_size // tp)
    |
  [RowParallel: matmul -> reduce-scatter]
    |
输出: (seq // sp, hidden_size)

关键点: - gate_up_proj 的 gather_output=False:输出保持 TP 分片状态,无需 all-gather - down_proj 的 input_is_parallel=True:告知输入已经是 TP 分片 - SwiGLU 激活在 TP 分片上局部计算,无需通信

与其他模块的关系

  • 被 parallel_decoder.py 的解码层使用
  • 使用 parallel_linear.py 的 MergedColumnParallelLinear

小结

LLaMA MLP 使用 SwiGLU 激活(SiLU(gate) * up),通过合并 gate 和 up 投影减少通信。TP 并行下,gate_up_proj(列切分)和 down_proj(行切分)构成一对,中间的激活计算完全本地化。