跳转至

parallel_rmsnorm.py — Qwen2 并行 RMSNorm

文件路径

verl/models/qwen2/megatron/layers/parallel_rmsnorm.py

文件概述

Qwen2 的 RMSNorm 实现。与 LLaMA 版本几乎相同,唯一区别是直接在顶层导入 fused_rms_norm_affine(LLaMA 版在 forward 内延迟导入)。

关键代码讲解

from apex.normalization.fused_layer_norm import fused_rms_norm_affine

class ParallelQwen2RMSNorm(nn.Module):
    def __init__(self, config: Qwen2Config, megatron_config):
        normalized_shape = (config.hidden_size,)
        self.weight = nn.Parameter(torch.ones(normalized_shape))
        self.variance_epsilon = config.rms_norm_eps
        if megatron_config.sequence_parallel:
            sp_utils.mark_parameter_as_sequence_parallel(self.weight)

    def forward(self, hidden_states):
        return fused_rms_norm_affine(
            input=hidden_states, weight=self.weight,
            normalized_shape=self.normalized_shape,
            eps=self.variance_epsilon, memory_efficient=True,
        )

小结

使用 Apex 融合算子的 RMSNorm,支持序列并行参数标记。与 LLaMA 版本功能完全相同。