parallel_rmsnorm.py — LLaMA 并行 RMSNorm¶
文件路径¶
verl/models/llama/megatron/layers/parallel_rmsnorm.py
文件概述¶
实现 LLaMA 的 RMSNorm(均方根归一化),使用 NVIDIA Apex 的融合算子加速,支持序列并行。
关键代码讲解¶
class ParallelLlamaRMSNorm(nn.Module):
def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig):
normalized_shape = (config.hidden_size,)
self.normalized_shape = torch.Size(normalized_shape)
self.weight = nn.Parameter(torch.ones(self.normalized_shape))
self.variance_epsilon = config.rms_norm_eps
# 序列并行:标记 weight 为 SP 参数
if megatron_config.sequence_parallel:
sp_utils.mark_parameter_as_sequence_parallel(self.weight)
def forward(self, hidden_states):
from apex.normalization.fused_layer_norm import fused_rms_norm_affine
return fused_rms_norm_affine(
input=hidden_states,
weight=self.weight,
normalized_shape=self.normalized_shape,
eps=self.variance_epsilon,
memory_efficient=True, # 节省反向传播内存
)
RMSNorm 公式¶
\[\text{RMSNorm}(x) = \frac{x}{\sqrt{\text{mean}(x^2) + \epsilon}} \cdot \gamma\]
相比 LayerNorm,RMSNorm 去掉了均值中心化,只做方差归一化,计算更简单。
序列并行兼容性¶
mark_parameter_as_sequence_parallel(self.weight)标记 weight 参数,使得在优化器中正确处理梯度同步- RMSNorm 在序列维度上是逐元素操作,可以直接在 SP 分片上计算,不需要额外通信
与其他模块的关系¶
- 被
parallel_decoder.py使用(input_layernorm和post_attention_layernorm) - 被
modeling_llama_megatron.py使用(最后的norm)
小结¶
ParallelLlamaRMSNorm 是一个简洁的 RMSNorm 实现,依赖 Apex 的融合算子提供高效计算和内存优化。序列并行通过参数标记实现,不需要修改前向计算逻辑。