modeling_llama_megatron.py — LLaMA Megatron 主模型¶
文件路径¶
verl/models/llama/megatron/modeling_llama_megatron.py
文件概述¶
这是 LLaMA 模型在 Megatron 框架下的完整实现。它定义了 6 个模型类,覆盖了三种使用场景:
1. 带 padding 的基础版本:ParallelLlamaModel + ParallelLlamaForCausalLM
2. 去 padding (RmPad) 版本:ParallelLlamaModelRmPad + ParallelLlamaForCausalLMRmPad + ParallelLlamaForValueRmPad
3. 去 padding + 流水线并行 (PP) 版本:ParallelLlamaModelRmPadPP + ParallelLlamaForCausalLMRmPadPP + ParallelLlamaForValueRmPadPP
关键代码讲解¶
1. 基础模型 ParallelLlamaModel¶
class ParallelLlamaModel(nn.Module):
def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig):
# 使用 Megatron 的 VocabParallelEmbedding(词表按 TP 切分)
self.embed_tokens = tensor_parallel.VocabParallelEmbedding(
num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, **embedding_kwargs
)
# Transformer 解码层
self.layers = nn.ModuleList(
[ParallelLlamaDecoderLayer(config, megatron_config) for _ in range(config.num_hidden_layers)]
)
self.norm = ParallelLlamaRMSNorm(config, megatron_config)
这是标准的 Transformer 解码器结构,使用了 Megatron 的 TP 并行嵌入。
2. 因果语言模型 ParallelLlamaForCausalLM¶
class ParallelLlamaForCausalLM(nn.Module):
def __init__(self, config, megatron_config):
self.model = ParallelLlamaModel(config, megatron_config)
# lm_head 使用 ColumnParallelLinear,输出按 TP 切分
self.lm_head = tensor_parallel.ColumnParallelLinear(
input_size=config.hidden_size, output_size=config.vocab_size,
bias=False, gather_output=False, ...
)
def forward(self, input_ids, attention_mask, position_ids):
hidden_states = self.model(input_ids, attention_mask, position_ids)
logits = self.lm_head(hidden_states)[0]
# 从所有 TP rank 收集完整的 logits
logits = tensor_parallel.gather_from_tensor_model_parallel_region(logits)
return CausalLMOutputWithPast(logits=logits.float())
3. 去 Padding 版本 (RmPad)¶
去 padding 是核心优化:将 batch 中的 padding token 去除,只计算有效 token。
class ParallelLlamaForCausalLMRmPad(nn.Module):
def forward(self, input_ids, attention_mask, position_ids):
batch_size, sequence_length = input_ids.shape
# 1. 去除 padding,得到紧凑表示
input_ids, indices, cu_seqlens, max_seqlen_in_batch, *_ = unpad_input(
input_ids.unsqueeze(dim=-1), attention_mask
) # (total_nnz, 1)
# 2. 如果启用序列并行,pad 到 TP 的倍数
if self.megatron_config.sequence_parallel:
input_ids = sp_utils.pad_to_sequence_parallel(input_ids)
# 3. 前向传播
input_ids = input_ids.transpose(0, 1) # (1, total_nnz+pad)
outputs = self.model(input_ids=input_ids, ...)
# 4. 去掉 SP padding + 恢复原始 padding
logits = self._forward_head(hidden_states)
if self.megatron_config.sequence_parallel:
logits = logits[:cu_seqlens[-1]]
logits = pad_input(logits, indices, batch_size, seqlen=sequence_length)
数据流变化:(batch, seq) -> unpad -> (total_nnz,) -> SP pad -> (total_nnz_padded,) -> 模型计算 -> SP unpad -> pad_input -> (batch, seq)
4. Value 模型¶
Value 模型继承 CausalLM,但将 lm_head 改为输出单个标量值:
class ParallelLlamaForValueRmPad(ParallelLlamaForCausalLMRmPad):
def _init_head(self, config):
# 输出维度为 1(标量值)
self.lm_head = nn.Linear(in_features=config.hidden_size, out_features=1, bias=False)
# 标记为序列并行参数
sp_utils.mark_parameter_as_sequence_parallel(self.lm_head.weight)
def _forward_head(self, hidden_states):
logits = self.lm_head(hidden_states)
if self.megatron_config.sequence_parallel:
logits = tensor_parallel.gather_from_sequence_parallel_region(logits)
return logits
5. 流水线并行版本 (PP)¶
PP 版本的核心区别:每个 PP stage 只包含部分层。
class ParallelLlamaModelRmPadPP(nn.Module):
def __init__(self, config, megatron_config, pre_process, post_process):
# pre_process=True 的 stage 才有 embedding
if pre_process:
self.embed_tokens = tensor_parallel.VocabParallelEmbedding(...)
# 只创建属于本 PP stage 的层
pp_rank = mpu.get_pipeline_model_parallel_rank()
if vpp_size is not None:
# 虚拟 PP:交错分配层
offset = vpp_rank * (num_layers // vpp_size) + (pp_rank * num_layer_vpp_chunk)
else:
offset = pp_rank * num_layer_per_pp
for i in range(self.num_layer_this_model):
layer = ParallelLlamaDecoderLayerRmPad(config, megatron_config, layer_idx=offset + i)
self.layers.add_module(f"{i}", layer)
# post_process=True 的 stage 才有 final norm
if post_process:
self.norm = ParallelLlamaRMSNorm(config, megatron_config)
def set_input_tensor(self, input_tensor):
"""PP 中间 stage 通过此方法接收上一个 stage 的输出"""
self.input_tensor = input_tensor
def forward(self, input_ids, ...):
if self.pre_process:
hidden_states = self.embed_tokens(input_ids)
else:
hidden_states = self.input_tensor # 来自上一个 PP stage
for decoder_layer in self.layers:
hidden_states = decoder_layer(hidden_states, ...)
if self.post_process:
hidden_states = self.norm(hidden_states)
return hidden_states
辅助函数¶
def _make_causal_mask(input_ids_shape, dtype, device):
"""创建因果注意力掩码(下三角矩阵)"""
def _expand_mask(mask, dtype, tgt_len=None):
"""将 [bsz, seq_len] 的 mask 扩展为 [bsz, 1, tgt_len, src_len]"""
核心类列表¶
| 类名 | 父类 | 特点 |
|---|---|---|
ParallelLlamaModel |
nn.Module | 基础 Transformer,带 padding |
ParallelLlamaForCausalLM |
nn.Module | CausalLM,带 padding |
ParallelLlamaModelRmPad |
nn.Module | 去 padding,使用 flash_attn_varlen |
ParallelLlamaForCausalLMRmPad |
nn.Module | 去 padding 的 CausalLM |
ParallelLlamaForValueRmPad |
...RmPad | Value 模型,输出标量 |
ParallelLlamaModelRmPadPP |
nn.Module | 去 padding + PP |
ParallelLlamaForCausalLMRmPadPP |
nn.Module | 去 padding + PP 的 CausalLM |
ParallelLlamaForValueRmPadPP |
...RmPadPP | Value 模型 + PP |
与其他模块的关系¶
- 使用
layers/子目录中的并行层组件 - 被 verl 的训练流程调用(作为 actor/critic 模型)
- 与
checkpoint_utils/配合进行权重加载/保存
小结¶
这个文件提供了 LLaMA 在 Megatron 下的完整模型定义,按需支持三种并行模式。去 padding (RmPad) 通过 unpad_input/pad_input 减少无效计算,PP 版本通过 pre_process/post_process 标志和 set_input_tensor 方法实现流水线并行。Value 模型通过继承 CausalLM 并重写 head 来输出标量值。