跳转至

qwen3_vl.py — Qwen3-VL 视觉语言模型适配

文件路径

verl/models/transformers/qwen3_vl.py

文件概述

为 Qwen3-VL 视觉语言模型提供 PPO 训练适配。整体结构与 qwen2_vl.py 类似,但有两个重要区别: 1. DeepStack 视觉嵌入:Qwen3-VL 的视觉编码器额外输出多层深度堆叠 (deepstack) 嵌入 2. 视频时间编码:使用时间戳分隔视频帧,每帧独立编码 3. MoE Bug 修复:修复了 transformers 中 Qwen3-VL MoE 的一个 bug

关键代码讲解

get_rope_index() -- 视频编码的区别

与 Qwen2-VL 不同,Qwen3-VL 使用时间戳分隔视频帧:

def get_rope_index(processor, input_ids, image_grid_thw=None, video_grid_thw=None, ...):
    # Qwen3-VL 特有:将视频按帧展开
    if video_grid_thw is not None:
        video_grid_thw = torch.repeat_interleave(video_grid_thw, video_grid_thw[:, 0], dim=0)
        video_grid_thw[:, 0] = 1  # 每帧 t=1

关键区别:Qwen2-VL 将整个视频作为一个 3D 块编码(t=帧数, h, w),而 Qwen3-VL 将每帧拆开,用时间戳 token 分隔,每帧只编码空间维度。

_get_input_embeds() -- DeepStack 嵌入

def _get_input_embeds(model, input_ids, ...):
    if pixel_values is not None:
        # Qwen3-VL 视觉编码器返回两个值
        image_embeds, deepstack_image_embeds = model.visual(pixel_values, grid_thw=image_grid_thw)
        inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)

    # 处理 deepstack 嵌入:用于模型内部的多层视觉特征注入
    visual_pos_masks = None
    deepstack_visual_embeds = None
    if image_mask is not None and video_mask is not None:
        visual_pos_masks = image_mask | video_mask
        deepstack_visual_embeds = []
        for img_embed, vid_embed in zip(deepstack_image_embeds, deepstack_video_embeds):
            embed_joint = img_embed.new_zeros(visual_pos_masks.sum(), img_embed.shape[-1])
            embed_joint[image_mask_joint, :] = img_embed
            embed_joint[video_mask_joint, :] = vid_embed
            deepstack_visual_embeds.append(embed_joint)

    # 返回字典而非元组(与 Qwen2-VL 不同)
    return {
        "inputs_embeds": inputs_embeds,
        "attention_mask": attention_mask,
        "visual_pos_masks": visual_pos_masks,
        "deepstack_visual_embeds": deepstack_visual_embeds,
    }

DeepStack 是 Qwen3-VL 的特色:视觉编码器除了输出最终嵌入,还输出中间层的特征,这些特征会在语言模型的不同层中注入,增强视觉理解能力。

MoE Bug 修复

def patch_qwen3_vl_moe_sparse_moe_block_forward():
    """修复 transformers 4.57.3 中的 bug:
    原始代码错误地使用 torch.zeros_like(hidden_states)
    应该使用 torch.zeros_like(router_logits)"""

    def patched_forward(self, hidden_states):
        router_logits = self.gate(hidden_states)
        # BUG FIX: 原来是 routing_weights.to(hidden_states.dtype)
        routing_weights = routing_weights.to(router_logits.dtype)
        router_weights = torch.zeros_like(router_logits).scatter_(1, router_indices, routing_weights)

PPO 前向函数

与 Qwen2-VL 类似,提供 normal/torch/triton 三种后端:

def qwen3_vl_base_forward(self, input_ids, ...):
    input_kwargs = _get_input_embeds(self, input_ids, ...)
    kwargs.update(input_kwargs)  # 将 deepstack 等信息传给 language_model
    return self.language_model(input_ids=None, **kwargs)

核心类/函数列表

名称 作用
get_rope_index() Qwen3-VL 的 3D 位置编码生成
_get_input_embeds() 视觉嵌入注入(含 DeepStack)
Qwen3VLCausalLMOutputForPPO PPO 专用输出结构
qwen3_vl_base_forward() 基础前向(处理视觉输入)
forward_with_normal_backend() 标准前向
forward_with_torch_backend() PyTorch 融合前向
forward_with_triton_backend() Triton 融合前向
patch_qwen3_vl_moe_sparse_moe_block_forward() MoE bug 修复补丁

与 Qwen2-VL 的主要区别

方面 Qwen2-VL Qwen3-VL
视觉编码输出 单一嵌入 嵌入 + DeepStack 多层嵌入
视频编码 整块 3D (t, h, w) 按帧拆分,时间戳分隔
_get_input_embeds 返回 元组 字典(含 visual_pos_masks 等)
MoE 支持 无 有(含 bug 修复)
注意力适配 独立实现 复用 qwen2_vl 的 flash attention

与其他模块的关系

  • 被 monkey_patch.py 调用
  • 视觉嵌入注入模式与 qwen2_vl.py 类似
  • 注意力层复用 Qwen2-VL 或通用的 flash attention

小结

Qwen3-VL 在 Qwen2-VL 基础上引入了 DeepStack 多层视觉特征和按帧的视频编码。代码结构类似但细节更复杂,特别是 DeepStack 嵌入的处理需要额外跟踪视觉 token 的位置掩码。