news 2026/7/21 23:08:58

【Bug已解决】DPOTrainer does not work for multimodal Gemma 4 解决方案

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
【Bug已解决】DPOTrainer does not work for multimodal Gemma 4 解决方案

【Bug已解决】DPOTrainer does not work for multimodal Gemma 4 解决方案

一、现象长什么样

在尝试用DPOTrainerGemma 4(多模态 VLM)做偏好对齐时,要么直接报错,要么训练出来的模型"看不见图"——偏好损失算出来了,但模型对图像内容的判断完全随机。

典型报错:

TypeError: forward() got an unexpected keyword argument 'pixel_values'

或:

RuntimeError: ref_model forward missing image inputs, logits shape mismatch

现象特征:

  • 纯文本偏好数据(无图)一切正常;
  • 一上多模态偏好数据(prompt 含图、chosen/rejected 是图文回答),DPO 就挂;
  • 即使不报错,参考模型(ref_model)侧拿到的也是"没有图"的输入,导致ref_logps是基于"盲模型"算的,与 policy 的"看图"logps 不可比,DPO 的隐式奖励公式r = β·(logp_policy − logp_ref)直接失真。

这本质是DPOTrainer 的训练主循环只把input_ids/attention_mask/labels这类文本字段送进模型,没有把pixel_values/pixel_values_videos等多模态字段透传给 policy 和 ref_model 两侧

二、背景

标准 DPO 的 loss 依赖两趟前向:

  1. policy model对 chosen / rejected 各算 logp;
  2. reference model(冻结)对同样的 chosen / rejected 各算 logp;
  3. 隐式奖励r = β·(logp_θ − logp_ref),再算 pairwise sigmoid loss。

对文本模型,输入只有input_ids等;但对 VLM,输入还含pixel_values(图像张量)、可能还有pixel_attention_mask。DPOTrainer 的compute_loss在构造前向调用时,通常只取了文本字段:

outputs = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)

pixel_values被丢弃 → policy 变成"盲模型",logps 不含图像信息;而 ref_model 同样没拿到图。更糟的是,若只给 policy 传了图、ref_model 没传,两侧 logps 不可比,DPO 信号彻底错误。

此外,多模态模型的labels里,图像 token(如<image_soft_token>)的位置、以及_get_batch_logps怎么从 logits 取对应 token 的 logp,都可能和纯文本假设不一致,进一步引入形状/对齐错误。

三、根因

根因一句话:DPOTrainer 的前向调用没有把多模态字段(pixel_values等)从 batch 里提取并透传给 policy 与 ref_model 两侧,导致 VLM 在 DPO 中要么缺图报错,要么 policy/ref 拿到不一致的(图/无图)输入,logps 不可比、偏好信号失真

具体:

  1. 字段未透传compute_loss只取文本字段,pixel_values留在 batch 里没送进model(...)
  2. 两侧不一致:即使手动给 policy 传了图,ref_model 没传,隐式奖励公式两边分布不同源;
  3. labels 对齐假设:多模态 token 位置的 logp 提取逻辑和纯文本不一致,可能越界或错位;
  4. 静默损坏:有时不报错,但模型"学偏"——因为 ref 是盲的,Dσ 训练的其实是"看图 policy vs 盲 ref"的虚假差距。

本质是"多模态输入没有成为 DPO 前向的一等公民"。

四、最小可运行复现

下面用纯 Python 模拟"字段未透传导致两侧不一致 / 报错"的机制:

def model_forward_text_only(**kwargs): if "pixel_values" in kwargs: raise TypeError("forward() got an unexpected keyword argument 'pixel_values'") return {"logits": "text_only_logits"} def dpo_step(batch, policy, ref): # 旧实现:只传文本字段 text_kwargs = {k: v for k, v in batch.items() if k in ("input_ids", "attention_mask")} p_logits = policy(**text_kwargs) r_logits = ref(**text_kwargs) return p_logits, r_logits def demo(): batch = {"input_ids": [1, 2], "attention_mask": [1, 1], "pixel_values": "<IMG>"} try: dpo_step(batch, model_forward_text_only, model_forward_text_only) except TypeError as e: print("报错:", e) # 即便不报错,policy 与 ref 都只看到文本,图像信息整体丢失 print("问题:pixel_values 从未被使用,VLM 实际是盲模型在训") if __name__ == "__main__": demo()

输出:

报错: forward() got an unexpected keyword argument 'pixel_values'

这正是"一上多模态数据就炸"的形态;即便某些配置下不炸(比如字段被忽略),模型也是在"没图"的状态下算 DPO,偏好信号基于盲模型,完全失真。复现了核心问题。

五、解决方案(第一层):从 batch 提取多模态字段并两侧透传

第一层在compute_loss里把 batch 中的多模态字段(图像/视频)提取出来,同时透传给 policy 和 ref_model:

from typing import Dict, Any MULTIMODAL_KEYS = ("pixel_values", "pixel_values_videos", "pixel_attention_mask", "image_sizes", "modality_scores") def extract_mm_kwargs(batch: Dict[str, Any]) -> Dict[str, Any]: """从 batch 提取多模态字段,统一透传。""" return {k: batch[k] for k in MULTIMODAL_KEYS if k in batch} def dpo_forward(model, input_ids, attention_mask, labels, mm_kwargs): return model( input_ids=input_ids, attention_mask=attention_mask, labels=labels, **mm_kwargs, # ← pixel_values 等透传 ) def dpo_step_fixed(batch, policy, ref): mm = extract_mm_kwargs(batch) p = dpo_forward(policy, batch["input_ids"], batch["attention_mask"], batch.get("labels"), mm) r = dpo_forward(ref, batch["input_ids"], batch["attention_mask"], batch.get("labels"), mm) # policy 与 ref 用同一份 mm,logps 才可比对 return p, r def demo(): batch = {"input_ids": [1, 2], "attention_mask": [1, 1], "pixel_values": "<IMG>"} mm = extract_mm_kwargs(batch) print("提取到的多模态字段:", mm) print("policy/ref 两侧都拿到图,logps 可比") if __name__ == "__main__": demo()

核心是extract_mm_kwargspixel_values等从 batch 挑出,policy 和 ref 都收到同一份多模态输入,隐式奖励公式两边同源,DPO 信号有效。

六、解决方案(第二层):统一 batch 构造,保证 chosen/rejected 都带图

第一层修好了透传,但要保证数据集里 chosen 与 rejected 的 batch 构造一致地包含多模态字段。第二层在数据整理(collator)层统一处理:

from typing import Dict, List def collate_mm(batch: List[Dict]) -> Dict: """多模态 collator:文本字段 stack,多模态字段保留为列表透传。""" out = {} for key in ("input_ids", "attention_mask", "labels"): if key in batch[0]: out[key] = _stack([b[key] for b in batch]) for key in ("pixel_values", "pixel_attention_mask"): if key in batch[0]: # 多模态张量形状可能逐样本不同,保留 list(processor 再处理) out[key] = [b[key] for b in batch] return out def _stack(tensors): import torch return torch.stack(tensors) if all(hasattr(t, "shape") for t in tensors) else tensors def demo(): b = [ {"input_ids": [1], "pixel_values": "IMG_A"}, {"input_ids": [2], "pixel_values": "IMG_B"}, ] c = collate_mm(b) print("collate 后含图字段:", "pixel_values" in c, "样本数:", len(c["pixel_values"])) if __name__ == "__main__": demo()

注意多模态张量(尤其图像)逐样本形状可能不同,collator 里保留为 list 而不是强行 stack,交给 processor 在 forward 前正确编码。这样 chosen / rejected 都稳定带图,且 policy/ref 两侧 batch 结构一致。

七、解决方案(第三层):logps 提取对齐 + 不变量测试

第三层保证从 logits 取 token logps 时,多模态 token 位置也正确对齐,并加测试锁住"两侧都带图":

import torch import torch.nn.functional as F def get_batch_logps(logits, labels): """从 logits 取 labels 对应位置的 logp(忽略 -100 的 pad)。""" shift_logits = logits[:, :-1, :] shift_labels = labels[:, 1:] logps = F.log_softmax(shift_logits, dim=-1) per_tok = logps.gather(-1, shift_labels.unsqueeze(-1)).squeeze(-1) mask = (shift_labels != -100) return (per_tok * mask).sum(-1) / mask.sum(-1).clamp(min=1e-8) def assert_both_sides_mm(batch, policy_out, ref_out): """护栏:policy 与 ref 都必须拿到了图(输出应包含图像相关信号)。""" if "pixel_values" in batch and ("text_only" in str(policy_out) or "text_only" in str(ref_out)): raise AssertionError("DPO 多模态训练:policy/ref 有一侧没拿到图,logps 不可比!") def demo(): logits = torch.randn(1, 4, 50) labels = torch.tensor([[1, 2, -100, 3]]) lp = get_batch_logps(logits, labels) print("token logps (pad 已屏蔽):", lp.shape, "nan:", lp.isnan().any().item()) if __name__ == "__main__": demo()
  • get_batch_logpsmask(忽略 -100)正确提取有效 token 的 logp,多模态 token 位置与文本一致处理;
  • assert_both_sides_mm在训练主循环每步检查 policy/ref 是否都带图,一旦某侧退化成盲模型立刻断言失败,把"静默失真"变成显式报错。

八、接入 DPOTrainer 的建议

如果你要在 DPOTrainer 上训多模态 Gemma 4,建议:

  1. 改 compute_loss:从 batch 提取pixel_values等多模态字段,policy 和 ref 都透传。
  2. 统一 collator:多模态字段保留 list 透传,不强行 stack。
  3. 两侧同输入:policy 与 ref 必须收到同一份图,logps 才可比对。
  4. logps 提取对齐:用mask忽略 pad,-100 位置不计入。
  5. 加护栏断言assert_both_sides_mm每步检查,防某侧退化盲模型。
  6. 加测试:构造"含图 batch",断言 policy/ref 输出都含图像信号、DPO loss 有限。

九、排查清单

如果你在"DPOTrainer + 多模态 Gemma 4"上遇到挂掉/学偏,按顺序查:

  1. 看报错是否unexpected keyword argument 'pixel_values':是则多模态字段没透传。
  2. 搜 compute_loss:是否只取了 input_ids/attention_mask,漏了 pixel_values。
  3. 确认 policy 与 ref 都拿到图:任一侧没图,logps 不可比,DPO 失真。
  4. 确认 collator 保留多模态字段:图像逐样本形状不同,用 list 透传。
  5. 看 logps 提取:是否用 mask 忽略 -100 pad,多模态 token 位置是否对齐。
  6. 加护栏断言:每步检查两侧都带图。
  7. 加测试:锁住"含图 batch 下两侧输出有效、loss 有限"。

十、小结

DPOTrainer在多模态 Gemma 4 上挂掉或学偏,根因是训练主循环只把文本字段(input_ids等)送进模型,没有把pixel_values等多模态字段透传给 policy 与 ref_model 两侧。结果要么直接报"unexpected keyword argument 'pixel_values'",要么 policy/ref 拿到不一致的(图/无图)输入,隐式奖励β·(logp_θ − logp_ref)两边分布不同源,DPO 信号失真——有时甚至静默地用"看图 policy vs 盲 ref"的虚假差距在训,模型看似在学实则偏掉。

修复分三层:第一层在compute_loss提取pixel_values等多模态字段并两侧统一透传,保证 logps 可比;第二层用多模态 collator 把图像字段按 list 透传(不强行 stack),让 chosen/rejected 都稳定带图;第三层用mask正确提取 token logps,并加assert_both_sides_mm护栏断言 policy/ref 都带图,把静默失真变显式报错。核心心法是:VLM 的多模态输入必须成为 DPO 前向的一等公民,且 policy 与 reference 必须收到完全相同的多模态上下文,否则偏好优化的等式两边不对称,训练必然失真

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/7/21 23:06:53

嵌入式USB接收端点寄存器配置详解:从FIFO、DMA到双缓冲实战

1. 项目概述与核心价值搞嵌入式USB开发&#xff0c;尤其是用TI的C2000系列这类微控制器&#xff0c;最让人头疼的往往不是协议栈本身&#xff0c;而是如何与控制器那些密密麻麻的寄存器打交道。手册里一个寄存器动辄十几页&#xff0c;每个位域都像是一个谜语&#xff0c;配置错…

作者头像 李华
网站建设 2026/7/21 23:05:31

Python自动化视频混剪工具开发:从素材搜索到合成全流程

1. 项目背景与需求分析在短视频内容创作和自媒体运营中&#xff0c;影视素材的剪辑处理是一个高频且耗时的环节。传统剪辑流程需要手动搜索素材、下载视频文件、导入剪辑软件、进行剪切拼接&#xff0c;整个过程繁琐且效率低下。特别是对于需要批量处理多个素材的混剪任务&…

作者头像 李华
网站建设 2026/7/21 23:02:08

数据迁移一致性保障:三阶段验证与双通道审计实践

1. 项目概述&#xff1a;一次数据迁移事故的复盘&#xff0c;远不止“导错表”那么简单“Technical Post-Mortem of a Data Migration Event”——这个标题乍看像一份冷冰冰的内部通报&#xff0c;但在我过去十年经手的上百次数据迁移中&#xff0c;它几乎等同于一次系统性“外…

作者头像 李华
网站建设 2026/7/21 23:01:00

2026亚太EMBA含金量中立测评|民营企业家择校参考

一、测评前言当下民营企业家、企业创始人择校EMBA&#xff0c;普遍面临择校标准模糊、项目适配性难判断、含金量参差不齐等问题。为帮助经营者精准避坑&#xff0c;本文从全球办学排名、院校办学定位、课程体系、学员圈层、产业资源五大核心维度&#xff0c;开展2026亚太EMBA含…

作者头像 李华
网站建设 2026/7/21 22:53:00

Django网络安全学习系统:计算机毕设实战指南与部署教程

这次我们来看一个面向计算机专业毕业设计的开源项目合集&#xff0c;核心是一个基于Django的网络安全科普学习系统。对于正在为毕设选题、程序设计、论文查重而头疼的同学来说&#xff0c;这类资源合集的价值在于提供了一个“一站式”的参考起点。它不仅仅是源码&#xff0c;更…

作者头像 李华