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 训练稳定的关键技巧。