AI CLIP踩坑记录
2024 年初接了个电商商品图分类的活。常规路子是攒标注、训分类头、调参,一个新类目就得重训一轮,周期少说两周。
OpenAI 的 CLIP 论文当时很火,号称不用微调就能零样本分类。我的第一反应和大多数人一样:NLP 里的零样本都一堆坑,多模态凭什么突然神了?但还是想试试,万一真能省掉标注呢。
第一次接触 CLIP
那是 2024 年初,我需要为一个电商项目做商品图片分类。传统做法是准备标注数据、训练分类模型、迭代调参。周期至少两周,而且每增加一个新类别就要重训一遍。
偶然看到 OpenAI 的 CLIP 论文,号称"无需任何微调就能做零样本分类"。我当时第一反应是:这又是营销话术吧?零样本这种东西在 NLP 里都还有一堆坑,怎么到了多模态就突然神了?
但还是试试看。先搭个最简单的环境:
# 安装依赖
pip install torch torchvision transformers ftfy regex tqdm
pip install git+https://github.com/openai/CLIP.git
跑个最基础的 demo:
import clip
import torch
from PIL import Image
device = "cuda" if torch.cuda.is_available() else "cpu"
model, preprocess = clip.load("ViT-B/32", device=device)
# 加载一张图
image = preprocess(Image.open("product.jpg")).unsqueeze(0).to(device)
text = clip.tokenize(["a photo of a red dress", "a photo of a blue shirt", "a photo of black pants"]).to(device)
with torch.no_grad():
image_features = model.encode_image(image)
text_features = model.encode_text(text)
logits_per_image, logits_per_text = model(image, text)
probs = logits_per_image.softmax(dim=-1).cpu().numpy()
print("Label probs:", probs)
输出结果让我愣了一下——分类准确率比我想象中高太多。就这样,我被 CLIP 的能力说服了,决定深入了解它到底是怎么做到的。
对比学习:CLIP 的核心思想
CLIP 的本质是对比学习(Contrastive Learning),但不是传统意义上那种单模态的对比。
简单说:把图片和文本分别编码成向量,然后让"匹配"的图文对在向量空间里靠得更近,让"不匹配"的图文对互相远离。
训练时,一个 batch 里有很多图文对。比如 batch size 是 64,就有 64 张图和 64 段描述。对于其中一张图 A 和它对应的描述 A’,我们希望它们的相似度最高;而对于描述 A’ 和其他 63 张图,相似度都要低。
这就像是在开派对,你得记住自己带的人,别跟别人混在一起。
这个图的重点是:真正要拉近的只有"正确配对",其他所有组合都是要被推开的负样本。
实际训练时,CLIP 用的 loss 函数是对比损失的一个变种,叫 symmetric cross entropy loss。简单理解:同时优化"图找文"和"文找图"两个方向。
# 伪代码,展示对比损失的核心思想
def contrastive_loss(image_features, text_features, temperature=0.07):
# 归一化
image_features = image_features / image_features.norm(dim=1, keepdim=True)
text_features = text_features / text_features.norm(dim=1, keepdim=True)
# 计算相似度矩阵
logits = (image_features @ text_features.T) / temperature
# 对角线是正样本,其他是负样本
batch_size = image_features.shape[0]
labels = torch.arange(batch_size)
# 两个方向的交叉熵
loss_i = F.cross_entropy(logits, labels)
loss_t = F.cross_entropy(logits.T, labels)
return (loss_i + loss_t) / 2
这个 temperature 参数很有意思。它控制的是分布的"尖锐程度",值越小,模型越自信;值越大,模型越保守。CLIP 论文里用的是 0.07,这是在 WebImageText 数据上调出来的经验值。
实际部署遇到的坑
理论上理解了,实际用起来还是要踩坑。
显存问题
CLIP 的 ViT-B/32 模型不算大,但推理时如果 batch size 太大会爆显存。我第一次在生产环境部署时,默认用了 32 的 batch size,结果在 T4 上直接 OOM。
解决办法很直接:调小 batch size 或者用 gradient checkpointing(如果需要训练)。但对于纯推理场景,调小 batch size 就够了:
# 原来
with torch.no_grad():
image_features = model.encode_image(images) # 32张图
# 改成
with torch.no_grad():
image_features = []
for i in range(0, len(images), 8):
batch = images[i:i+8]
image_features.append(model.encode_image(batch))
image_features = torch.cat(image_features)
预处理也别忽视。CLIP 要求输入图片是 [3, 224, 224] 的 float tensor,值域在 [0, 1] 之间,还要做标准化。如果你用的是 OpenCV 读取的 BGR 格式,记得转 RGB:
import cv2
import numpy as np
image = cv2.imread("product.jpg")
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
image = Image.fromarray(image)
image = preprocess(image)
文本提示词的选择
零样本分类的效果高度依赖你怎么写提示词。我一开始用的是很简单的描述:
prompts = ["dress", "shirt", "pants", "shoes"]
结果准确率只有 60% 左右。后来参考了论文里的做法,加上"A photo of"前缀:
prompts = [
"A photo of a dress",
"A photo of a shirt",
"A photo of pants",
"A photo of shoes"
]
准确率一下提升到 85%+。
再后来发现,提示词还可以更精细。比如对于红裙子,可以写:“A photo of a red dress, fashion, clothing”;对于场景分类,可以写:“A photo of a beach, ocean, sand, summer”。
这些额外信息不是堆砌,而是在给模型更多"钩子"去匹配。CLIP 训练时见过大量类似描述,它会把这些词和对应的视觉特征关联起来。
数据规模偏差
这是我踩过最大的坑。CLIP 是用 WebImageText 训练的,4 亿对图文对。这些数据主要来自网络,天然偏向某些类别和风格。
我做过一个产品分类任务,用 CLIP 跑零样本,发现电子产品分类很差,但服装、家居类就很好。查了一下,WebImageText 里电子产品相关的图文对确实少,而且大多是营销图,不是实拍图。
解决办法有几个:
- 如果有少量标注数据,可以微调 CLIP 的文本编码器
- 用更细粒度的提示词描述
- 直接训练一个小模型,用 CLIP 的输出作为特征
我选了第三种,效果最稳定。CLIP 作为特征提取器,后面接一个简单的分类层:
import torch.nn as nn
class CLIPClassifier(nn.Module):
def __init__(self, clip_model, num_classes):
super().__init__()
self.clip_model = clip_model
self.classifier = nn.Linear(512, num_classes) # CLIP 的特征维度是 512
def forward(self, images):
with torch.no_grad():
features = self.clip_model.encode_image(images)
return self.classifier(features)
这样只需要训练最后的分类层,很快就能收敛。
零样本之外的更多玩法
零样本分类是最直接的用法,但 CLIP 能做的远不止这些。
图文检索
这是 CLIP 的原生能力。给定一张图,从大量文本里找到最匹配的;或者反过来。
# 图到文检索
def image_to_text_search(query_image, text_database, model):
image_feature = model.encode_image(preprocess(query_image).unsqueeze(0))
text_features = model.encode_text(clip.tokenize(text_database))
similarities = (image_feature @ text_features.T).squeeze(0)
best_idx = similarities.argmax().item()
return text_database[best_idx]
我做过一个内部知识库的图片检索功能,用 CLIP 直接就能跑通,连训练都不需要。用户上传一张截图,系统自动从文档库找出最相关的页面。
异常检测
这是个很有意思的应用。正常商品图片会和描述高度匹配,而异常图片(比如假货、瑕疵品)在 CLIP 的空间里会和所有文本描述的距离都比较远。
def anomaly_detection(image, normal_descriptions, threshold=0.25):
image_feature = model.encode_image(preprocess(image).unsqueeze(0))
text_features = model.encode_text(clip.tokenize(normal_descriptions))
similarities = (image_feature @ text_features.T).squeeze(0)
max_sim = similarities.max().item()
return max_sim < threshold
阈值需要根据实际数据调,但逻辑很简单:如果连最匹配的描述都不够"匹配",那这张图就有问题。
图像生成引导
CLIP 还能用来引导图像生成。比如用 Stable Diffusion 时,用 CLIP 的文本编码器来确保生成结果更贴合描述。
这个我没做过实际项目,但原理上就是:生成过程中不断计算生成图像和目标描述的 CLIP 相似度,用它作为 loss 的一部分。
一些实践中的经验
用了 CLIP 大半年,积累了一些经验:
模型选择不是越大越好。ViT-B/32 在大多数任务上够用了,ViT-L/14 虽然精度更高,但推理慢很多、显存占用也大。如果你的应用对延迟敏感,优先考虑小模型。
多语言支持有局限。CLIP 原生只支持英文,用中文提示词效果会差很多。如果必须用中文,可以考虑:
- 用机器翻译把中文转成英文
- 找支持中文的 CLIP 变体(比如 Chinese-CLIP)
预训练很重要。CLIP 的强大来自于它见过 4 亿对图文。如果你的任务和这些数据差异太大(比如医疗影像),零样本效果可能一般。这时候要么微调,要么找领域特定的模型。
后处理很关键。CLIP 给出的是概率分布,但实际应用中往往要做二次处理。比如设置阈值过滤低置信度结果,或者用规则修正明显错误的分类。
最后
CLIP 最大的价值不是它有多准,而是它改变了我们处理多模态问题的方式——不再需要为每个任务从头训练一个模型,而是用预训练的跨模态表示作为起点。
但这不代表 CLIP 是万能的。数据偏差、领域适应性、语言限制,这些都是真实存在的问题。技术选择从来都是在成本、效果、可维护性之间找平衡。
现在回头看,那段时间折腾 CLIP 的经历让我明白:工具再强大,最终还是要回到具体场景、具体问题、具体限制里去用。从来没有"银弹",只有"合适"。
这大概就是技术的魅力所在吧。
版权声明: 本文首发于 指尖魔法屋-AI CLIP踩坑记录(https://blog.thinkmoon.cn/post/254-clip-contrastive-learning-multimodal-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。