跳转至

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 的融合算子提供高效计算和内存优化。序列并行通过参数标记实现,不需要修改前向计算逻辑。