十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

模态不平衡成因与对策:结合特权信息与模态平衡的训练方法

模态不平衡成因与对策:结合特权信息与模态平衡的训练方法 多模态大模型MLLM在训练过程中有一个非常隐蔽、但直接影响最终效果的问题叫做模态不平衡。很多同学训练视觉问答模型时会发现loss 下降很快准确率却一直卡在某个平台期仔细观察之后才发现模型根本没有“认真看”图像而是靠文本先验在猜答案。这篇文章就围绕这个问题展开梳理模态不平衡的成因并介绍一种结合模态平衡策略与**特权信息Privileged Information**的训练思路用来提升模型的视觉推理能力。全文包含完整可运行的 PyTorch 示例代码适合已经有一定深度学习基础、正在折腾多模态大模型训练的读者。1. 背景与核心概念1.1 多模态大模型为什么会出现“模态不平衡”多模态大模型Multimodal Large Language Model简称 MLLM通常由三部分组成视觉编码器、文本编码器、跨模态融合层或语言模型解码器。代表性的工作包括 CLIP、BLIP、LLaVA、Qwen-VL 等。这类模型的目标是让模型同时理解图像和文本完成图文检索、视觉问答、图像描述、视觉推理等任务。但在训练过程中我们经常会发现一个现象文本侧的信息过于强势视觉侧的特征被压制。这个问题可以用一个很直观的方式描述你把图像输入换成一个随机噪声图模型的预测结果几乎不变说明模型已经“学会”了不看图。图像信息没有真正参与决策多模态退化成了“单模态文本模型”。模态不平衡的成因是多方面的语言先验太强模型本身是从大规模纯文本语料中预训练出来的文本编码器对答案类别的覆盖能力远超视觉编码器。模型发现只要看问题文本就能猜个大概视觉输入对降低 loss 的贡献不大。梯度更新不均衡视觉编码器通常使用较深的 CNN 或 ViT文本编码器路径更短、梯度更容易回传。在共享优化器的情况下文本侧参数获得的梯度更新更充分视觉侧参数长期“吃不饱”。融合层结构偏向文本很多 Cross-Attention 融合层直接用文本特征作为 Query视觉特征作为 Key/Value。如果视觉特征本身没有对齐到与文本特征相似的空间Attention 的权重很快就会偏向文本 Query 自身。数据分布差异训练数据里文本描述和问题天然带有答案线索模型很容易走捷径。模态不平衡带来的直接结果是模型在评测集上看起来“还行”但一旦遇到训练分布之外的图像、新颖的组合、指代不清的问题性能就会明显下降。因为模型没有真正学会视觉推理它只是把问题文本映射到了回答文本。1.2 特权信息训练阶段的“额外参考书”特权信息Privileged Information来自 Learning Using Privileged InformationLUPI范式核心思想是训练阶段可以为模型提供一些推理阶段不会出现的额外信息模型通过蒸馏或辅助监督的方式把这些额外信息中学到的能力压缩回主模型从而让主模型在推理时即使没有这些信息也能表现得更好。在多模态大模型中常见的特权信息包括图像的详细描述文本caption。OCR 文本。目标检测框和类别标签。深度图、分割掩码。人工标注的推理步骤或结构化知识。例如训练一个视觉问答模型时输入是一张图片和一个问题“图中有几只鸟”常规输出是类别数或答案。如果我们在训练时额外给一句描述“图中有 3 只红色的小鸟站在树枝上”这句话就是特权信息。模型需要学会把图像内容和这段描述对齐从而获得更强的视觉语义表征。推理时我们只输入图片和问题模型也能输出正确答案。特权信息本质上解决的是“视觉监督信号不足”的问题。它用一段与图像强相关的额外文本把视觉特征拉到一个更容易被下游任务使用的语义空间相当于给视觉编码器配了一位“私教”。1.3 模态平衡与特权信息如何互相配合模态平衡解决的是“模型看不见图”的问题特权信息解决的是“视觉特征缺乏强监督”的问题。两者可以叠加起来模态平衡策略调整视觉塔和文本塔的梯度更新幅度防止文本侧过度主导。特权分支用额外描述文本提供更强的视觉对齐信号帮助视觉编码器在有限的梯度预算里学到更有区分度的特征。推理时去掉特权输入主分支通过蒸馏保留特权分支学到的能力保证推理效率不受影响。接下来我会先用原理拆解的方式说明这两个技术点的内在逻辑再给出一个可运行的工程示例让你能直接复现整个训练流程。2. 环境准备与版本说明本节给出示例运行环境。不同环境之间可能存在依赖版本差异请以实际项目为准。2.1 基础环境示例项目使用 Python PyTorch Hugging Face Transformers推荐环境如下操作系统Linux / macOS / Windows建议 Linux Python3.9 或 3.10 PyTorch2.x2.0 及以上 Transformers4.30 及以上建议 4.36 以上 CUDA11.8 或 12.1如使用 GPU 显存建议 16GB 以上使用 ViT-Base Bert-Base 时如果没有 GPU也可以用小 batch size 在 CPU 上运行但训练速度会非常慢建议至少使用单张 8GB 显存以上的 GPU。2.2 安装依赖创建虚拟环境并安装依赖python -m venv venv source venv/bin/activate pip install --upgrade pip pip install torch transformers如果训练过程需要更完善的日志和指标管理可以额外安装pip install tensorboard datasets2.3 示例项目结构本文的示例项目结构如下mllm_modality_balance/ ├── model.py ├── train.py └── README.md代码思路是用 ViT 作为视觉编码器、BERT 作为文本编码器构造一个简化版视觉问答模型并在训练中同时加入模态平衡和特权信息机制。3. 核心原理拆解3.1 模态不平衡背后的梯度问题从优化角度看模态不平衡本质上是不同模态编码器的梯度更新幅度差异过大。假设总损失为L L_main L_aux其中L_main来自主分类头L_aux可能来自对比损失、辅助损失等。反向传播时视觉塔和文本塔各自接收到的梯度范数并不一致。文本塔通常更容易拿到较大的梯度。原因之一是跨模态融合层的梯度回传路径不同。视觉特征从图像经过多层 Transformer 编码再通过 Cross-Attention 融合文本特征则可以直接通过残差连接和自注意力更新路径更短。另一个原因是预训练权重的差异BERT 已经具备很强的语义推理能力而视觉塔的 ViT 在图文任务上还需要大量调整。如果两者在同一个优化器里平等更新视觉塔的更新很容易被文本塔“掩盖”。在实际训练中我们建议定期统计visual branch 参数的平均梯度范数。text branch 参数的平均梯度范数。两者的比值。如果比值长期偏离 1说明存在模态不平衡需要处理。3.2 模态平衡的几种实现思路模态平衡没有唯一的标准实现常见策略包括损失加权为视觉相关损失分配更高权重比如增加对比损失的权重。思路简单但权重需要人工调。梯度缩放在反向传播后按比例缩放不同分支的梯度。这是本文示例采用的方式可以动态计算缩放系数。学习率错峰不同分支使用不同的学习率视觉塔学习率高一些文本塔学习率低一些。冻结策略先冻结文本塔训练视觉塔和融合层再统一微调。投影层调整重点调整视觉投影层和文本投影层的初始化方式使两者初始输出分布更接近。Token 级 Reweight在 Transformer 层内对视觉 token 的注意力权重做额外增强让图像信息在自注意力中有更大的影响力。本文示例采用梯度缩放方案因为它的计算方式直观而且能在每个 step 动态适应训练过程。需要说明的是这是一个教学级简化版本生产项目中可能需要被更复杂的策略替代或补充。3.3 特权信息的训练范式特权信息训练的目标不是让模型在推理时继续依赖额外输入而是在训练时借助额外信息提升视觉表示质量。具体实现上本文采用一个“双分支 蒸馏”结构主分支输入图像 问题文本输出预测 logits。特权分支输入图像 问题文本 特权描述文本输出预测 logits。蒸馏损失把特权分支的 logits 作为软标签让主分支的 logits 向它学习。推理阶段只使用主分支。这样模型可以在训练时获得额外的语义帮助推理时却不增加任何输入成本和延迟。对于视觉推理任务特权分支通常比主分支更容易收敛因为它拿到了更多“答案线索”。通过蒸馏主分支不需要在推理阶段使用这些线索也能学到近似能力。这正是特权信息提升视觉推理性能的关键机制。4. 完整实战案例下面我们用一个简化版视觉问答任务按步骤实现“模态平衡 特权信息”的完整训练流程。示例代码可以复制后直接运行重点是理解实现思路。4.1 构建 model.py模型结构首先创建model.py定义视觉编码器、文本编码器、跨模态融合层以及主模型。# 文件路径mllm_modality_balance/model.py import torch import torch.nn as nn from transformers import ViTModel, BertModel class VisualEncoder(nn.Module): def __init__(self, model_namegoogle/vit-base-patch16-224-in21k, out_dim768): super().__init__() self.vit ViTModel.from_pretrained(model_name) self.proj nn.Linear(self.vit.config.hidden_size, out_dim) def forward(self, pixel_values): last_hidden self.vit(pixel_values).last_hidden_state cls_feat last_hidden[:, 0, :] return self.proj(cls_feat) class TextEncoder(nn.Module): def __init__(self, model_namebert-base-uncased, out_dim768): super().__init__() self.bert BertModel.from_pretrained(model_name) self.proj nn.Linear(self.bert.config.hidden_size, out_dim) def forward(self, input_ids, attention_mask): output self.bert(input_idsinput_ids, attention_maskattention_mask) pooled output.pooler_output return self.proj(pooled) class CrossModalFusion(nn.Module): def __init__(self, hidden_size768, num_heads8): super().__init__() self.cross_attn nn.MultiheadAttention(hidden_size, num_heads, batch_firstTrue) self.norm1 nn.LayerNorm(hidden_size) self.mlp nn.Sequential( nn.Linear(hidden_size * 2, hidden_size * 4), nn.GELU(), nn.Dropout(0.1), nn.Linear(hidden_size * 4, hidden_size), ) self.norm2 nn.LayerNorm(hidden_size) def forward(self, img_feat, text_feat): query text_feat.unsqueeze(1) key_value img_feat.unsqueeze(1) attn_out, _ self.cross_attn(query, key_value, key_value) out self.norm1(query attn_out) out out.squeeze(1) out torch.cat([out, text_feat], dim-1) out self.mlp(out) return self.norm2(out) class MLLMWithBalance(nn.Module): def __init__(self, num_classes10): super().__init__() self.visual_encoder VisualEncoder() self.text_encoder TextEncoder() self.fusion CrossModalFusion() self.classifier nn.Linear(768, num_classes) # 特权分支复用文本编码器使用独立的融合层和分类头 self.privilege_fusion CrossModalFusion() self.privilege_classifier nn.Linear(768, num_classes) def forward(self, pixel_values, input_ids, attention_mask, privilege_idsNone, privilege_maskNone): img_feat self.visual_encoder(pixel_values) text_feat self.text_encoder(input_ids, attention_mask) fused self.fusion(img_feat, text_feat) logits self.classifier(fused) if privilege_ids is not None: priv_text_feat self.text_encoder(privilege_ids, privilege_mask) priv_fused self.privilege_fusion(img_feat, priv_text_feat) priv_logits self.privilege_classifier(priv_fused) return logits, priv_logits, img_feat, text_feat return logits, None, img_feat, text_feat这里有几个设计细节值得说明VisualEncoder和TextEncoder各自带一层proj把特征统一投影到 768 维空间方便后续融合。CrossModalFusion中文本特征作为 Query视觉特征作为 Key/Value这是目前多模态模型中常见的 Cross-Attention 设计。特权分支复用了同一个text_encoder因为问题文本和特权描述文本本质上是同一种模态。这样既减少了参数量也让主分支和特权分支共享底层语义编码能力。num_classes需要根据你的任务标签数量调整。4.2 构建 train.py数据集与训练逻辑接下来是核心的train.py。示例使用一个模拟的视觉问答数据集输入为随机图像张量、问题文本、特权描述文本输出为类别标签。你可以在真实项目中替换为实际图片数据集。# 文件路径mllm_modality_balance/train.py import argparse import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader, Dataset from transformers import BertTokenizer from model import MLLMWithBalance class SimulatedVQADataset(Dataset): def __init__(self, size200, num_classes10, max_text_len24): self.size size self.num_classes num_classes self.tokenizer BertTokenizer.from_pretrained(bert-base-uncased) self.max_text_len max_text_len self.answers [fclass_{i} for i in range(num_classes)] def __len__(self): return self.size def __getitem__(self, idx): pixel_values torch.randn(3, 224, 224) label idx % self.num_classes question fwhat is in image {idx}? privilege fthere is one object, and the category is {self.answers[label]} q_enc self.tokenizer( question, paddingmax_length, truncationTrue, max_lengthself.max_text_len, return_tensorspt, ) p_enc self.tokenizer( privilege, paddingmax_length, truncationTrue, max_lengthself.max_text_len, return_tensorspt, ) return { pixel_values: pixel_values, input_ids: q_enc[input_ids].squeeze(0), attention_mask: q_enc[attention_mask].squeeze(0), privilege_ids: p_enc[input_ids].squeeze(0), privilege_mask: p_enc[attention_mask].squeeze(0), label: torch.tensor(label, dtypetorch.long), } def collate_fn(batch): return { pixel_values: torch.stack([item[pixel_values] for item in batch]), input_ids: torch.stack([item[input_ids] for item in batch]), attention_mask: torch.stack([item[attention_mask] for item in batch]), privilege_ids: torch.stack([item[privilege_ids] for item in batch]), privilege_mask: torch.stack([item[privilege_mask] for item in batch]), labels: torch.stack([item[label] for item in batch]), } def contrastive_loss(img_feat, text_feat, temperature0.07): img_feat F.normalize(img_feat, dim-1) text_feat F.normalize(text_feat, dim-1) logits img_feat text_feat.t() / temperature labels torch.arange(logits.size(0)).to(logits.device) loss ( F.cross_entropy(logits, labels) F.cross_entropy(logits.t(), labels) ) / 2 return loss def compute_grad_norms(model): visual_norm 0.0 text_norm 0.0 for name, param in model.named_parameters(): if param.grad is None: continue if visual_encoder in name: visual_norm (param.grad.norm(2).item() ** 2) elif text_encoder in name: text_norm (param.grad.norm(2).item() ** 2) return visual_norm ** 0.5, text_norm ** 0.5 def balance_modality_gradients(model): v_norm, t_norm compute_grad_norms(model) eps 1e-8 total v_norm t_norm eps v_scale total / (2 * v_norm eps) t_scale total / (2 * t_norm eps) # 防止极端情况下梯度爆炸 v_scale min(v_scale, 10.0) t_scale min(t_scale, 10.0) for name, param in model.named_parameters(): if param.grad is None: continue if visual_encoder in name: param.grad * v_scale elif text_encoder in name: param.grad * t_scale return v_norm, t_norm torch.no_grad() def evaluate(model, dataloader, device, noise_imageFalse): model.eval() correct 0 total 0 for batch in dataloader: pixel_values batch[pixel_values].to(device) if noise_image: pixel_values torch.randn_like(pixel_values) input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) labels batch[labels].to(device) logits, _, _, _ model(pixel_values, input_ids, attention_mask, None, None) preds logits.argmax(dim-1) correct (preds labels).sum().item() total labels.size(0) return correct / total if total 0 else 0.0 def main(): parser argparse.ArgumentParser() parser.add_argument(--epochs, typeint, default3) parser.add_argument(--batch_size, typeint, default8) parser.add_argument(--lr, typefloat, default2e-5) parser.add_argument(--num_classes, typeint, default10) parser.add_argument(--data_size, typeint, default200) parser.add_argument(--device, typestr, defaultcuda if torch.cuda.is_available() else cpu) args parser.parse_args() device torch.device(args.device) train_dataset SimulatedVQADataset(sizeargs.data_size, num_classesargs.num_classes) train_loader DataLoader( train_dataset, batch_sizeargs.batch_size, shuffleTrue, collate_fncollate_fn, ) val_dataset SimulatedVQADataset(size80, num_classesargs.num_classes) val_loader DataLoader( val_dataset, batch_sizeargs.batch_size, shuffleFalse, collate_fncollate_fn, ) model MLLMWithBalance(num_classesargs.num_classes).to(device) optimizer torch.optim.AdamW(model.parameters(), lrargs.lr, weight_decay0.01) for epoch in range(args.epochs): model.train() total_loss 0.0 total_v_norm 0.0 total_t_norm 0.0 for batch in train_loader: pixel_values batch[pixel_values].to(device) input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) privilege_ids batch[privilege_ids].to(device) privilege_mask batch[privilege_mask].to(device) labels batch[labels].to(device) logits, priv_logits, img_feat, text_feat model( pixel_values, input_ids, attention_mask, privilege_ids, privilege_mask, ) loss_main F.cross_entropy(logits, labels) loss_priv F.cross_entropy(priv_logits, labels) with torch.no_grad(): soft_label F.softmax(priv_logits.detach() / 2.0, dim-1) loss_kl F.kl_div( F.log_softmax(logits / 2.0, dim-1), soft_label, reductionbatchmean, ) loss_cl contrastive_loss(img_feat, text_feat) loss loss_main 0.3 * loss_priv 0.3 * loss_kl 0.1 * loss_cl optimizer.zero_grad() loss.backward
返回列表