共计 3562 个字符,预计需要花费 9 分钟才能阅读完成。
为什么需要多模态大模型?
在 AI 领域,单模态模型(如纯文本 BERT 或纯图像 ResNet)已无法满足现实需求。比如医疗场景中,医生需要同时分析 CT 影像和患者病史文本;电商平台要处理商品图片与描述文字的关联检索。单模态模型存在三个核心痛点:

- 信息割裂 :X 光片和诊断报告分开处理会丢失关键关联
- 交互缺失 :用户用语言描述图片搜索时,单模态系统难以理解
- 表征局限 :文本和视觉特征空间不一致导致跨模态任务效果差
主流架构横评
| 模型 | 参数量 | 训练数据 | 支持模态 | 典型应用场景 |
|---|---|---|---|---|
| CLIP | 400M | 4 亿图文对 | 图像 + 文本 | 零样本分类 |
| BLIP | 224M | 1.29 亿图文对 | 图像 + 文本 | 图像描述生成 |
| Flamingo | 80B | 27 亿图文对 + 视频字幕 | 图像 / 视频 + 文本 | 开放式视觉问答 |
选择建议 :
– 轻量级任务选 BLIP
– 需要强泛化能力用 CLIP
– 视频场景考虑 Flamingo
手把手实现双塔模型
基础结构搭建
import torch
from transformers import BertModel, BertTokenizer
from torchvision.models import resnet50
class DualTower(torch.nn.Module):
def __init__(self):
super().__init__()
self.vision_encoder = resnet50(pretrained=True)
self.text_encoder = BertModel.from_pretrained('bert-base-uncased')
# 投影到共同特征空间
self.vision_proj = torch.nn.Linear(2048, 256)
self.text_proj = torch.nn.Linear(768, 256)
def forward(self, images, input_ids, attention_mask):
# 视觉特征提取
img_features = self.vision_encoder(images) # [bs, 2048]
img_emb = self.vision_proj(img_features) # [bs, 256]
# 文本特征提取
text_outputs = self.text_encoder(
input_ids=input_ids,
attention_mask=attention_mask
)
text_emb = self.text_proj(text_outputs.last_hidden_state[:,0,:]) # [bs, 256]
return img_emb, text_emb
关键组件实现
跨模态注意力层
class CrossAttention(torch.nn.Module):
"""query 来自模态 A,key/value 来自模态 B"""
def __init__(self, embed_dim=256, num_heads=8):
super().__init__()
self.multihead_attn = torch.nn.MultiheadAttention(
embed_dim=embed_dim,
num_heads=num_heads,
batch_first=True
)
def forward(self, query, key_value):
# query: [bs, seq_len_q, dim]
# key_value: [bs, seq_len_kv, dim]
attn_output, _ = self.multihead_attn(
query=query,
key=key_value,
value=key_value
)
return attn_output
对比损失计算
def contrastive_loss(img_emb, text_emb, temperature=0.07):
"""
img_emb/text_emb: [bs, 256] 已 L2 归一化
temperature: 控制困难样本的权重
"""
# 计算相似度矩阵
logits = torch.matmul(img_emb, text_emb.t()) / temperature # [bs, bs]
# 对角线元素是正样本对
labels = torch.arange(logits.shape[0], device=img_emb.device)
# 对称的对比损失
loss_img = torch.nn.functional.cross_entropy(logits, labels)
loss_text = torch.nn.functional.cross_entropy(logits.t(), labels)
return (loss_img + loss_text) / 2
实战避坑指南
数据预处理
- 图像处理 :
- 不要直接 resize 到固定尺寸,先保持长宽比缩放再中心裁剪
-
推荐使用 ImageNet 的均值和标准差归一化 ([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
-
文本处理 :
- 用与模型匹配的 tokenizer(如 BERT 模型用 BertTokenizer)
- 注意设置 max_length 并统一 padding 到相同长度
训练技巧
混合精度训练问题 :
# 正确使用姿势
scaler = torch.cuda.amp.GradScaler()
with torch.cuda.amp.autocast():
img_emb, text_emb = model(images, input_ids, attention_mask)
loss = contrastive_loss(img_emb, text_emb)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
遇到梯度爆炸时:
1. 检查 loss scaling 是否过大
2. 添加梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
进阶优化方案
LoRA 微调
# 在原线性层旁添加低秩适配器
class LoRALayer(torch.nn.Module):
def __init__(self, in_dim, out_dim, rank=8):
super().__init__()
self.lora_A = torch.nn.Linear(in_dim, rank, bias=False)
self.lora_B = torch.nn.Linear(rank, out_dim, bias=False)
def forward(self, x):
return self.lora_B(self.lora_A(x))
# 修改原模型参数
original_linear = model.text_encoder.encoder.layer[0].attention.self.query
model.text_encoder.encoder.layer[0].attention.self.query = torch.nn.Sequential(
original_linear,
LoRALayer(original_linear.out_features, original_linear.out_features)
)
分布式训练
# 初始化 DDP
import torch.distributed as dist
dist.init_process_group(backend='nccl')
model = torch.nn.parallel.DistributedDataParallel(model)
# 数据分片
sampler = torch.utils.data.distributed.DistributedSampler(dataset)
dataloader = DataLoader(dataset, batch_size=64, sampler=sampler)
开放性问题思考
-
时序对齐 :视频中的动作、音频波形和字幕文本如何建立时间轴对应关系?是否需要在 Transformer 中引入相对位置编码?
-
模态缺失 :当测试时缺少某个模态(如只有图片没有文本),模型如何保持性能?可否通过生成伪模态特征来解决?
-
评估指标 :现有的跨模态检索指标(如 Recall@K)是否能真实反映医疗等专业场景的需求?是否需要设计领域特定的评估体系?
总结
多模态模型不是简单的模块拼装,需要深入理解跨模态交互的本质。建议初学者:
1. 先用现成模型(如 HuggingFace 的 BLIP 实现)跑通 Pipeline
2. 重点分析 attention map 可视化结果
3. 逐步尝试修改模型结构
在实际业务中,往往需要根据数据特点调整模型结构。比如医疗影像需要更精细的局部特征提取,可以尝试将 CNN backbone 替换为更密集的 UNet 结构。
