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(行切分)构成一对,中间的激活计算完全本地化。