跳转至

glm4v.py — GLM-4V 视觉语言模型适配

文件路径

verl/models/transformers/glm4v.py

文件概述

为 GLM-4V(智谱 AI)视觉语言模型提供 PPO 训练适配。整体框架与 qwen2_vl.py 非常相似(M-RoPE + varlen attention),但在视频帧处理上有独特设计。

关键代码讲解

get_rope_index() -- 视频帧编码的独特方式

def get_rope_index(processor, input_ids, image_grid_thw=None, video_grid_thw=None, ...):
    image_token_id = processor.tokenizer.convert_tokens_to_ids("<|image|>")
    video_start_token_id = processor.tokenizer.convert_tokens_to_ids("<|begin_of_video|>")
    video_end_token_id = processor.tokenizer.convert_tokens_to_ids("<|end_of_video|>")

    # 先分类每个 token 的类型:text / image / video
    input_token_type = []
    video_check_flg = False
    for token in input_tokens:
        if token == video_start_token_id:
            video_check_flg = True
        elif token == video_end_token_id:
            video_check_flg = False

        if token == image_token_id and not video_check_flg:
            input_token_type.append("image")
        elif token == image_token_id and video_check_flg:
            input_token_type.append("video")
        else:
            input_token_type.append("text")

独特设计:GLM-4V 使用 <|begin_of_video|> 和 <|end_of_video|> 标记视频区域。在视频标记内的 <|image|> token 被视为视频帧,在外部的被视为普通图像。

视频帧递增编码

for modality_type, start_idx, end_idx in input_type_group:
    if modality_type == "video":
        t = video_frame_num  # 帧号递增
        # ...
        for t_idx in range(llm_grid_t):
            t_index = torch.tensor(t_idx).view(-1, 1).expand(-1, h * w).flatten()
            # 每帧独立编码空间位置
        video_frame_num += 1  # 帧号递增
    else:
        video_frame_num = 1  # 遇到非视频内容重置帧号

与 Qwen2-VL 将整个视频作为一个 3D 块不同,GLM-4V 逐帧编码,帧号随视频帧递增,遇到文本重置。

注意力和前向函数

注意力层(glm4v_attn_forward)和前向函数(forward_with_torch/triton_backend)的模式与 Qwen2-VL 完全相同:

def glm4v_attn_forward(self, hidden_states, attention_mask, position_ids, position_embeddings, ...):
    # 与 qwen2_vl_attn_forward 结构相同
    cos, sin = position_embeddings
    query_states, key_states = apply_multimodal_rotary_pos_emb(...)

    attn_output = _custom_flash_attention_forward(
        ..., position_ids=position_ids, ...
    )

_get_input_embeds()

与 Qwen2-VL 基本相同,使用 masked_scatter 注入视觉嵌入:

def _get_input_embeds(model, input_ids, 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, 1176), ...)  # GLM-4V 特定 patch 维度

process_position_ids()

def process_position_ids(position_ids):
    if position_ids.ndim != 3 or position_ids.size(0) != 4:
        raise ValueError("position_ids should be (4, batch_size, seq_length)")
    return position_ids  # GLM-4V 不需要像 Qwen2-VL 那样丢弃 text position ids

GLM-4V 的 position_ids 是 4D(text + t/h/w),直接透传。

核心类/函数列表

名称 作用
get_rope_index() GLM-4V 的 3D 位置编码(含视频帧递增)
_custom_flash_attention_forward() 自定义 Flash Attention(处理非单调位置)
glm4v_attn_forward() GLM-4V 注意力层替换函数
_get_input_embeds() 视觉嵌入注入
Glm4vCausalLMOutputForPPO PPO 专用输出结构
forward_with_normal_backend() 标准前向
forward_with_torch_backend() PyTorch 融合前向
forward_with_triton_backend() Triton 融合前向

与 Qwen2-VL 的区别

方面 Qwen2-VL GLM-4V
视频标记 <\|video_pad\|> token <\|begin_of_video\|>...<\|end_of_video\|> 包裹
帧编码 整块 3D 编码 逐帧递增编码
position_ids 丢弃 text 维度 保留完整 4D
token 分类 按 token ID 直接分类 使用状态机(video_check_flg)

与其他模块的关系

  • 被 monkey_patch.py 调用
  • 共享 qwen2_vl.py 的 _custom_flash_attention_forward 设计模式
  • 依赖 verl.utils.ulysses 的 AlltoAll 通信

小结

GLM-4V 的适配在整体架构上与 Qwen2-VL 相似(M-RoPE + varlen attention),但在视频帧处理上更加灵活。使用视频标记对(begin/end)来界定视频区域,并对帧号递增编码,更自然地处理视频中的多帧场景。