公司动态
VQA模型高分却未看图?ReKey稀疏微调让模型真正用上视觉证据
VQA 模型在榜单上拿高分不一定说明模型真的在“看图”。VQAVisual Question Answering视觉问答要同时理解图像和问题输出一个答案。但很多模型会在训练集里学到答案频率只要问题中出现“有没有”“什么颜色”等词就输出高频答案图像只被当成可有可无的装饰。这种“记答案”的行为会让分数好看却不能泛化到真实场景。ReKey 是一种试图改变这种局面的技术思路不重新训练整个模型而是定位真正决定答案的那一小块证据只更新与这一小块相关的参数从而把模型的注意力拉回图像本身。接下来的内容会从 VQA 为什么会出现“高分但没在看图”的现象讲起拆解 ReKey 的核心机制再给出一个可以自己动手的最小实验流程最后补充验证指标、常见问题和生产落地建议。1. VQA 的高分为什么不可信先区分“看图”和“记答案”1.1 VQA 任务的定义以及它为什么容易走捷径VQA 的输入既有图像又有文本输出是答案。这个任务天然有两条解题路径一条是真正理解图像中的对象、关系、颜色、动作再回答问题另一条是仅从问题文本中推断答案甚至依赖训练集的分布统计。例如VQA 数据集中“这个物体是什么颜色”这类问题训练样本里“红色”出现的频率可能远高于其他颜色。模型如果把“颜色红色”作为一个强先验就能在不少样本上得到正确答案并且不需要真正区分红色物体。这不只是训练技巧问题而是数据偏差与模型归纳偏向共同作用的结果。VQA 的评测通常只看整体准确率一个模型只要在大部分样本上猜中高频答案分数就会很高。于是出现了一个尴尬的现象模型在榜单上排名靠前但换一组分布不同的图片后性能立刻崩塌。这也是为什么“高分”和“真正在看图”不能直接画等号。1.2 一个微型例子反事实样本立刻暴露“记答案”假设训练集中大量出现这类样本问题图里的香蕉是什么颜色答案黄色。模型很容易学会一个捷径看到“香蕉”和“颜色”就输出“黄色”。如果给出一张已经腐烂发黑的香蕉图片正确答案是“黑色”或“棕色”依赖记忆的模型仍然可能输出“黄色”。这里要判断模型是在看图还是记答案不需要特别复杂的指标只要对图像做反事实修改看答案是否跟着变化。如果答案不变说明模型并没有真正使用图像证据。这个小例子说明准确率只能衡量“在当前数据集上是否答对”不能衡量“是否在正确推理”。1.3 只看准确率的盲区为什么需要归因和定位为了判断模型是否在记忆常见做法是看注意力图或归因分数注意力图反映模型在回答问题时看的是图像哪个区域。归因分数反映每个输入 token 或像素对最终答案的贡献。中间层激活的梯度可以反映哪些通道参与了最终决策。只靠准确率无法判断模型是否“看对地方”。这也是 ReKey 这一类方法的前提如果想让模型真正看图先要知道模型当前把注意力放在了哪里然后只修改决定答案的关键通道。2. 模型为什么更容易“记答案”语言先验、注意力塌缩与长尾分布2.1 语言先验问题文本本身就是一条捷径VQA 模型可以从问题文本中学到很强的先验。例如问题模板训练集可能的高频答案风险图里有没有__?没有忽略小目标这是什么颜色?红色 / 黄色与具体对象无关图里有几只__?2忽视数量变化这个人在这做什么?跑步依赖动作模板当问题里出现常见的动词、名词、疑问词组合时模型可以不用认真编码图像只凭文本模式就估计一个合理答案。如果这个估计在训练集上频繁命中损失函数不会强制模型去看图像于是语言先验被保留下来。2.2 注意力塌缩视觉证据没有被真正使用在 Transformer 架构里跨模态注意力会计算问题 token 与图像 token 的相似度。理想情况下问题中的“香蕉”应该和图像中的香蕉区域形成高注意力权重。但训练过程中如果文本先验已经足够降损失模型就没有动力去优化图像注意力权重。时间一长图像 token 的注意力权重可能呈现两种异常均匀分布所有图像区域都被低权重扫过集中于背景、文字说明或高频对象的区域而不是问题真正指向的区域。这就是“注意力塌缩”。视觉编码器可能已经提取了丰富特征但后续层根本不使用它们模型本质上退化成一个文本分类器。2.3 长尾答案分布低频正确答案被训练损失压制VQA 数据集的答案分布并不均匀。常见答案在训练样本里占比高模型初期就会优先学到这些答案。交叉熵损失对每个样本一视同仁但梯度更新会被高频类别主导。低频答案往往需要更细的视觉证据才能判断。例如区分“跑车”和“轿车”需要看轮廓和细节而“有没有人”这种问题只需要检测到人形区域。由于高频答案梯度占优低频正确答案所依赖的网络子结构得不到充分更新。最终模型变成一个“高频答案分类器”。2.4 三个“作弊”通道叠加后的表现语言先验、注意力塌缩、长尾分布不是互斥的它们会叠加出现问题文本上有强先验模型“感觉知道”答案图像区域被弱注意模型“随便看看”高频答案主导梯度模型“只学常见项”。这三个现象共同造成 VQA 模型的高分失真。ReKey 要治理的是这些通道交汇的位置也就是“到底哪一部分中间层参与编造了答案”。3. ReKey 的核心思想定位“决定答案的那一小块”3.1 从“全量微调”到“定点更新”传统微调会调整所有层的参数包括那些与记忆偏差无关、甚至负责通用视觉感知的层。全量微调有两个问题参数量大训练和存储成本高容易在目标任务上过拟合破坏预训练模型已经学到的通用视觉能力。ReKey 的策略不同先通过归因手段识别哪些中间表示对当前答案的预测贡献最大然后只更新这些表示背后的参数。这样参数数量少、过拟合风险低同时保留模型已经具备的视觉理解能力。3.2 “Key”在 ReKey 中的含义Transformer 里有 QQuery、KKey、VValue。在视觉问答场景中问题可以理解为 Query它要去图像中检索相关信息图像特征作为 Key 和 Value。如果某个视觉 token 与当前问题语义相关它的 Key 就应该和 Query 高度匹配。VQA 模型如果记忆答案通常意味着匹配关系出了问题模型总是匹配到高频答案区域而不是问题指向的目标。ReKey 的核心动作就是重新调整这些关键的 Key 向量让 Query 更准确地检索到真实证据因此叫 ReKey。这里的“一小块”可以是一个视觉 token、一个注意力头也可以是一个低秩子空间关键是范围远小于全量参数。3.3 ReKey 的三步流程ReKey 的落地可以拆成三个步骤步骤目标常见方法1. 影响定位找到哪些中间表示对最终答案贡献最大梯度范数、注意力分数、积分梯度2. 子集筛选选出 top-k 个关键块按影响分数排序组合成参数掩码3. 稀疏微调只更新关键块对应参数其余冻结requires_grad、优化器参数过滤影响定位是 ReKey 最重要的环节。定位不准后续筛选和微调都会失效。因此实际项目中通常不会只使用单个样本的梯度而是从一个验证 batch 或多个训练 batch 中统计影响分数再做平均。3.4 为什么只更新“关键块”能缓解“记答案”“记答案”的决策根据并不均匀分布在整个网络里。实践中的观察是只有少数注意力头和少数视觉 token 对错误答案有主导作用。如果把它们找出来并重新标定就能破坏语言先验的捷径。ReKey 相当于做了一次外科手术定位到病灶只处理病灶。模型在保留原有视觉语义的同时被迫改变对问题模板和图像证据的匹配方式从而减少对高频答案的依赖。这也是它和全量微调、LoRA 等参数高效微调方法最大的区别ReKey 不只是“减少参数”而是“定向干预”。4. 从概念到工程准备环境并搭一个 VQA 基线4.1 环境依赖要复现 ReKey 的最小实验需要准备以下环境Python 3.8 以上PyTorch 2.xTransformers 库用于加载文本编码器和部分多模态模型Pillow用于图像读取与预处理datasets用于管理小规模训练集tqdm用于展示训练进度。可以创建一个虚拟环境并安装依赖python -m venv rekey_env source rekey_env/bin/activate pip install torch transformers datasets pillow tqdm版本号建议在安装时确认不要直接锁定一个不可验证的旧版本。如果使用 GPU还需要额外确认 CUDA 版本与 PyTorch 匹配。4.2 数据准备带反事实标注的小规模训练集ReKey 的验证需要“反事实测试集”也就是同一类问题、同一类对象但图像内容发生改变。为了让实验快速跑通可以自己构造一个小型数据集。数据格式可以是 JSON{ train: [ { image_id: train_0001, image_path: images/train_0001.jpg, question: What color is the banana?, answer: yellow, counterfactual_answer: black } ], val: [ { image_id: val_0001, image_path: images/val_0001.jpg, question: What color is the apple?, answer: red, counterfactual_answer: green } ] }counterfactual_answer字段是关键。它表示同一问题在图像内容被修改后应该输出的正确答案。这个字段将在后续评估中使用用来判断模型是否真的根据图像改变答案。实际项目中反事实样本可以由人工标注也可以通过图像编辑工具自动生成。4.3 最小 VQA 模型结构为了说明 ReKey 的机制先搭一个最小的 VQA 基线模型。这个模型由三部分组成视觉编码器把图片编码成向量文本编码器把问题编码成向量融合层和分类头输出答案分布。以下是简化版 PyTorch 代码import torch import torch.nn as nn from transformers import BertModel, ResNetModel class MiniVQA(nn.Module): def __init__(self, num_answers256): super().__init__() self.vision ResNetModel.from_pretrained(microsoft/resnet-18) self.text BertModel.from_pretrained(bert-base-uncased) self.proj nn.Linear(512 768, 512) self.classifier nn.Linear(512, num_answers) def forward(self, images, input_ids, attention_mask): image_feat self.vision(pixel_valuesimages).pooler_output text_feat self.text(input_idsinput_ids, attention_maskattention_mask).pooler_output fused torch.cat([image_feat, text_feat], dim-1) fused torch.relu(self.proj(fused)) logits self.classifier(fused) return logits这里只是示意。实际项目中要替换为更完整的视觉编码器、更大的答案词表以及处理“自由文本答案”的解码器。如果使用 Qwen3-VL 这类多模态大模型可以将这个最小结构替换成对应的跨模态注意力和语言模型头但 ReKey 的定位逻辑仍然适用。4.4 基线评估先确认模型确实出现“记答案”行为在应用 ReKey 之前先跑一次基线评估。评估不能只看整体准确率需要按问题类型拆分def evaluate_basic(model, dataloader): model.eval() total 0 correct 0 stats {} with torch.no_grad(): for batch in dataloader: logits model(**batch) pred logits.argmax(dim-1) match pred.eq(batch[answer_ids]) for i, qtype in enumerate(batch[question_type]): stats.setdefault(qtype, [0, 0]) stats[qtype][0] match[i].item() stats[qtype][1] 1 total len(match) correct match.sum().item() print(整体准确率:, correct / total) for qtype, (correct_count, count) in stats.items(): print(qtype, correct_count / count)如果“颜色类”问题准确率很高但后续反事实测试准确率很低说明模型很可能在记高频颜色答案而不是在看图中的具体颜色。这个结论就是启动 ReKey 的依据。5. 实现 ReKey只更新决定答案的那一小块参数5.1 基于梯度的影响定位ReKey 的第一步是定位关键参数。这里用梯度范数作为影响分数简单且容易实现。思路对一批样本前向传播计算损失反向传播得到每个参数的梯度然后用梯度范数衡量该参数对当前答案的影响程度。import torch import torch.nn.functional as F def compute_parameter_importance(model, dataloader, num_batches4): model.zero_grad() importance {} for i, batch in enumerate(dataloader): if i num_batches: break logits model(**batch) loss F.cross_entropy(logits, batch[answer_ids]) loss.backward() for name, param in model.named_parameters(): if param.grad is not None: norm param.grad.detach().norm().item() importance[name] importance.get(name, 0.0) norm model.zero_grad() for name in importance: importance[name] / num_batches return importance这里指定num_batches是为了降低单个 batch 的随机性。影响分数会按参数名记录。如果两个 batch 的定位结果差异过大需要增加采样 batch 的数量。5.2 筛选 top-k 关键块并生成参数掩码拿到每个参数的影响分数后按分数从高到低排序选出前 k 个参数名作为“关键块”。def select_topk_params(importance_dict, k20): sorted_items sorted(importance_dict.items(), keylambda x: x[1], reverseTrue) selected [name for name, _ in sorted_items[:k]] return selected这里k指的是参数名数量不是参数量。实际工程中更合理的做法是控制“可训练参数量占总参数量”的比例例如选择 top 5% 或 top 10%。筛选之后生成掩码的方式通常有两种对选中的参数设置requires_gradTrue其余设为False在优化器的参数列表里只传入选中的参数。第二种方式更干净因为它只影响优化器不会改变模型内部的 forward 路径。推荐使用第二种方式。5.3 冻结无关参数只更新关键子集下面给出冻结参数并构建优化器的示例def prepare_rekey_optimization(model, selected_param_names, lr1e-4): trainable_params [] for name, param in model.named_parameters(): if name in selected_param_names: trainable_params.append(param) else: param.requires_grad False optimizer torch.optim.Adam(trainable_params, lrlr) return optimizer这样只有选中的参数拥有优化器状态和梯度更新。其余参数虽然在前向中仍参与计算但不会更新。需要注意很多模型里存在 LayerNorm 的 gamma/beta 和各类 bias这些参数虽然名字不在影响力 top-k 中但在训练中承担偏移校准作用。如果完全不更新它们模型可能难以稳定训练。实践中可以强制把 bias 和 LayerNorm 参数也加入可训练集合。5.4 训练循环与验证ReKey 的训练循环和普通微调没有本质区别但验证逻辑要调整。下面是一个简化训练循环def train_rekey(model, optimizer, train_loader, val_loader, epochs3): for epoch in range(epochs): model.train() for batch in train_loader: logits model(**batch) loss F.cross_entropy(logits, batch[answer_ids]) optimizer.zero_grad() loss.backward() optimizer.step() val_acc evaluate_counterfactual(model, val_loader) print(fepoch {epoch}, val counterfactual acc: {val_acc:.4f})这里的evaluate_counterfactual会在下一章解释。ReKey 的核心关注点是反事实准确率能否提升而不是简单的训练集准确率。6. 验证效果怎么判断模型真的在“看图”而不是“记答案”6.1 验证维度ReKey 是否有效不能只看整体准确率。建议从四个维度验证维度指标方法整体准确率Acc常规验证集反事实能力Counterfactual Acc修改图像后答案是否正确语言先验依赖Question-only Acc只输入问题不进图像注意力正确性Attention Overlap注意力区域与标注目标区域的重叠度只输入问题不进图像也就是 question-only baseline能直接反映模型对文本先验的依赖程度。ReKey 后的模型question-only 准确率应该下降说明它不再单靠问题文本也能答对而是依赖图像。6.2 反事实测试脚本反事实测试的核心逻辑是对同一问题构造两张不同的图模型应该给出不同且正确的答案。def evaluate_counterfactual(model, dataloader): model.eval() correct 0 total 0 with torch.no_grad(): for batch in dataloader: # batch 中包含原始图和反事实图 original_logits model(imagesbatch[images], input_idsbatch[input_ids]) cf_logits model(imagesbatch[cf_images], input_idsbatch[input_ids]) original_pred original_logits.argmax(dim-1) cf_pred cf_logits.argmax(dim-1) original_correct original_pred.eq(batch[answer_ids]) cf_correct cf_pred.eq(batch[cf_answer_ids]) correct (original_correct cf_correct).sum().item() total len(batch[answer_ids]) return correct / total这个指标同时要求模型在原始图上答对在反事实图上也要答对。如果模型只是记忆了训练集中的高频答案它难以同时满足两个条件。6.3 语言先验测试语言先验测试可以这样设计把“What color is the banana?” 改为 “What color is the thing?”保持图像不变看模型答案是否因为问题缺少“banana”而剧烈变化。如果模型对“thing”这种无明确指向的问题也能基于图像颜色输出正确答案说明它开始依赖视觉证据。ReKey 的目标之一就是让模型在文本信息不足时仍然能看图像。6.4 注意力可视化与归因确认最后用视觉归因方法确认注意力区域是否变化。最简单的做法是取跨模态注意力权重平均后缩放到原图大小再和目标区域框计算 IoU。如果 ReKey 前的注意力集中在背景或高频区域ReKey 后应当更集中于问题所指对象的区域。这一步虽然不贡献数值指标但对排查“模型到底在看哪里”非常有效。7. 常见问题与排查路径7.1 冻结参数后 BatchNorm 仍在更新现象模型训练不稳定或验证集表现异常。原因BatchNorm 的 running statistics 不会因为requires_gradFalse而停止更新。即使所有参数都被冻结BatchNorm 在前向传播时仍会根据当前 batch 更新均值方差。处理方式在冻结大部分参数时确认需要更新统计信息的 BatchNorm 是否保留如果不想更新将模型切到eval()模式或者在构建模型时设置track_running_statsFalse。检查方式打印 BatchNorm 层的running_mean训练前后对比是否变化。7.2 梯度定位不稳定关键块每次都不一样现象同一数据和超参数两次定位得到不同的 top-k 参数集合。可能原因梯度范数本身方差大采样 batch 数量太少模型没有收敛到稳定区域。检查方式固定随机种子重复两次影响定位计算所选参数集合的重叠比例。处理方式增加采样 batch 数量或者改用积分梯度、注意力时间序列平均等更稳定的归因方法。如果重叠比例仍然很低应该先稳定基线模型再应用 ReKey。7.3 top-k 选得太大或太小效果不升反降这是 ReKey 最需要调的超参数。给出一个参考表现关键块比例可能表现常见问题过小例如 0.1%模型几乎不更新参数被冻结过多loss 不下降适中例如 5%-10%loss 下降反事实测试提升需要消融确定具体数值过大例如 50% 以上接近全量微调过拟合风险高失去 ReKey 意义推荐做法是画一条“反事实准确率随关键块比例变化”的曲线选择峰值附近较小的 k 值。7.4 训练后整体准确率没有下降但反事实准确率也没提升这可能是验证集设计问题。如果验证集和训练集同分布ReKey 不一定能体现优势。需要专门构建分布外验证集例如用不同数据集用图像编辑生成的新对象组合用只含文本先验的困难样本。另外还要检查训练损失是否下降。如果损失没有下降说明关键块定位过窄或者冻结参数过多。此时可以适当扩大 top-k或强制加入 LayerNorm 和 bias。7.5 显存不足ReKey 冻结了大量参数可以降低优化器状态和梯度存储但前向传播和反向传播的中间激活仍然占用显存。处理方式# 启用 gradient checkpointing model.gradient_checkpointing_enable()同时可以减小 batch size、使用混合精度训练pip install torch.cuda.amp或者使用显存优化工具例如 DeepSpeed ZeRO。对于显存不是瓶颈的服务器直接使用单卡 A100 或 4090 也能跑通小规模验证。8. 最佳实践与扩展方向8.1 什么时候适合用 ReKey什么时候不适合适合 ReKey 的场景已经有一个不错的预训练 VQA 模型但希望缓解记忆偏差微调预算有限无法承受全量参数更新需要可解释性想知道模型修改了哪部分有少量反事实标注样本能够验证干预效果。不适合的场景从头训练一个大型多模态模型ReKey 不是训练范式没有归因工具也没有留存梯度信息的内部代码数据量太少无法稳定定位关键块。8.2 ReKey 与 LoRA、Prefix Tuning 的对比方法更新范围可解释性显存开销适用场景全量微调所有参数弱高数据充足算力充足LoRA低秩矩阵中低通用参数高效微调Prefix Tuning前缀 token 参数中低语言模型生成任务ReKey影响分数最高的关键块强可调定位记忆偏差并干预ReKey 和 LoRA 不是互斥的。可以先使用 ReKey 定位关键注意力头再在这个子集上使用 LoRA 进行适配。这样既能减少参数又能保留可解释性。8.3 与多模态大模型结合当前多模态大模型例如 Qwen3-VL 等同样存在“答案记忆”风险。大规模预训练后模型会记住大量常见视觉知识和语言关联。在特定 VQA 任务上如果只做全量指令微调仍然可能强化高频答案。ReKey 可以尝试作为大模型微调的后处理过程先通过梯度或注意力归因找出哪些跨模态注意力头对错误答案贡献最大然后只在这些头上做低秩微调。这样既能改变模型对特定问题模板的响应又不破坏大模型已经具备的开放世界知识。8.4 可复用清单在把 ReKey 应用于生产或研究时建议逐项检查是否准备了反事实验证集而不是只看整体准确率影响定位是否在多个 batch 上稳定而不是单次采样参数掩码是否覆盖应该更新的 LayerNorm 和 bias是否对 top-k 比例做了消融实验是否同时计算 question-only 准确率确认语言先验下降是否保存了冻结参数列表和训练日志方便复现是否保留全量微调或 LoRA 基线用于对比是否检查了 BatchNorm running statistics 的更新行为是否记录注意力可视化结果确认模型真的在看目标区域是否在分布外数据上做了测试而不是只依赖原始验证集。这份清单可以直接用于代码评审或发布前的质量检查。8.5 下一步学习建议如果第一次接触 ReKey建议先不要直接套用大模型。可以构造一个几千样本的小型 VQA 数据集故意让训练集中某一类答案占多数训练一个小模型复现“记答案”现象。然后实现基于梯度的 ReKey比较全量微调、LoRA 和 ReKey 在反事实测试上的差异。这个过程能帮助你理解定位、筛选、冻结这三步的实际影响。之后再迁移到 Qwen3-VL 之类的多模态模型上观察它的跨模态注意力头分布用 ReKey 的思路做定向干预。VQA 的高分是否可信最终取决于验证方式是否足够严格。ReKey 给出的不是万能答案而是一种更接近“外科手术”的干预路径找到决定答案的那一小块只更新那一小块然后通过反事实测试证明模型真的在看图。