跳转至

qwen2_vl.py — Qwen2-VL / Qwen2.5-VL 视觉语言模型适配

文件路径

verl/models/transformers/qwen2_vl.py

文件概述

为 Qwen2-VL 和 Qwen2.5-VL 视觉语言模型提供完整的适配支持,包括: 1. 3D 位置编码 (M-RoPE):为图像/视频生成时间-高度-宽度的三维位置 ID 2. 自定义 Flash Attention:处理 M-RoPE 带来的非单调位置 ID 3. 视觉嵌入注入:将视觉编码器输出替换到文本嵌入中 4. PPO 前向函数:Triton/Torch 两种融合计算后端

关键代码讲解

M-RoPE 位置编码 -- get_rope_index()

VLM 的核心挑战是:文本 token 只有一维位置,而图像/视频 token 有三维位置(时间 T、高度 H、宽度 W)。

def get_rope_index(processor, input_ids, image_grid_thw=None, video_grid_thw=None, ...):
    position_ids = torch.ones(3, input_ids.size(0), ...)  # (3, seqlen)

    for _ in range(image_nums + video_nums):
        if ed_image < ed_video:
            t, h, w = image_grid_thw[image_index]
            # 图像:t=1, h 和 w 根据 patch 大小计算
        else:
            t, h, w = video_grid_thw[video_index]
            # 视频:t=帧数

        # 文本区域:三个维度使用相同的递增位置
        llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)

        # 视觉区域:三个维度各自独立编码
        t_index = torch.arange(llm_grid_t).view(-1, 1).expand(-1, llm_grid_h * llm_grid_w).flatten()
        h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten()
        w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten()
        llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + text_len + st_idx)

关键理解:返回的 position_ids 形状为 (3, seqlen),三个维度分别编码时间、高度、宽度。对于纯文本 token,三个维度值相同。

自定义 Flash Attention -- _custom_flash_attention_forward()

M-RoPE 导致 position_ids 非单调递增(视觉 token 的位置在多维上"跳跃"),标准 Flash Attention 无法处理。

def _custom_flash_attention_forward(query_states, key_states, value_states, ...):
    sp_size = get_ulysses_sequence_parallel_world_size()
    if sp_size > 1:
        # Ulysses SP: AlltoAll 通信
        query_states = gather_seq_scatter_heads(query_states, seq_dim=1, head_dim=2)
        # ...
        # 还需要 all_gather position_ids!
        position_ids_lst = [torch.empty_like(position_ids) for _ in range(sp_size)]
        position_ids = dist.all_gather(position_ids_lst, position_ids, group=...)
        position_ids = torch.cat(position_ids_lst, dim=-1)

    # 非单调位置 -> 使用 varlen 版本
    if not (torch.diff(position_ids, dim=-1) >= 0).all():
        q, k, v, (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = \
            prepare_fa2_from_position_ids(query_states, key_states, value_states, position_ids)
        attn_output = flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
    else:
        attn_output = _flash_attention_forward(...)

注意:对于 VLM,在 Ulysses SP 中除了 QKV 要做 AlltoAll,position_ids 也需要 all_gather,因为 varlen attention 需要完整的位置信息来确定子序列边界。

视觉嵌入注入 -- _get_input_embeds()

def _get_input_embeds(model, input_ids, attention_mask, pixel_values, ...):
    inputs_embeds = model.get_input_embeddings()(input_ids)  # 文本嵌入

    if pixel_values is not None:
        image_embeds = model.visual(pixel_values, grid_thw=image_grid_thw)  # 视觉编码
        mask = input_ids == model.config.image_token_id
        inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)  # 替换

    # 对于纯文本数据(无图像),创建虚拟视觉输入以保持计算图连通
    if pixel_values is None and pixel_values_videos is None:
        pixel_values = torch.zeros((16, patch_dim), ...)
        image_embeds = model.visual(pixel_values, grid_thw=...)
        inputs_embeds += 0.0 * image_embeds.mean()  # 梯度可以流过视觉编码器

设计亮点:纯文本数据时构造虚拟视觉输入并加上 0.0 * mean(),确保视觉编码器参数有梯度流过,避免 FSDP 训练中的同步问题。

PPO 前向函数

与 dense_common.py 类似,提供三种后端:

def forward_with_normal_backend(self, input_ids, ...):
    """直接输出 logits,不计算 log_probs"""
    outputs = qwen2_vl_forward(self, input_ids, **kwargs)
    logits = self.lm_head(outputs[0])
    return Qwen2VLCausalLMOutputWithPast(logits=logits, ...)

def forward_with_torch_backend(self, input_ids, ...):
    """使用 PyTorch 融合计算 log_probs + entropy"""
    fused_linear_for_ppo = FusedLinearForPPO()
    log_probs, entropy = fused_linear_for_ppo.forward(...)

def forward_with_triton_backend(self, input_ids, ...):
    """使用 Triton kernel 融合计算 log_probs + entropy"""
    log_probs, entropy = linear_cross_entropy(...)

核心类/函数列表

名称 作用
get_rope_index() 生成 3D M-RoPE 位置 ID
_custom_flash_attention_forward() 处理非单调位置 ID 的 Flash Attention
qwen2_vl_attn_forward() Qwen2-VL 注意力层替换函数
_get_input_embeds() 视觉嵌入注入
Qwen2VLCausalLMOutputForPPO PPO 专用输出数据结构
forward_with_normal_backend() 标准前向(输出 logits)
forward_with_torch_backend() PyTorch 融合前向
forward_with_triton_backend() Triton 融合前向
prepare_fa2_from_position_ids() 从 position_ids 构造 varlen attention 参数
process_position_ids() 处理 4D position_ids 的版本兼容

与其他模块的关系

  • 被 monkey_patch.py 中的 apply_monkey_patch() 和 patch_forward_with_backends() 调用
  • 依赖 verl.utils.ulysses 的 AlltoAll 原语
  • 使用 flash_attn_varlen_func 处理 varlen attention

小结

Qwen2-VL 的适配是所有 VLM 中最完整的参考实现。核心挑战是处理多模态位置编码带来的非单调性,以及在 Ulysses SP 中正确同步位置信息。_get_input_embeds 中的虚拟视觉输入设计是保证 FSDP 训练稳定的关键技巧。