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 版本功能完全相同。