【Bug已解决】DPOTrainer does not work for multimodal Gemma 4 解决方案
【Bug已解决】DPOTrainer does not work for multimodal Gemma 4 解决方案
一、现象长什么样
在尝试用DPOTrainer对Gemma 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 依赖两趟前向:
- policy model对 chosen / rejected 各算 logp;
- reference model(冻结)对同样的 chosen / rejected 各算 logp;
- 隐式奖励
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 不可比、偏好信号失真。
具体:
- 字段未透传:
compute_loss只取文本字段,pixel_values留在 batch 里没送进model(...); - 两侧不一致:即使手动给 policy 传了图,ref_model 没传,隐式奖励公式两边分布不同源;
- labels 对齐假设:多模态 token 位置的 logp 提取逻辑和纯文本不一致,可能越界或错位;
- 静默损坏:有时不报错,但模型"学偏"——因为 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_kwargs把pixel_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_logps用mask(忽略 -100)正确提取有效 token 的 logp,多模态 token 位置与文本一致处理;assert_both_sides_mm在训练主循环每步检查 policy/ref 是否都带图,一旦某侧退化成盲模型立刻断言失败,把"静默失真"变成显式报错。
八、接入 DPOTrainer 的建议
如果你要在 DPOTrainer 上训多模态 Gemma 4,建议:
- 改 compute_loss:从 batch 提取
pixel_values等多模态字段,policy 和 ref 都透传。 - 统一 collator:多模态字段保留 list 透传,不强行 stack。
- 两侧同输入:policy 与 ref 必须收到同一份图,logps 才可比对。
- logps 提取对齐:用
mask忽略 pad,-100 位置不计入。 - 加护栏断言:
assert_both_sides_mm每步检查,防某侧退化盲模型。 - 加测试:构造"含图 batch",断言 policy/ref 输出都含图像信号、DPO loss 有限。
九、排查清单
如果你在"DPOTrainer + 多模态 Gemma 4"上遇到挂掉/学偏,按顺序查:
- 看报错是否
unexpected keyword argument 'pixel_values':是则多模态字段没透传。 - 搜 compute_loss:是否只取了 input_ids/attention_mask,漏了 pixel_values。
- 确认 policy 与 ref 都拿到图:任一侧没图,logps 不可比,DPO 失真。
- 确认 collator 保留多模态字段:图像逐样本形状不同,用 list 透传。
- 看 logps 提取:是否用 mask 忽略 -100 pad,多模态 token 位置是否对齐。
- 加护栏断言:每步检查两侧都带图。
- 加测试:锁住"含图 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 必须收到完全相同的多模态上下文,否则偏好优化的等式两边不对称,训练必然失真。
