
CLIP零样本分类概率不对角对比学习机制拆解与3步诊断法【免费下载链接】CLIPCLIP (Contrastive Language-Image Pretraining), Predict the most relevant text snippet given an image项目地址: https://gitcode.com/GitHub_Trending/cl/CLIPCLIPContrastive Language-Image Pre-Training把一张图和一组候选文本扔进同一个前向输出一个相似度矩阵再用 argmax 选出最相关的文本片段——这就是它能零样本分类的全部家底。但当你按 README 的 zero-shot 示例跑完看到类似这样的输出时Top predictions: snake: 65.31% turtle: 12.29% sweet_pepper: 3.83% lizard: 1.88% crocodile: 1.75%top-1 和 top-2 只差 4.6 倍top-5 加起来也才 85%——这个软的概率分布不是玄学而是 CLIP 主模型 里一个可学习温度系数和一组 L2 归一化特征共同作用的结果。把CLIP的 forward 从头读一遍会发现它短得出奇顺着相似度矩阵、特征对齐、零样本提示展开这三层把概率链路拆完每层配一段十来行的诊断脚本就能精确定位问题出在温度、特征还是提示上。概率不对角先看 logit_scale 和它乘的那一行CLIP的 forward 做的事可以概括为两个编码器各自把图像、文本映射到同一个向量空间点积得到相似度再乘一个可学习的温度系数softmax 之后就是 README 里打印的那些概率。def forward(self, image, text): image_features self.encode_image(image) text_features self.encode_text(text) # 归一化特征点积从此等价于余弦相似度 image_features image_features / image_features.norm(dim1, keepdimTrue) text_features text_features / text_features.norm(dim1, keepdimTrue) # 余弦相似度当 logits乘上可学习的温度系数 logit_scale self.logit_scale.exp() logits_per_image logit_scale * image_features text_features.t() logits_per_text logits_per_image.t() return logits_per_image, logits_per_text这段代码在 clip/model.py 的CLIP.forward里。三个事实决定了输出的形状第一两边特征都做 L2 归一化点积从此等于余弦相似度被钉死在 [-1, 1]第二logit_scale在 CLIP 构造器 里初始化为torch.ones([]) * np.log(1 / 0.07)是个可学习参数.exp()之后才是真正的温度第三forward 返回的是 logits 而非 loss——对比损失batch 内图像侧、文本侧各做一次交叉熵再平均需要训练方自己补这个仓库只给了推理侧的模型与接口。logit_scale 是温度的倒数旋钮它越大softmax 越尖锐正样本对与负样本对的概率差距被放大越小概率越趋向均匀分布。图像编码器和文本编码器彼此独立跨模态之间唯一可学习的桥梁就是这一个标量所以输出形态不对时它是第一个要看的参数。import torch, clip from PIL import Image model, preprocess clip.load(ViT-B/32, devicecpu) image preprocess(Image.open(CLIP.png)).unsqueeze(0) text clip.tokenize([a diagram, a dog, a cat]) with torch.no_grad(): logits, _ model(image, text) probs logits.softmax(dim-1).squeeze(0).numpy() # 诊断一温度系数与概率尖锐度 print(logit_scale(exp):, model.logit_scale.exp().item()) print(probs:, probs) print(top1-top2 gap:, (probs[0] - probs[1]).round(4))现象可能原因排查方向概率接近均匀分布如 0.25/0.25/0.25/0.25温度太小或特征间余弦相似度差异不足打印logit_scale.exp()检查两条特征模长是否为 1相似度矩阵对角线不突出特征没拉开正负样本温度只是放大了噪声走下一节的对齐检查logit_scale 训练后期持续超过 20过拟合风险噪声对也被推到高分冻结该参数或加权重衰减诊断信号如果logit_scale.exp()落在 5~20 之间、但对角线依然不突出问题就不在温度而在特征继续往下看。特征对齐L2 归一化与 EOT token 的两处对齐图 1 左半部分画的就是这个矩阵的来历N 张图、N 条文本各自过编码器得到一个 N×N 的点积矩阵对角线是正样本对。这个矩阵能不能对角占优取决于两处对齐。图像侧很简单encode_image返回的向量在 forward 里被 L2 归一化落在单位球面上余弦相似度直接可比。文本侧多一步——CLIP的文本输入是 77 个 token 的定长序列padding 用 0 填充而文本特征不是取序列末尾位置而是先定位到最后一个有效 tokenBPE 词表里的/w再投影def encode_text(self, text): x self.token_embedding(text).type(self.dtype) # [batch_size, n_ctx, d_model] x x self.positional_embedding.type(self.dtype) x x.permute(1, 0, 2) # NLD - LND x self.transformer(x) x x.permute(1, 0, 2) # LND - NLD x self.ln_final(x).type(self.dtype) # x.shape [batch_size, n_ctx, transformer.width] # 取 eot embedding 位置的特征eot_token 是序列中编号最大的 token x x[torch.arange(x.shape[0]), text.argmax(dim-1)] self.text_projection return xtext.argmax(dim-1)这个写法初看费解但拆开就是词表顺序的巧合分词器 的SimpleTokenizer把词表构造成 256 个字节加 256 个带/w后缀的字节再依次追加 BPE 合并结果最后才是|startoftext|49406和/w49407。/w的 id 是整张词表里最大的而 padding 全是 0id 0于是 argmax 恒等于最后一个真实 token 的位置——padding 永远不会赢。这个位置过完 Transformer 携带了整句的语义再乘text_projection投进共享空间。对齐质量用一行就能量化正样本对的对角线余弦相似度应该明显高于该行的平均值。with torch.no_grad(): img model.encode_image(image) / model.encode_image(image).norm(dim-1, keepdimTrue) txt model.encode_text(text) / model.encode_text(text).norm(dim-1, keepdimTrue) sim img txt.t() print(diag:, sim.diagonal().round(3).tolist()) print(row means:, sim.mean(dim1).round(3).tolist())诊断信号如果对角线没有比行均值高出一个明显量级例如差值 0.1说明特征空间没把正负样本分开温度系数再大也只是放大噪声另外单独打印img.norm(dim-1)应当恒等于 1若不等说明你在模型外手写相似度时漏了归一化。提示词工程为什么细粒度分类要先展开模板零样本分类的完整调用长这样注意第 5 步用的是写死的常数 100.0 而不是logit_scaleimport clip, torch model, preprocess clip.load(ViT-B/32, devicecpu) image preprocess(Image.open(CLIP.png)).unsqueeze(0) with torch.no_grad(): image_features model.encode_image(image) text_features model.encode_text( torch.cat([clip.tokenize(fa photo of a {c}) for c in [snake, turtle, lizard, crocodile]]) ) # 归一化后点积即余弦相似度乘固定温度 100 再 softmax image_features / image_features.norm(dim-1, keepdimTrue) text_features / text_features.norm(dim-1, keepdimTrue) similarity (100.0 * image_features text_features.T).softmax(dim-1) print(similarity[0].topk(4))similarity的形状是[1, 4]一行四列每个候选文本一个余弦相似度乘 100 拉大 gap 后 softmax。README 的 CIFAR-100 例子就是把 4 个换成 100 个、再topk(5)。这里温度是常数而不是参数因为零样本推理时模型是冻结的100.0 只负责把概率拉开不改变排序。零样本分类的主要失败模式是提示敏感换个措辞概率分布就变。仓库里 data/prompts.md 给 26 个数据集都准备了classes与templates两份清单CIFAR-100 的 18 个模板覆盖了a photo of a {}.、a blurry photo of a {}.、a black and white photo of a {}.、a photo of a small {}.等变体。把类名逐一替换进每个模板每个类就展开出多条文本特征classes [apple, snake, turtle, lizard] templates [a photo of a {}., a blurry photo of a {}., a photo of the {}.] prompts [t.format(c) for c in classes for t in templates] # 4 类 x 3 模板 12 条 with torch.no_grad(): text_features model.encode_text(clip.tokenize(prompts)) text_features / text_features.norm(dim-1, keepdimTrue) # 每个类把 3 条提示的 softmax 概率求平均削弱单模板偏置 scores (100.0 * image_features text_features.T) scores scores.view(len(classes), len(templates), -1).softmax(dim-1).mean(dim1) print(dict(zip(classes, scores[0].tolist())))平均之后单模板的偏置被摊薄一张偏糊的图对a blurry photo模板得分偏高但另外两个模板会把它拉回来。Prompt_Engineering_for_ImageNet.ipynb 里对 ImageNet 的多模板集成就是这个思路的完整版。诊断信号同一张图换两套措辞相近的模板top-1 类别就翻盘说明你踩在提示敏感上直接上多模板平均若类别本身就是细粒度的蛇/蜥蜴/鳄鱼单模板的 65% 顶天平均也救不回语义上的重叠。落地清单与延伸一个五步排查流程按顺序走核对输入链路图像过没过 clip.py 里_transform返回的预处理Resize → CenterCrop → ToTensor → 归一化文本过没过clip.tokenize两处的版本不一致会让相似度整体失真。打印encode_image输出向量的norm(dim-1)确认模长恒为 1不等说明模型外的手写流程漏了归一化。打印相似度矩阵确认对角线正样本对比行均值高出一个明显量级差距 0.1 就转特征侧排查。确认温度在位训练时logit_scale.exp()落在 5~20零样本推理时用的是常数 100.0 而非参数。输出仍偏软上多模板平均data/prompts.md 的模板起步Prompt_Engineering_for_ImageNet.ipynb 参考集成方式。配套资源clip/clip.pyclip.load提供 RN50、ViT-B/16、ViT-L/14 等 8 个预训练模型的下载与加载available_models()可列出全部型号。data/prompts.md26 个数据集的类名与提示模板对照直接拿来当零样本推理的起点。tests/test_consistency.pyJIT 与非 JIT 两条推理路径的一致性校验排查部署差异时可以复用它的思路。回看开头那个 65.31% 的 snakeCLIP 的 forward 里只有一条 matmul 和一个可学习温度系数概率输出不对角的原因就藏在这 20 行里而诊断这件事本身就是打印出这一行里的中间量——温度系数、归一化模长、相似度矩阵的对角线。机制拆到这里调参就不再是碰运气而是对着数改。延伸阅读Radford et al., Learning Transferable Visual Models From Natural Language SupervisionarXiv:2103.00020CLIP 原始论文对比损失形式与 400M 图文对的训练细节都在其中推荐精读第 3 节的损失推导。model-card.md官方模型卡写清了 34 个评测基准上的表现区间与已知局限部署前建议通读。【免费下载链接】CLIPCLIP (Contrastive Language-Image Pretraining), Predict the most relevant text snippet given an image项目地址: https://gitcode.com/GitHub_Trending/cl/CLIP创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考