【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 各算 logpreference 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_idsinput_ids, attention_maskattention_mask, labelslabels)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_modelfrom 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_idsinput_ids, attention_maskattention_mask, labelslabels, **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 用同一份 mmlogps 才可比对 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]: # 多模态张量形状可能逐样本不同保留 listprocessor 再处理 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(min1e-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 必须收到完全相同的多模态上下文否则偏好优化的等式两边不对称训练必然失真。