跳转至

monkey_patch.py — 统一 Monkey Patch 入口

文件路径

verl/models/transformers/monkey_patch.py

文件概述

这是 transformers 适配层中 最核心 的文件。它提供了 apply_monkey_patch 函数,在运行时动态替换 HuggingFace 模型的注意力层和前向函数,从而实现:

  1. Ulysses 序列并行:在注意力计算前后插入 AlltoAll 通信
  2. 融合 kernel:替换前向函数以使用 fused linear cross entropy
  3. TiledMLP:将 MLP 分片计算以降低显存
  4. VLM 适配:为视觉语言模型(Qwen2-VL, Qwen3-VL, GLM-4V, KimiVL)做特殊处理
  5. 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 如何适配各种模型的关键。