微信场景多模态Embedding训练实战:从双塔模型到16G显存优化
2026/9/4 4:28:16 网站建设 项目流程

1. 微信场景下的多模态 Embedding,到底在解决什么问题

先说个我在实际落地中反复被问到的点:多模态 Embedding 不是把一张图和一段文字丢进同一个模型,然后拼个向量就完事。微信这种体量的产品,每天产生的是视频号里的短视频、公众号图文、朋友圈状态、小程序商品、聊天中随手发的图片和语音,这些内容之间天然存在跨模态的语义关联。用户可能搜一张靴子的图片,实际想找的是公众号里那篇穿搭攻略;用户可能在视频号刷到一个做菜片段,转头就想搜同款食材的图文教程。这种“用 A 模态表达、要 B 模态结果”的需求,在微信里太常见了。

所以多模态 Embedding 的核心目标,是把不同模态的内容统一映射到同一个向量空间里,让图片、文本、视频帧、音频片段之间的相似度可以计算、可以比较、可以检索。模型输出的向量质量直接决定了搜索排序、推荐召回、内容去重、素材库管理的效果上限。如果向量空间没对齐,后面做什么都白搭。

再说说适用人群。这篇文章不是写给只看理论的人看的,而是给那些真正要动手训练模型的人。我假设你有一定的 PyTorch 使用基础,跑过一些简单的 CV 或 NLP 模型,但还没有完整做过一个多模态对比学习的训练管线。文章里我会聊聊在 16G 显存这种常见的单卡环境下,怎么把模型训起来、怎么构造训练数据、怎么处理 batch 内负样本、怎么评估向量质量,以及我在实际项目中踩过的那些文档里不会写的坑。

训练多模态 Embedding 本质上是个工程问题,不完全是算法问题。你需要的不是一篇顶会论文复现,而是一套能落地、能迭代、能上线的方案。

2. 方案选型:为什么我不推荐一上来就上多模态大模型

训练这块,最常见的误区就是“哪个模型强用哪个”。网上热词里总看得到“16g显存多模态模型推荐”“多模态大模型”这类搜索词,很多人上来就想微调一个 Qwen2-VL 或者 LLaVA 这类大模型来产出向量。我的看法很直接:除非你的业务场景极其复杂,否则这是个性价比很低的做法。

2.1 Embedding 模型和生成式多模态模型的本质差异

生成式多模态大模型的目标是理解并生成内容,它的输出是 token 序列,内部表示虽然包含语义,但它的设计目标不是为了产生一个紧凑、可比较的向量。你要用它做检索,还得专门抽取倒数第二层的隐状态,然后在上面再接 pooling 层,再做 normalization,效果还不一定好。因为大模型在训练时并没有被显式地约束“相似的输入要映射到相近的向量”,它学到的表示是任务驱动的,不是度量驱动的。

而 Embedding 模型则不同,它的核心是度量学习。训练目标就是让正样本对(比如一张靴子图和它的标题文本)在向量空间里距离近,让负样本对距离远。整个网络结构的设计、损失函数的选择、样本构造的方式,全部围绕向量空间的对齐展开。

2.2 在 16G 显存下,我推荐的结构组合

我的建议是采用双塔结构(Two-Tower),而不是单塔融合。图像塔用 CLIP/ViT 系列的视觉编码器,文本塔用 BERT 或者更轻量的 Sentence-BERT 结构,两个塔的输出投影到同一个 512 维或 768 维的向量空间,然后通过对比学习拉近正样本。这套方案在显存占用上非常友好。

以我的实测经验,在 16G 显存的单卡上,图像塔用 ViT-B/16(大约 86M 参数),文本塔用 6 层 BERT(约 66M 参数),投影层用两层 MLP,输入图片分辨率 224x224,文本长度截断到 64 token,batch size 可以开到 256,配合混合精度训练,显存占用大概在 11G 到 13G 之间,还有余量做梯度累积。如果换成 ViT-L/14,batch size 就要降到 64 左右,否则会直接 OOM。

这里有一个很关键的设计:两个塔的输出向量,必须经过 L2 归一化再做相似度计算。原因不复杂,归一化之后,向量内积就等于余弦相似度,值域固定在 [-1, 1],配合温度系数缩放,训练会稳定很多。很多人训练的时候不归一化,向量模长分布不一致,损失值波动很大,那个坑我踩过。

2.3 什么时候才需要考虑更大的模型

有一个常见的热搜词是“多模态融合算法”,这个方向确实是前沿,但在训练 Embedding 模型的场景下,我建议你先问自己三个问题:业务是否需要细粒度的跨模态对齐(比如像素级别的图文对应)?现有双塔召回的指标是否已经到达瓶颈?你有多少标注数据?

如果数据量不足十万对,双塔结构已经足够;如果检索精度不够,优先排查的应该是负样本挖掘和难例挖掘,而不是换更大的模型。我在项目里见过太多人花两周去微调一个大模型,最后 recall 只提升了 0.5%,而调整损失函数里的温度系数和难负样本策略,一夜之间涨了 3 个点。

当然,如果是离线分析场景,或者想做多模态内容生成,那大模型有它的价值。但做检索召回,Embedding 模型是更精准的武器。

3. 数据构造:微信场景下的图文对、难负样本与滑动窗口策略

数据是整个训练管线里最耗时间、也最决定上限的部分。一个多模态 Embedding 模型的效果不是看网络结构有多花哨,而是看喂进去的图文对质量有多高。这点在微信的业务场景里尤其明显,因为用户产生的数据噪音非常大。

3.1 图文对的来源与清洗

对于微信生态,常见的数据来源有三类。第一类是公众号文章的配图和标题、段落文本。一篇文章里的多张配图和周围文本天然构成弱相关的图文对,这里需要做位置窗口过滤,一般来说取图片前后各 200 字以内的文本,超出窗口的文本大概率跟图没什么关系。第二类是视频号的视频帧和封面标题,抽帧频率建议每 2 秒一帧,然后和该片段的字幕对齐。第三类是朋友圈的图片和文案,这类数据噪音最大,因为用户发文案很随意,经常是“今天天气真好”配一张自拍,这种图文对信号非常弱,我一般会先过滤掉所有短文案(字数小于 15 的)。

数据清洗里面有一个容易被忽略的问题:重复图片。用户在不同时间发的同一张图(比如一张表情包反复用),在数据集中会形成大量近乎重复的样本。如果不做去重,模型会把这些图片的向量聚集在一起,导致检索结果被表情包霸屏。我处理的方法是从数据里随机抽样一部分图片,用预训练的 CLIP 模型离线提取特征,然后做局部敏感哈希(LSH)粗筛,再对候选集计算精确相似度,把相似度大于 0.95 的合并成一组,只保留最早出现的样本。这一步能去掉大约 20% 到 30% 的重复数据。

3.2 难负样本的挖掘机制

对比学习训练中,负样本的选择决定了模型能学到多精细的区分能力。如果负样本都是些“靴子和汽车”这种明显不相关的,模型学到的只是粗粒度的语义差异,它分不清“靴子和高跟鞋”的差别。这就是难负样本挖掘要解决的问题。

我的做法是三阶段挖掘。第一轮,先用一个预训练的 CLIP 模型作为初始化,跑一遍全量数据的向量检索,对每个正样本图文对,找出“文本描述相近但配图完全不同”的图文对,以及“图片相近但文本描述完全不同”的图文对,作为困难负样本。第二轮,在训练进行到一半时,冻结模型权重,重新对训练集做一次向量检索,用“当前模型认为难分但实际上是负样本”的样本,替换掉一部分上一轮的难负样本。第三轮是动态难负样本挖掘,每个 batch 内除了常规的 in-batch 负样本外,还会从一个小规模的难负样本池里在线采样,把它们也拉进这个 batch 参与计算。

这个机制说起来不复杂,但工程实现上有几个细节值得注意。难负样本池不能太大,我习惯控制在训练集总量的 2% 左右,太大了每次采样会有 IO 压力,太小了覆盖度不够。每训练 500 步左右就重新挖掘一次,不要每个 epoch 都挖,那会训练到过拟合困难的样本上。

3.3 滑动窗口滤波在训练数据预处理中的角色

热词里出现了“滑动窗口滤波模型”,这本来是一个信号处理的术语,但我觉得它用在图文序列对齐上非常贴切。在视频场景里,一段视频的不同帧和语音文本存在时序对应关系,如果用全局匹配,容易把文字和毫无关系的画面配对。我的做法是设定一个 3 秒到 5 秒的滑动窗口,在窗口内计算视频帧特征和文本特征的相似度,保留相似度最高的那对作为正样本,窗口滑动步长设置为 1 秒,允许相邻窗口之间存在重叠。这样做的好处是,即使视频画面切换很快,文本和画面的对齐关系也能保持稳定。

对于图文对来说,同样可以用滑动窗口的思路来处理长文本和长图片序列。当一篇公众号文章包含多张配图时,我会在文章中按照位置滑动一个 300 字左右的窗口,每个窗口只和距离最近的配图建立潜在匹配关系,然后再通过 CLIP 模型初筛,把分数达到阈值的窗口-配图对纳入训练集。这个策略简单有效,能大幅减少弱相关图文对的引入。

4. 训练细节与完整流程

数据准备好之后,就到了最核心的环节——模型训练。这一节我会按照实际操作的顺序,从损失函数的选择、batch 配置、训练脚本到 16G 显存的优化,把完整流程过一遍。

4.1 损失函数与温度系数

多模态 Embedding 训练最常用的损失函数是 InfoNCE(或者叫 NT-Xent 损失)。这个损失函数的直觉理解是这样的:你有一批图文对(比如一个 batch 里有 256 张图和对应的 256 段文本),对每一张图来说,正确的文本是它在 batch 里的匹配项,其他 255 条文本都是负样本;对每一段文本来说同理。然后模型要最大化正样本对之间的相似度,同时最小化负样本对之间的相似度。

公式我不展开推导了,但有个参数你必须重点调:温度系数(temperature)。它控制模型对负样本的“狠心程度”。温度越低,模型越关注那些相似度高的难负样本,区分粒度越细,但容易训练不稳定;温度越高,模型对负样本的惩罚越宽松,训练稳定,但区分能力弱。常见设置是 0.07 或 0.1,我在微信场景的数据上实测,0.05 到 0.08 这个区间效果最好。

代码实现上,PyTorch 里可以直接基于 cosine similarity 计算,我不建议手动写循环,用矩阵运算一次性算完。下面是我常用的一个简化版 InfoNCE 实现,batch size 为 N,特征维度为 d:

import torch import torch.nn.functional as F def info_nce_loss(image_embeds, text_embeds, temperature=0.07): # image_embeds: [N, d], text_embeds: [N, d] # 两个嵌入都已经是 L2 归一化后的 # 计算图文相似度矩阵 [N, N] logits = image_embeds @ text_embeds.T / temperature # 对角线是正样本,构造标签 labels = torch.arange(logits.shape[0], device=logits.device) # 计算双向损失:图片到文本、文本到图片 loss_img = F.cross_entropy(logits, labels) loss_txt = F.cross_entropy(logits.T, labels) return (loss_img + loss_txt) / 2

这个实现里有几个细节注意。logits 矩阵的 shape 是 [N, N],在 batch 为 256 时,这是一个 256x256 的矩阵,显存开销很小,但计算量不小,所以 batch 大小受限的主要是双塔的前向计算,而不是对比损失本身。我习惯把梯度的计算放在两个方向上,也就是对称损失。为什么?因为一个 batch 里,图片到文本的匹配和文本到图片的匹配难度往往不对称,对称损失能让模型双向对齐都学到。

4.2 完整的训练脚本结构

下面是我实际用来跑训练的一个骨架脚本。它不完整,但包含了最核心的训练循环、混合精度、梯度累积和日志输出,你可以直接在这个基础上改造。

import os import argparse import torch from torch import nn from torch.cuda.amp import autocast, GradScaler from torch.utils.data import DataLoader from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR from transformers import CLIPVisionModel, BertModel from dataset import WeChatMultimodalDataset from losses import info_nce_loss def parse_args(): parser = argparse.ArgumentParser() parser.add_argument("--batch_size", type=int, default=256) parser.add_argument("--lr", type=float, default=2e-5) parser.add_argument("--epochs", type=int, default=10) parser.add_argument("--temperature", type=float, default=0.07) parser.add_argument("--accum_steps", type=int, default=1) parser.add_argument("--output_dir", type=str, default="./checkpoints") parser.add_argument("--grad_clip", type=float, default=1.0) return parser.parse_args() def main(): args = parse_args() # 双塔模型 image_encoder = CLIPVisionModel.from_pretrained("openai/clip-vit-base-patch16") text_encoder = BertModel.from_pretrained("bert-base-chinese") # 投影层:把各自维度映射到 512 维 img_proj = nn.Sequential( nn.Linear(image_encoder.config.hidden_size, 1024), nn.GELU(), nn.Linear(1024, 512) ) txt_proj = nn.Sequential( nn.Linear(text_encoder.config.hidden_size, 1024), nn.GELU(), nn.Linear(1024, 512) ) model = nn.ModuleDict({ "image_encoder": image_encoder, "text_encoder": text_encoder, "img_proj": img_proj, "txt_proj": txt_proj }) model.cuda() train_dataset = WeChatMultimodalDataset(...) train_loader = DataLoader( train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True, pin_memory=True ) optimizer = AdamW(model.parameters(), lr=args.lr, weight_decay=0.02) scheduler = CosineAnnealingLR(optimizer, T_max=args.epochs * len(train_loader)) scaler = GradScaler() global_step = 0 for epoch in range(args.epochs): for batch in train_loader: images = batch["image"].cuda() texts = batch["text"].cuda() attention_mask = batch["attention_mask"].cuda() with autocast(): img_feats = image_encoder(pixel_values=images).pooler_output txt_feats = text_encoder(input_ids=texts, attention_mask=attention_mask).pooler_output img_embeds = F.normalize(img_proj(img_feats), dim=-1) txt_embeds = F.normalize(txt_proj(txt_feats), dim=-1) loss = info_nce_loss(img_embeds, txt_embeds, args.temperature) loss = loss / args.accum_steps scaler.scale(loss).backward() if (global_step + 1) % args.accum_steps == 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) scaler.step(optimizer) scaler.update() optimizer.zero_grad() scheduler.step() if global_step % 100 == 0: print(f"Epoch {epoch} | Step {global_step} | Loss: {loss.item():.4f}") global_step += 1

这个脚本看起来不复杂,但每个部件背后都有讲究。比如为什么用 CLIPVisionModel 而不是直接用 torchvision 的 ResNet?因为 CLIP 的视觉编码器本身经过大规模图文对比学习预训练,做初始化起点比 ImageNet 分类预训练的 ResNet 好得多,收敛速度快,最终效果也明显更好。文本塔为什么用 BERT?因为在中文场景下,BERT 类模型的语义理解能力比轻量模型稳定,如果你还要同时支持英文内容,可以直接换成 XLM-RoBERTa,不需要改结构。

4.3 16G 显存的优化策略与实测配置

很多朋友一看到 batch size 256,第一反应是“我显卡不够”。别急,16G 显存完全够用,关键是做好三件事:混合精度、梯度累积、梯度检查点(gradient checkpointing)。

混合精度我用的是 torch.cuda.amp 自带的自动混合精度,脚本里已经写了。这套机制的原理一句话总结:前向和反向计算中,对显存占用大的部分用 FP16,对数值敏感的部分(比如 loss 缩放)保持在 FP32。实测下来,训练速度能提升大约 40%,显存占用能降低 30% 到 40%。

梯度累积解决的是“batch 小了负样本不够”的问题。如果你的卡只能跑 batch size 32,那就在 8 步中累积梯度,等价于一个 256 的 batch。这里有个关键细节:累积梯度时,loss 需要除以累积步数。脚本里已经写了 loss = loss / args.accum_steps,这一步很多人会忘,忘掉之后的学习率等效于放大了 accum_steps 倍,模型容易训飞。

梯度检查点适合在显卡实在不够、还想上 ViT-L 的情况。它的原理是不保存中间激活值,反向传播时重新计算。代价是训练速度会慢 20% 左右。我的建议是,如果你用的是 ViT-B,不需要梯度检查点,浪费速度不划算;如果上 ViT-L,那就开。

我实测过几组配置,整理成表格,你可以对照自己的卡来选:

模型组合batch size梯度累积显存占用备注
ViT-B/16 + 6层BERT1282约 9G推荐,速度质量均衡
ViT-B/16 + 6层BERT2561约 13G显存刚好,不建议再加序列长度
ViT-L/14 + 6层BERT644约 14G开启梯度检查点
ViT-L/14 + 12层BERT328约 15.5G极限配置,不推荐

这里面有个容易忽略的额外显存消耗,来自对比学习的 logits 矩阵。当 batch size 为 256 时,logits 是 256x256,占用不大;但当 batch size 达到 1024 时,logits 矩阵就占了 8MB 显存,前向计算时梯度图还会翻几倍,所以在超大 batch 下,对比损失本身也会成为显存瓶颈。

5. 推理、向量检索与模型部署的落地细节

训练不是终点,模型训完要上线跑检索才有价值。微信场景下的多模态 Embedding 应用,主要是在向量数据库里做相似度检索。这里面的工程坑不少。

5.1 推理时的向量处理流程

模型训练完之后,推理时有一个容易被忽视的约束:训练时输入是“一批图 + 一批文本”,推理时往往是“一张图 + 一段文本”或者“一段文本 + 一个候选图片集”。由于训练时用了 batch 内的对比学习,模型对 batch 统计有一定依赖性,推理时 batch 变小会导致向量分布轻微偏移。解决方法是推理时对向量做标准化(L2 normalize),并且尽量保持和训练一致的输入预处理流程。

另外,推理输入的图片分辨率要和训练一致。我见过有人训练时用 224x224,推理时为了省时间改成 112x112,结果向量质量和检索效果跳水。这不是模型的问题,是输入分布不一致的问题。

推理阶段还有一个细节:文本截断长度。训练时我截断到 64 token,推理时必须保持同样的截断策略。如果推理时允许长文本输入,模型会看到训练时从未见过的 token 位置,pooler_output 的质量会退化。

5.2 向量降维与量化

Embedding 做完之后是 512 维的 float 向量,一个向量占 2KB。如果线上有 1 亿条内容,那就是 200GB 的原始向量,存储压力很大。所以实际工程中我一般会做降维和量化。

降维用 PCA。先在验证集上计算所有向量的 PCA 主成分,保留前 128 维或 256 维,一般能保留 95% 以上的检索性能,但存储省一半以上。这里要注意,PCA 的变换矩阵必须在训练集或独立的验证集上拟合,不能在整个数据库上拟合,否则会泄漏信息,影响后续增量数据的处理。

量化用 Product Quantization(PQ)。把 512 维向量切成 16 个子空间,每个子空间 32 维,用 KMeans 聚类成 256 个中心点,这样每个子空间只需要 1 个字节存储中心点索引,最终一个向量只需要 16 字节。配合倒排索引(IVF),检索速度会快很多。但 PQ 是有损压缩,如果你们的业务对精度要求很高,我建议用 OPQ 做旋转优化,可以在同样码率下提升不少召回率。

5.3 与向量数据库的对接

现在市面上有 Milvus、Faiss、Qdrant、Weaviate 等主流向量检索库。对于微信这种规模,我倾向于用 Milvus 或者基于 Faiss 自建检索服务。Faiss 的优势是底层库性能稳定,支持 GPU 加速,适合技术团队自己封装;Milvus 的优势是运维省心,自带数据管理、索引构建、监控告警。

不管用哪个,有几件事必须做。一是建立独立的评估集,每个版本上线前都跑一遍召回率对比,防止模型迭代后某个指标回退。二是设置相似度阈值,多模态检索里跨模态的相似度天然低于单模态检索,阈值低了会返回一堆噪音,阈值高了会漏召回,这个阈值需要用线上日志去统计分布来定,不能拍脑袋。

另外,如果业务里还需要关联 RAG 流程,那 Embedding 向量不仅仅是做多模态召回,它还要和文本 Chunk 的向量做混合检索。热词里出现“deepseek embedding ragflow”,说明不少人会使用 RAGFlow 这类系统。RAGFlow 默认用他们自己的 Embedding,但可以替换成你训练好的模型。接入方式是按照 RAGFlow 的 Embedding 接口实现一个自定义类,把模型部署成一个 HTTP 服务,然后配置指向这个服务。实测下来,替换自己训练的多模态 Embedding 之后,在包含图片描述的文档检索场景下,召回质量明显好于通用 Embedding 模型,因为通用模型没见过你的业务数据分布。

6. 评估指标与调优心得

模型训完,怎么判断它好不好?只看 loss 下降是不够的,loss 降低只代表模型拟合了训练数据,不代表检索效果好。我从实际项目中总结出一套评估方案,如果只能看三个指标,那就看 Recall@K、MRR 和向量分布的可视化。

6.1 三个必须看的指标

Recall@K 衡量的是:对每个查询,模型返回的前 K 个结果中,有多少比例覆盖了人工标注的正确结果。K 一般取 10 或 50。这个指标对召回场景最直观。

MRR(Mean Reciprocal Rank)衡量的是:正确结果排在返回列表第几位。第一位得 1 分,第二位得 0.5,第三位得 0.33,等等。MRR 对排序质量更敏感,搜索场景尤其关注它。

向量分布可视化是调优的辅助手段。用 t-SNE 或者 UMAP 把验证集的向量降到二维,然后按内容类别着色。一个训练良好的模型,同类内容的向量应该自然聚成团,不同类之间有清晰的分界。如果看到类别之间乱成一锅粥,说明模型没学到区分特征,这时候调损失函数和难负样本比调网络结构有用。

6.2 调优过程中的几个关键判断

在微信业务数据上,我遇到过一个很典型的现象:模型训练到第 3 个 epoch 之后,训练 loss 还在下降,但验证集 Recall@10 不再上升了。这是明显的过拟合信号。我的做法是从第 3 个 epoch 开始,每个 epoch 都保存一次 checkpoint,最后根据验证集指标选择最优的那一版,而不是训练到最后一个 epoch。另外,把温度系数在训练中后期从 0.07 慢慢提高到 0.1,也可以缓解过拟合,因为更高的温度会让模型对负样本的关注度降低,减少对训练集噪音的拟合。

调优难负样本池的比例也值得记录。负样本池占比从 0% 增加到 2% 时,Recall@10 能提升 3 到 5 个点;但继续增加到 5%,收益反而下降,因为太难的负样本会让模型训练不稳定,loss 震荡明显。我最后固定在了 1.5% 到 2% 之间。

一个高效的技巧是多模态模型蒸馏。如果你有条件使用一个更强的多模态大模型作为教师模型,可以拿它来给图文对生成更精细的相似度分数,然后用这个分数作为软标签,蒸馏到小 Embedding 模型上。这个方向有专门的算法,不是简单的把大模型的输出向量拿过来对齐,而是要去蒸馏大模型对样本对的相对关系。具体操作上,把教师模型对每个 batch 内所有图文对的相似度矩阵算出来,用温度 T=2 做 softmax 得到分布,让小模型的相似度分布去拟合这个软标签分布。训练速度会慢一些,但效果提升很值得。

6.3 和通用模型的对比

训练完成之后,对比是必须做的。我通常会拿我们训练好的模型和我们之前的通用多模态模型做一个同测试集对比。评价维度包括:单模态内的图文检索、跨模态检索、跨语言检索、视频帧检索。

这里有个很实用的结论:通用模型在通用领域的图文匹配上可能更好,但在业务专属场景(比如垂直领域的专业名词、用户的特定表达习惯)下,微调后的模型明显占优。通用模型的优势是泛化能力好,冷启动能力强,适合上线初期没有业务数据的情况;一旦积累了十万级以上高质量的图文对,训练自己的模型就能显著拉开差距。

7. 常见问题与排查技巧实录

最后这部分,我把自己训练多模态 Embedding 模型过程中踩过的坑整理成一个排查手册。你可能不会一次全遇到,但遇到的时候至少知道从哪个方向排查。

7.1 显存 OOM 的排查顺序

显存不够是最常见的问题,但很多人一看到 OOM 就盲目减小 batch size,这是不对的。正确的排查顺序是:先看是不是输入分辨率太大,把 224 改成 160,显存能省一大半;再看序列长度,如果你把文本截断在 128 token 而实际大部分文本只有 20 个 token,那就在数据加载时动态截断到实际长度的 90% 分位;然后看混合精度是否开了,梯度检查点是否开了;最后才考虑减小 batch size。

我用一个实际案例说明:一个朋友复现我的配置,在 16G 卡上跑不起来,报 OOM。远程一看,batch size 设的是 256,图片预处理时 resize 到 384x384,文本序列长度是 512。三个因素叠加,显存爆了很正常。把分辨率改回 224、序列长度改为 64 之后,显存从 14G 降到了 8G,问题直接解决。

7.2 loss 剧烈震荡或变 NaN

loss 震荡的原因很多,最常见的是学习率太大、温度系数太低、难负样本太多。排查思路是按顺序排除:把学习率降一个数量级(比如从 2e-5 降到 2e-6),看 loss 是否稳定;如果稳定了,说明是学习率问题;然后把温度系数提高到 0.1,看是否稳定;再检查难负样本池占比,如果超过 3%,降低到 1%。

NaN 的问题一般出在混合精度上。fp16 的表示范围有限,遇到极端值容易溢出。解决方法是检查数据预处理:文本 token 里不能有 NaN 的 attention mask;图片像素值必须在 [0,1] 或 [0,255] 且类型正确;确认投影层输出之后做了 L2 归一化,归一化会把数值压缩到 [-1,1],可以防止后续矩阵乘法溢出;最后看 GradScaler 是否正常工作,如果 loss_scale 在训练中不断下降,说明有梯度溢出,需要检查模型里是否有数值不稳定的层。

7.3 模型效果一直上不去,怎么排查

这是最让人头疼的情况。模型训练没报错,loss 也下降了,但检索效果就是不行。我的排查路线是:先检查训练数据的质量,随机抽 200 条图文对,人工看一遍,如果发现很多图文根本不相关,那就先去修数据。这一步能解决 60% 的“效果不行”问题;数据没问题,再去查负样本构造,是不是负样本太简单了,模型没有压力去学区分能力,如果是,做一轮难负样本挖掘;负样本也没问题,就看双塔之间是否对齐,把训练集里图文对相似度的分布画出来,如果分布峰值接近 1,说明模型把所有东西都映射到了同一个区域,这是模型退化的典型特征,需要调整投影层结构或者加深投影 MLP。

这里有一个我常用的“体检”脚本,用 100 对标注图文对,看模型在同 batch 内的相似度排序,如果正样本对的平均排序低于前 10%,说明模型有严重问题,需要回退到模型结构层面的检查。

7.4 模型推理速度太慢

多模态 Embedding 在线上推理时,图像塔是计算瓶颈。一张图过 ViT-B 的前向大约需要 15ms 到 30ms(取决于 GPU 型号),文本塔只需要 2ms 到 5ms。优化思路有几个:图像塔用 TensorRT 或者 ONNX Runtime 做加速,实测能把单张推理时间压到 8ms 以内;对图片做缓存,同一个 URL 的图片特征只算一次,命中缓存直接返回;如果对延迟要求极高,可以考虑把视觉塔替换成更轻量的结构,比如 MobileViT。

微信场景里还有个特殊情况:图片可能被压缩过,EXIF 旋转可能没被正确处理,导致推理时图片是横着的。这个在数据预处理时就要处理好,统一按 EXIF 信息旋转,再 resize。我见过一个 case,线上有 3% 的图片带着旋转信息没处理,那 3% 的检索结果全乱了。

8. 后续可以扩展的方向

多模态 Embedding 训练到能上线,只是第一步。按照我的经验,后面有几个值得继续投入的方向。

一个方向是在线难负样本挖掘。我用的是离线定期挖掘,更激进的方案是在每轮迭代中实时挖掘困难负样本并缓存,这样模型训练时看到的负样本永远比当前模型能力略难一点,收敛快,上限高。缺点是工程复杂度高,且需要注意不要把标签错误的样本引入训练集。

另一个方向是加入音频模态。微信场景里视频号包含大量语音信息,如果把音频也映射到同一向量空间,可以实现“听到某段音频就搜到相关画面”的体验。做法是再加一个音频塔,用预训练的 Wav2Vec 或者 CLAP 做编码器,训练策略和图文对比基本一样。

还有多模态模型蒸馏这条路。热词里出现“模型蒸馏”,这个是当前业界常用的见效方法。如果你有一个更强的大模型(比如 Qwen2-VL 这类多模态大模型),可以在离线批量推理时让它产出高质量的伪标签,再蒸馏到小模型上。这个方向对资源的要求高一些,但效果往往比只靠自己标注数据训练要好,值得投入。

最后,向量索引的更新策略也很关键。微信这种每天产生大量新内容的场景,增量索引的时效性很重要。我现在的方案是:新内容入库时先抽取向量,写入一个临时的“新鲜索引”,每隔半小时和主索引做一次合并。这样既能保证新内容的实时召回,又不会因为频繁合并索引导致检索性能下降。实际线上,这个方案能让新内容的可搜时间从小时级降到分钟级,这个体验差异用户能明显感知到。

多模态 Embedding 这个方向,没有一个放之四海而皆准的答案。模型结构、数据配比、负样本策略、温度系数、评估指标,每个环节都需要针对自己的数据分布反复调。我自己训第一版的时候,光在温度系数和难负样本池配比上就实验了将近一周。但门槛不高,只要有清晰的数据思路和一套可迭代的训练管线,结果一定会比直接用通用模型好。希望这篇实践记录能让你少走一些弯路。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询