monkey_patch.py — 统一 Monkey Patch 入口¶
文件路径¶
verl/models/transformers/monkey_patch.py
文件概述¶
这是 transformers 适配层中 最核心 的文件。它提供了 apply_monkey_patch 函数,在运行时动态替换 HuggingFace 模型的注意力层和前向函数,从而实现:
- Ulysses 序列并行:在注意力计算前后插入 AlltoAll 通信
- 融合 kernel:替换前向函数以使用 fused linear cross entropy
- TiledMLP:将 MLP 分片计算以降低显存
- VLM 适配:为视觉语言模型(Qwen2-VL, Qwen3-VL, GLM-4V, KimiVL)做特殊处理
- PrefixGrouper:支持前缀分组注意力
关键代码讲解¶
核心入口函数 apply_monkey_patch¶
def apply_monkey_patch(
model: PreTrainedModel,
ulysses_sp_size: int = 1,
use_remove_padding: bool = True,
use_fused_kernels: bool = False,
fused_kernels_backend: str = None,
use_prefix_grouper: bool = False,
use_tiled_mlp: bool = False,
tiled_mlp_shards: int = 4,
):
这个函数根据模型类型(model.config.model_type)分发到不同的处理逻辑。
Ulysses 序列并行的核心:_ulysses_flash_attention_forward¶
def _ulysses_flash_attention_forward(
query_states, key_states, value_states, attention_mask, query_length,
*args, position_ids=None, **kwargs,
):
ulysses_sp_size = get_ulysses_sequence_parallel_world_size()
if ulysses_sp_size > 1 and position_ids is not None:
# 重复 KV 头以适配序列并行切分
repeats = max(ulysses_sp_size // key_states.size(2), 1)
key_states = repeat_kv(key_states, repeats)
value_states = repeat_kv(value_states, repeats)
# AlltoAll: 序列维度聚合,注意力头维度切分
# (bsz, seq_len/n, n_head, head_dim) -> (bsz, seq_len, n_head/n, head_dim)
query_states = gather_seq_scatter_heads(query_states, seq_dim=1, head_dim=2)
key_states = gather_seq_scatter_heads(key_states, seq_dim=1, head_dim=2)
value_states = gather_seq_scatter_heads(value_states, seq_dim=1, head_dim=2)
# 执行 Flash Attention
attn_output = _flash_attention_forward(...)
if ulysses_sp_size > 1 and position_ids is not None:
# 反向 AlltoAll: 注意力头聚合,序列维度切分
attn_output = gather_heads_scatter_seq(attn_output, seq_dim=1, head_dim=2)
return attn_output
Ulysses 序列并行的核心思想: - 每个 GPU 只持有序列的 1/N(N 是并行度) - 在注意力计算前,通过 AlltoAll 通信将完整序列交换到每个 GPU,同时将注意力头分散 - 注意力计算后,再通过 AlltoAll 将结果换回来
VLM 输入切片 patch¶
def patch_vlm_for_ulysses_input_slicing(model_class: type):
"""为 VLM 模型的 decoder 打补丁,在第一次 forward 时切分输入"""
def ulysses_wrapped_decoder_forward(self, *args, **kwargs):
if slice_now:
call_kwargs["inputs_embeds"] = slice_input_tensor(inputs_embeds, dim=1, padding=False)
call_kwargs["position_ids"] = slice_input_tensor(position_ids, dim=-1, padding=False)
# ...
VLM 模型的特殊之处在于:视觉 embedding 需要在完整序列上计算后,再进行序列并行切分。
前向函数替换¶
def patch_forward_with_backends(model, use_fused_kernels, fused_kernels_backend):
if model.config.model_type in ["qwen2_5_vl", "qwen2_vl"]:
from verl.models.transformers.qwen2_vl import forward_with_torch_backend, forward_with_triton_backend
# ... 根据模型类型选择对应的 forward 实现
if fused_kernels_backend == "triton":
model.__class__.forward = forward_with_triton_backend_function
elif fused_kernels_backend == "torch":
model.__class__.forward = forward_with_torch_backend_function
核心类/函数列表¶
| 名称 | 作用 |
|---|---|
apply_monkey_patch() |
统一的 monkey patch 入口 |
_ulysses_flash_attention_forward() |
带 Ulysses SP 的 Flash Attention |
patch_vlm_for_ulysses_input_slicing() |
VLM 模型的输入切片适配 |
patch_forward_with_backends() |
根据后端替换前向函数 |
apply_prefix_grouper_patch() |
PrefixGrouper 注意力补丁 |
repeat_kv() |
重复 KV 头以适配 GQA/MQA |
与其他模块的关系¶
- 调用
llama.py,qwen2.py等的注意力替换函数 - 调用
dense_common.py,qwen2_vl.py等的前向替换函数 - 调用
tiled_mlp.py的 MLP 替换函数 - 依赖
verl.utils.ulysses提供的 AlltoAll 通信原语
小结¶
这是 verl 在 HuggingFace 模型上做优化的核心枢纽。通过运行时替换模型方法,避免了修改 transformers 源码,同时实现了序列并行、融合 kernel 等关键优化。理解这个文件是理解 verl 如何适配各种模型的关键。