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 的位置掩码。