ARTICLE DETAIL

资讯详情

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

CLIP图像文本跨模态检索:从对比学习原理到微调与工程落地

CLIP图像文本跨模态检索:从对比学习原理到微调与工程落地 简介基于 CLIP 模型的图像文本跨模态检索方案以 PDF 电子文档形式提供面向多模态学习、计算机视觉与自然语言处理交叉领域的研究者与学生针对图像与文本之间的语义鸿沟以及跨模态检索困难给出了完整的建模与实验思路。文档从数据预处理入手依次介绍图像裁剪、随机旋转、色域增强等图像增强方法以及小写转换、去除标点等文本增强操作并详细讲解视觉 Transformer 模型完成图像特征提取、文本 Transformer 模型完成文本特征提取以及按八比二划分训练集与测试集、通过交叉熵损失进行对比预训练、再完成零样本分类的全过程。同时文档还说明了如何利用损失率和召回率 RecallK 评估模型并分别对图像检索和文本检索任务给出最佳学习率范围与前五名检索结果。资源共 1 个文件为 PDF 类型压缩包大小 4.48MB。目前已有 238 人学习下载适合需要复现跨模态检索实验、学习 CLIP 应用方法或参考竞赛方案的研究人员使用。1. 基于 CLIP 模型的图像文本跨模态检索先别急着调参搞清楚它到底解决了什么说到基于 CLIP 模型的图像文本跨模态检索很多人的第一反应是这不就是给图片找句子、给句子找图片吗对也不对。真正的难点不是你用哪个模型而是你怎么定义“相似”——一张“戴着红色帽子的狗在草地上奔跑”的图和文本“一只狗在户外玩耍”到底算不算匹配CLIP 给出的答案是算因为它把图和文本都映射到同一个语义向量空间然后用余弦相似度衡量距离。这个思路直接绕过了传统方法里“先检测物体、再对齐属性”的复杂管线把跨模态检索简化成一个向量查最近邻的问题。但这篇笔记不打算只讲概念。我要带你从零走一遍CLIP 的原理边界在哪里、怎么用开源权重先跑通一个最小检索系统、怎么用中文数据微调、以及真正在业务里落地时那些网上没人明说的坑。适合谁看正准备用 CLIP 做以图搜图、图文匹配、数据清洗去重或者在多模态方向起步的工程师。你已经能用 Python 跑通基本模型调用但还没系统处理过跨模态检索的工程问题。我先说一句可能反直觉的话CLIP 的检索效果七成取决于你怎么处理数据三成才取决于模型本身。2. 为什么 CLIP 能跨模态检索对比学习与特征对齐的逻辑2.1 它不是“看图说话”而是把两个模态压到同一个向量空间想理解基于 CLIP 模型的图像文本跨模态检索得先接受一个前提CLIP 不是一个“理解了图像内容”的模型它是一个“学会了图文配对”的模型。它在训练时做的事情非常粗暴——拿 4 亿对图像文本数据让模型学会把配对的图文在向量空间里拉近、把不配对的推远。这就是对比学习Contrastive Learning的核心思想。具体训练目标是对称交叉熵损失也就是 InfoNCE 的一个变体。对于一个 batch 里的 N 对图文模型分别把图像编码成 image_features、把文本编码成 text_features然后计算一个 N×N 的相似度矩阵。矩阵对角线上的元素是正样本对其他位置都是负样本。理想情况下对角线上的相似度要远高于同行同列的其他值。# 伪代码理解 CLIP 训练目标的核心逻辑 import torch import torch.nn.functional as F def clip_loss(image_features, text_features, logit_scale): # image_features: [batch_size, embed_dim] # text_features: [batch_size, embed_dim] # logit_scale: 可学习的温度参数初始值通常为 ln(1/0.07) # 归一化计算余弦相似度矩阵 image_features F.normalize(image_features, dim-1) text_features F.normalize(text_features, dim-1) logits logit_scale * image_features text_features.T # [N, N] # 对称损失图像到文本、文本到图像各算一次 labels torch.arange(logits.shape[0], devicelogits.device) loss (F.cross_entropy(logits, labels) F.cross_entropy(logits.T, labels)) / 2 return loss这段代码里最关键的是logit_scale。它是一个可学习的标量参数作用是把相似度数值放大到合适的区间再算交叉熵。如果你自己微调时不小心把它初始化得太小比如 0.1模型会训不动因为 logits 的值域太窄softmax 之后几乎是个均匀分布梯度消失。常见做法是初始化为1 / 0.07的自然对数也就是大约 2.66然后让模型自己学。2.2 双塔结构为什么它天生适合做检索CLIP 的架构是两个编码器图像侧用 ViT 或 ResNet文本侧用 Transformer。两个塔完全独立只在算损失的时候才交互。这个设计对检索任务有天然优势你可以把所有的图像特征预先算好存进向量数据库查询时只需要算一次文本特征然后做向量检索就行。这和双塔召回在推荐系统里的用法如出一辙。需要提醒的是图像塔和文本塔输出的向量维度是一样的比如 512 维或 768 维但它们的向量空间不是天然对齐的必须经过上面的对比学习训练才行。如果你直接拿一个纯 ViT 模型和一个纯 BERT 模型各自输出向量去做余弦相似度结果基本是随机的。很多人踩坑就是以为“只要是向量就能比”其实跨模态检索的前提是对齐过的向量空间。2.3 CLIP 文本编码器的输入格式一个容易被忽略的细节从热词里看到有人问“clip文本编码节点怎么输入内容”这里统一说清楚。CLIP 的文本编码器不是拿原始字符串直接进去的它需要经过 tokenizer 处理成 token ids并且有固定的上下文模板。OpenAI 官方在训练时使用了比如“a photo of a {}”这样的模板推理时你也可以用但对检索来说模板的作用没想象中那么大。真正要注意的是输入长度上限通常是 77 个 token超出的部分会被截断。tokenizer 是 BPE 级别的词汇表大小是 49408和 GPT-2 基本一致。文本端还有一个特殊的开始符|startoftext|和结束符|endoftext|处理时别漏掉。from transformers import CLIPProcessor processor CLIPProcessor.from_pretrained(openai/clip-vit-base-patch32) inputs processor(text[a photo of a cat, 一只猫的照片], return_tensorspt, paddingTrue) # 输出包含 input_ids 和 attention_mask # 如果文本超过 77 个 token会用 [SEP] 截断截断位置在结尾这两行代码里paddingTrue会把 batch 里较短的文本补到和最长文本等长return_tensorspt是让输出变成 PyTorch Tensor。注意CLIPProcessor内部会自动组合图像和文本的预处理但如果你用了CLIPTokenizer和CLIPImageProcessor分开做要确保两边的padding和truncation策略一致否则 batch 推断时会报 shape 不匹配的错。3. 跑通最小跨模态检索环境准备、数据路径与推理3.1 安装与权重选择别盲目上 large 版本先落地。基于 CLIP 模型的图像文本跨模态检索最稳妥的起步方式是直接用 Hugging Face 上的开源权重。官方原版有clip-vit-base-patch32、clip-vit-base-patch16、clip-vit-large-patch14等。我的建议是第一轮跑通用 base 版本别上 large。Large 在推理时显存占用大约 2-3 倍速度慢一半不止但检索精度提升可能只有两三个点对验证流程来说不划算。# 创建一个干净的 conda 环境避免和已有项目冲突 conda create -n clip-retrieval python3.10 conda activate clip-retrieval pip install torch torchvision transformers sentencepiece ftfy regex # ftfy 和 regex 是 CLIP 文本预处理依赖的包缺了会报 ImportError注意sentencepiece不要漏装。很多人在 Hugging Face 上加载 CLIP 权重时报错就是因为缺这个包。另外PyTorch 版本建议 2.0 以上因为 transformers 新版对注意力 mask 的实现在旧版上可能行为不一致。3.2 准备一套最朴素的检索数据混合语言的坑既然标题是“图像文本跨模态检索”我们就做一个能跑的 demo给定一张图片从一批文本候选中找到它最匹配的句子。这里有一个常见的坑——CLIP 的中文能力很弱除非你用的是专门的中文 CLIP 权重比如OFA-Sys/chinese-clip-vit-base-patch16或经过中文微调的模型。直接用 OpenAI 原版权重做中文检索效果会让你怀疑人生。以下代码用中文 CLIP 权重做图文匹配。如果你是做英文场景把模型名换成openai/clip-vit-base-patch32即可。from transformers import CLIPProcessor, CLIPModel from PIL import Image import requests model_name OFA-Sys/chinese-clip-vit-base-patch16 model CLIPModel.from_pretrained(model_name) processor CLIPProcessor.from_pretrained(model_name) # 读取图片这里用一张猫的图片做测试 image Image.open(requests.get(https://demo.com/cat.jpg, streamTrue).raw) # 多语言混合的候选文本 candidate_texts [ 一只橘色的猫咪坐在窗台上, a dog playing in the park, 一个女孩在街道上骑自行车, 暖阳下的猫安静地看着窗外, ] inputs processor( textcandidate_texts, imagesimage, return_tensorspt, paddingTrue, ) with torch.no_grad(): outputs model(**inputs) # 这里拿到的 logits_per_image 就是图文相似度数值越大越匹配 probs outputs.logits_per_image.softmax(dim-1) for text, prob in zip(candidate_texts, probs[0]): print(f{text}: {prob:.4f})这个例子能跑通但有几个参数值得说明。logits_per_image是一个 1×N 的张量表示该图像和 N 条文本的相似度。如果你反过来做文本查图像应该看logits_per_text。实际使用中这两个方向的结果并不完全对称因为训练损失是对称的但推理时归一化方式会略有差异。3.3 相似度计算的玄学温度和归一化不能乱动检索时最常被忽略的是温度参数。模型内部的 logit_scale 在训练时已经被固定下来了推理时直接乘到相似度矩阵上。但你如果自己用特征向量算余弦相似度必须先把两个向量各自做 L2 归一化然后再点乘。如果不归一化直接用原始特征算点积结果会受向量模长影响排序和模型内部的输出不一致。# 正确的相似度计算方式先归一化再点乘 import torch.nn.functional as F image_feature outputs.image_embeds # [1, 512] text_features outputs.text_embeds # [N, 512] image_feature_norm F.normalize(image_feature, dim-1) text_features_norm F.normalize(text_features, dim-1) similarities (image_feature_norm text_features_norm.T).squeeze(0)到这里一个最小系统已经通了。但注意这个 demo 里的候选文本只有 4 条属于玩具级。真正做检索时文本候选可能有几十万条那就不能 for 循环去算了要用向量检索工具后面第 6 章会讲。4. 微调 CLIP 模型数据集构建、训练脚本与必调参数4.1 微调到底在调什么不是重新训练是对齐你的数据分布基于 CLIP 模型的图像文本跨模态检索在业务里几乎都要微调。原因很简单CLIP 预训练时用的是互联网图文对你的业务数据比如电商的商品图标题、医疗影像报告在分布上差异很大。微调的本质是用你的数据让模型重新对齐图文特征空间而不是教它认识新概念——那是 fine-tune 整个分类头才需要担心的事。微调时常见的做法是冻结图像塔和文本塔的底层参数只更新高层和对比投影层。但说实话如果你的数据量不足 1 万对我建议全部参数都冻结只训练投影层。这样能极大降低过拟合风险训练速度也快很多。数据量在 5 万对以上时可以考虑解冻文本塔的最后 2 层。4.2 数据格式与 DataLoader图像文本对到底要怎么组织微调最忌把数据组织成“一张图配一段字”就直接丢进去。你需要考虑的是同一张图可以有多个角度描述同一个描述可以对应多张图。数据清洗时优先保证一个原则——负样本不要出现在同一个 batch 里。比如你有 100 张几乎相同的商品图只是背景颜色不同它们的文本描述也高度相似。如果这些正样本对同时出现在同一个 batch 里模型会把它们互相当成负样本损失函数会非常困惑。from torch.utils.data import Dataset, DataLoader from PIL import Image class ImageTextPairDataset(Dataset): def __init__(self, paired_data, processor, image_dir): self.data paired_data # list of (image_path, text) self.processor processor self.image_dir image_dir def __len__(self): return len(self.data) def __getitem__(self, idx): image_path, text self.data[idx] image Image.open(f{self.image_dir}/{image_path}).convert(RGB) # 强制统一图片尺寸CLIP 的 processor 内部会做 resize 到 224x224 return self.processor(text[text], imagesimage, return_tensorspt, paddingTrue)这段代码是 Dataset 的核心骨架。有两个细节要注意convert(RGB)是为了处理灰度图和 RGBA 图否则某些图片会报通道数错误。return_tensorspt返回的是 tensor但由于图像和文本都套了[]中括号batch 维度是 1后续用 DataLoader 的collate_fn时要手动处理拼接不能用默认的 stack。4.3 训练脚本一个能直接用的微调流程微调的损失函数和预训练完全一样就是前面写的对称交叉熵。所以你不需要改损失函数只需要加载预训练权重把model设成训练模式然后数据传输到 GPU 上。import torch from torch.optim import AdamW from torch.cuda.amp import autocast, GradScaler model CLIPModel.from_pretrained(OFA-Sys/chinese-clip-vit-base-patch16) processor CLIPProcessor.from_pretrained(OFA-Sys/chinese-clip-vit-base-patch16) # 只训练投影层image_projection 和 text_projection for name, param in model.named_parameters(): if projection not in name: param.requires_grad False optimizer AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr1e-5) scaler GradScaler() model.cuda() for epoch in range(3): for batch in dataloader: images batch[pixel_values].cuda() input_ids batch[input_ids].squeeze(1).cuda() attention_mask batch[attention_mask].squeeze(1).cuda() with autocast(): outputs model(pixel_valuesimages, input_idsinput_ids, attention_maskattention_mask) logits outputs.logits_per_image # [batch, batch] labels torch.arange(len(images)).cuda() loss (F.cross_entropy(logits, labels) F.cross_entropy(logits.T, labels)) / 2 scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()这里面有三个参数需要重点说。第一个是学习率lr1e-5这是微调 CLIP 的安全区间超过 1e-4 很容易让模型遗忘预训练知识表现为 loss 下降很快但检索效果反而变差。第二个是autocast混合精度如果不加训练速度会慢很多而且显存占用翻倍。第三个是只训练投影层的策略这里的projection名字匹配在各版本 transformers 里都适用但如果你用了本地保存的模型权重命名可能有差异最好先model.named_parameters()打印出来确认。4.4 必调参数表照着这个范围粗调就好微调 CLIP 的或者说基于 CLIP 模型的图像文本跨模态检索方向微调的关键参数就那么几个我把经验范围列成表参数推荐范围说明学习率1e-6 ~ 5e-5冻结底层时用偏大值解冻全部时用偏小值Batch Size64 ~ 256对比学习吃 batch太小了负样本不够温度参数初始值2.66即 ln(1/0.07)可以训练更新别设成 1.0训练轮数2 ~ 5多了容易过拟合尤其投影层图像分辨率224base336large模型训练时的原始分辨率别自己乱改文本最大长度77CLIP 的 BPE 编码上限长了就截断我特别想强调 batch size。很多人在单卡上被迫用 16 或 32 的 batch然后发现检索效果怎么都不好。因为对比学习的负样本来自 batch 内其他样本batch 太小意味着负样本太少模型学不到精细的语义差异。如果你的卡放不下大 batch用梯度累积模拟也行但效果不如真正的 batch 大。5. 跨模态检索的 5 个踩坑记录现象、原因与解决5.1 中文长文本检索效果崩了永远检查 tokenizer 是否对现象用中文 CLIP 做检索短文本5-10 个字效果还行一旦文本超过 20 个字检索结果完全乱套甚至出现“和查询完全不相关的结果排第一”。原因中文 CLIP 的 tokenizer 基于中文 BPE但原始预训练数据里的长文本可能包含大量无意义的修饰词导致模型在长序列上的注意力分布被稀释。更关键的是CLIP 的训练文本都是短描述基本不超过一句完整的话模型对长文本泛化能力本来就弱。解决把输入文本预处理为“关键词堆叠”而非自然语言句子。比如“一只白色的狗在草地上奔跑旁边有小孩在玩耍”改为“白狗绿地奔跑小孩玩耍”。虽然不合语法但去除了冗余词检索精度能提升十几个百分点。这是我在实际项目里试过最有效的土办法。5.2 图像预处理的细节不一致导致特征偏移现象训练时用的图像是正方形缩放到 224推理时直接拿原始比例的图片丢进去检索效果下降了 3-5 个点。原因CLIP 处理器在做 resize 时默认是crop还是resize取决于你初始化时用的参数。如果推理时用了不同的图片预处理方式特征分布就会有偏移。解决代码里统一用processor.image_processor的默认配置不要自己手动写torchvision.transforms.Resize。如果你一定要自定义复制processor.image_processor里的配置逐项对应修改特别是do_center_crop和do_resize这两个开关。5.3 计算特征时忘了加torch.no_grad()现象特征抽取代码跑得没问题但推理阶段显存慢慢上涨最终 OOM而且系统的检索延迟忽高忽低。原因没有在特征抽取时禁用梯度计算PyTorch 会为每个操作构建计算图显存被激活图占满。解决所有推理时段的 forward 外面包torch.no_grad()并且顺手model.eval()。这两个缺一不可model.eval()影响 batch norm 和 dropouttorch.no_grad()影响梯度图。5.4 负样本选择不当导致微调效果“假好”现象微调后 loss 降到了 0.1 以下但用真实查询评测的 Recall10 只提升了一个点有时候甚至更差。原因训练数据里的负样本太简单了。比如你做商品检索正样本是“红色跑鞋”负样本是“蓝色跑鞋”——模型轻松就能区分颜色这种对比学不到高层语义。真正需要的是难负样本比如“红色跑鞋”和“红色休闲鞋”这样模型才被迫学品类边界。解决在训练 Loss 里引入难负样本挖掘策略每一轮从向量库里选出相似度排名 50-100 名的样本作为额外负样本把这部分 logits 加入交叉熵计算。代码上在labels那里做一点手脚把难负样本的标签设成-100忽略该位置的惩罚。5.5 微调和推理的模型版本不一致现象你微调了 A 版本的 CLIP但推理代码里还在用openai/clip-vit-base-patch32结果怎么调都复现不了你在训练时的验证效果。原因这不是玄学是路径问题。很多微调脚本会只存 model 权重不存 processor结果推理时你用的可能还是旧 processortokenizer 的词汇表不一样文本侧的特征就完全对不上。解决微调完直接把整个文件夹都保存下来包括 processor、tokenizer、model。用model.save_pretrained(save_dir)和processors.save_pretrained(save_dir)推理时统一用from_pretrained(save_dir)加载。这可以说是这个领域最常见的“自己坑自己”。6. 把 CLIP 接进业务系统批量向量化与近似最近邻检索6.1 批量抽取特征Batch Inference 的正确姿势当你真正要构建一个基于 CLIP 模型的图像文本跨模态检索系统时核心工作是把大量候选数据离线向量化。以商品检索为例100 万张图片用单卡 A100 抽取特征大约需要 3-4 小时。这里有两条优化路径第一用batch_size尽可能大让 GPU 利用率上去第二用DataLoader的num_workers把图像解码从 GPU 计算里剥离开。from torch.utils.data import DataLoader from tqdm import tqdm def extract_image_features(model, image_paths, processor, batch_size128): # 先把所有图像路径打包成 Dataset沿用前面的 Dataset 结构 dataset ImagePathDataset(image_paths, processor) loader DataLoader(dataset, batch_sizebatch_size, num_workers4, shuffleFalse) features [] model.eval() with torch.no_grad(): for batch in tqdm(loader): inputs {k: v.cuda() for k, v in batch.items()} embeds model.get_image_features(**inputs) # L2 归一化方便后续直接算余弦相似度 embeds torch.nn.functional.normalize(embeds, dim-1) features.append(embeds.cpu()) return torch.cat(features, dim0)这里get_image_features是 CLIP 模型专门用来抽取图像特征的接口不走完整的 forward少了一些无用的头计算速度快不少。这个代码还能进一步优化如果图片分辨率参差不齐可以在 Dataset 里做预缩放避免 processor 在 batch 推理时做动态 resize省掉不少 CPU 时间。6.2 向量检索库选型从暴力搜索到近似最近邻100 万条数据以内暴力搜索足够快——用 numpy 或者 torch 做矩阵乘法一次查询也就几十毫秒。但当你超过 1000 万条就必须上近似最近邻ANN了。常见的选择是faiss或hnswlib两者的区别在于faiss 的IndexIVFFlat在训练时需要你提供一部分样本做聚类适合大规模离线索引hnswlib 的 HNSW 索引构建慢一些但查询快适合在线高并发场景。import faiss import numpy as np all_image_embeds extract_image_features(...).numpy() # [N, 512] all_text_embeds extract_text_features(...).numpy() # [M, 512] # 建立 faiss 索引512 维向量用 IVF 方式nlist 按经验设为 sqrt(N) dim all_image_embeds.shape[1] nlist int(np.sqrt(all_image_embeds.shape[0])) quantizer faiss.IndexFlatIP(dim) # 内积索引因为特征已 L2 归一化 index faiss.IndexIVFFlat(quantizer, dim, nlist, faiss.METRIC_INNER_PRODUCT) index.train(all_image_embeds) index.add(all_image_embeds) # 查询一条文本向量返回 top10 图片 id scores, indices index.search(all_text_embeds[:1], k10)这个代码里的IndexFlatIP配合已归一化向量等价于余弦相似度检索。如果你用了IndexFlatL2但向量没做归一化结果排序会和余弦相似度不完全一致注意不要在线上环境混用。需要注意的是IndexIVFFlat的查询准确性取决于nprobe参数。nprobe 太小会漏掉最相似的簇导致召回率下降太大查询又慢。经验值是nprobe nlist / 10在这个基础上再×2 到×4 作为稳妥设置。如果你不能接受翻车风险直接用IndexFlatIP暴力搜100 万数据也就几百毫秒未必不可用。6.3 从“能跑”到“能用”一致性哈希与缓存最后提一个工程上很实际的问题——当你的系统同时支持图片查文本和文本查图片时图文两侧的特征必须来自同一个模型版本。一旦模型微调过特征空间变了所有已经入库的旧向量全部失效。我的习惯是给每个微调版本打一个版本号向量库里同时存model_version和feature两个字段。查询时如果请求的模型版本和向量库里的不一致先触发异步的重新特征抽取。这比你在代码里做热切换要省心得多因为特征抽取慢但按时重算在业务上往往是可以接受的。另外一个技巧是缓存文本特征。业务里的查询文本通常高度重复比如商品搜索里翻来覆去就那么几个关键词组合。给文本特征加一层 Redis 缓存key 是文本的哈希值value 是特征向量命中率高了之后检索 QPS 能翻倍。7. 最后一章微调后验证的“三件套”与一个必须养成的习惯模型微调完别急着上线。我每次做基于 CLIP 模型的图像文本跨模态检索项目最后都会跑一遍三件套验证人工抽检、量化评估、回归测试。人工抽检就是随机挑 50 个查询自己看结果量化评估算 Recall10 和 Mean Rank回归测试则是拿上一版模型的特征库跑同一批查询对比排序差异确保新版不出现大面积退步。def evaluate_recall_at_k(query_embeds, gallery_embeds, labels, k10): # query_embeds: [Q, D], gallery_embeds: [G, D] # labels: [Q, G] 的 0/1 矩阵1 表示匹配 scores query_embeds gallery_embeds.T # [Q, G] _, topk_idx scores.topk(k, dim-1) recall 0.0 for i in range(len(query_embeds)): positive set((labels[i] 1).nonzero(as_tupleTrue)[0].tolist()) hit len(positive.intersection(topk_idx[i].tolist())) recall hit / max(len(positive), 1) return recall / len(query_embeds)这段计算 Recall10 的代码看起来简单但有个坑labels矩阵如果太稀疏每个查询只能匹配 1 个目标召回率的数值会普遍偏低因为哪怕 top10 里只对一个也算一次命中。业务上更常看的其实是 Hits1 和 MRR倒数排名。评估时不要只看一个指标被老板问“检索效果到底行不行”时你能同时拿出 Recall10、Hits1 和 MRR 三个数字说服力会强得多。我的习惯是每次微调完都把验证结果、模型版本、数据版本、关键超参数记在一个 markdown 文件里哪怕当时不觉得有用两周后追问题的时候它就是你的后悔药。跨模态检索这个方向问题往往不出在模型而出在版本管理。希望你比我少走这个弯路希望这篇对你有帮助。本文还有配套的精品资源点击获取
返回列表