ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

CLIP提示学习实战:CoOp源码全链路解读与少样本调优

CLIP提示学习实战:CoOp源码全链路解读与少样本调优 前阵子把CoOp的源码从头到尾捋了一遍算是对提示学习这个方向在视觉语言模型上的落地方式有了更实的感受。网上讲CoOp思路的文章不少但能把代码一行一行说清楚的其实不多。这篇就把我读代码时的完整链路整理出来包括PromptLearner的构造、初始化、前向传播、训练更新以及我在复跑实验时踩过的几个坑希望能给正在啃这个仓库的朋友省点时间。这里先说明一下阅读前提你需要对CLIP的基本用法有一点概念至少要清楚图像编码器和文本编码器分别输出什么。如果完全不熟可以先花半小时跑一下OpenAI的CLIP官方demo手写一个“猫的照片”和图像算相似度的例子再回来看CoOp会顺畅很多。1. 先说清楚CoOp到底解决的是哪件事1.1 CLIP本身很能打但Prompt工程让它很别扭CLIP的原理说起来不复杂把图像和文本分别编码到同一个语义空间里然后用余弦相似度衡量“这张图”和“这句话”搭不搭。训练数据是海量的图文对所以它天然具备很强的零样本能力。比如你想让它做图像分类不需要训练一个分类头只需要构造一句文本提示即可。最常见的做法是构造一个模板a photo of a {class}把类别名填进去然后让CLIP分别计算图像特征和这些文本特征的相似度选分数最高的那个类别作为预测结果。这里{class}对应的文本提示写得好不好会直接影响准确率这就是所谓的Prompt Engineering。真实项目里你会发现这个环节非常考验运气和经验。比如“dog”这个类别a photo of a dog可能效果不错但换成a photo of a {class}, a type of pet效果可能更好到了医学影像数据集模板又得变成a histopathology slide of {class}。每个数据集都靠人工去试模板成本高、收益不稳定而且同一个模板在不同数据集上的表现差异可以非常大那种“换个词涨三个点”的情况并不少见。1.2 手工模板的上限迟早会被碰到人工设计模板还有一个更隐蔽的问题你只能离散地尝试有限几种写法无法确定当前写法是不是最优解。假设上下文一共有4个词的位置每个位置从3000个候选词里挑组合空间已经大到不可枚举人根本不可能靠直觉去穷举。所以CoOp的思路其实是把问题换了一个角度不要把提示词看作一串固定的文本而是把它看作一串可学习的连续向量用梯度下降去优化这组向量让模型在特定数据集上自己找到最合适的“提示语”。它把离散的模板搜索变成了连续空间里的参数优化。这也能解释为什么它叫Context Optimization优化的是上下文向量的值而不是文本本身。1.3 一句话概括CoOpCoOp就是把CLIP文本编码器输入的提示模板里的上下文词从可读的单词换成一组连续向量然后在下游任务的少量样本上通过反向传播来学习这组向量最终让CLIP在少样本场景下获得比手工模板更高的分类准确率。代码层面来说核心就是自定义了一个可学习的参数再把它插入到token embedding序列的对应位置里剩下的计算流程基本沿用CLIP。2. 先从整体结构看懂CoOp的训练闭环2.1 官方仓库的目录与核心文件CoOp的官方代码基于Dassl框架Dassl是一个风格化工具箱封装了数据加载、训练循环、配置系统。如果你之前没接触过第一眼看到目录结构可能会有点晕。但真正核心的就那么几个文件prompt_learner.py定义PromptLearner类负责把上下文向量初始化、拼接、输出文本特征这是整个CoOp最关键的文件。clip.py对OpenAI的CLIP模型做了适配重点是把TextEncoder单独拆出来方便接收来自PromptLearner的输入。coop.py训练器里面是完整的训练循环、损失计算、参数更新逻辑。linear_probe.py线性分类头的baseline实现用来和CoOp做对比实验。如果你只是想理解CoOp完全可以忽略Dassl里其他杂七杂八的工具模块专注读这三个文件就够了。我当时读的时候就是按这个顺序来的先看clip.py里的TextEncoder长什么样再看PromptLearner怎么把上下文拼进文本序列最后看trainer里反向传播回传到哪个参数。2.2 一个关键前提CLIP的文本编码器被拆分重写原版CLIP的encode_text方法内部会接收tokenized text经过embedding、位置编码、Transformer、LayerNorm、投影矩阵一整套流程走完输出文本特征。CoOp在clip.py里把这段逻辑单独封装成了TextEncoder类输入不再是token ids而是一个拼接好的embedding序列也就是prompts。你可能会问为什么不直接用原版接口而是非要重写一个TextEncoder因为原版接口输入是离散的token id而CoOp需要把可学习的连续向量插到embedding序列中间。如果你先让token id过一遍token_embedding再替换中间位置那就必须把embedding层和Transformer之间的流程拆开否则没法插入参数。代码里TextEncoder的forward接收的是prompts和tokenized_prompts两个参数前者是已经拼好的embedding输入后者是原始的token id用于找到EOT位置也就是序列结束标志的位置用来取最终特征。2.3 训练循环里到底谁在更新训练阶段的数据流是这样走的图像经过CLIP的图像编码器得到image features这个编码器全程冻结不参与梯度更新。类别名称通过PromptLearner转换成一整组文本特征的embedding再经过TextEncoder得到text features。image features和text features做余弦相似度得到logits再算交叉熵损失。反向传播时梯度只流回PromptLearner里的context参数CLIP本身不更新。这里有个容易让新手困惑的地方既然CLIP的文本编码器也被冻结那为什么每次前向还会重新计算文本特征因为context参数变了token embedding拼接后的序列也就变了所以即使Transformer权重不变每次输入到Transformer的序列内容也完全不同输出自然跟着变。梯度虽然不经过Transformer内部的参数但反向传播会一路传回context向量本身这个路径必须保持畅通。3. 手把手拆解PromptLearner的构造与初始化3.1 上下文向量的两种出生方式PromptLearner的构造函数是理解整个仓库的钥匙。它第一个决定性操作就是初始化context向量。官方代码支持两种方式。第一种从文本模板起步。如果配置里指定了CTX_INIT比如常见的a photo of a代码会把这个字符串拆成单词每个词用CLIP的token_embedding转成向量也就是取token embedding矩阵里对应位置的向量作为context向量的初始值。这个做法的好处是起点在语义上有意义模型从这个基础上再去调整训练初期比较稳定。第二种完全随机初始化。如果不给模板代码会创建一个形状为[n_ctx, ctx_dim]的空张量用标准差0.02的正态分布填充。这里的ctx_dim对应CLIP文本分支的embedding维度OpenAI的ViT-B/32模型是512维ViT-L/14则是768维这个维度必须和模型匹配否则后面拼接时会直接报错。实际使用中我发现从模板初始化通常收敛更快准确率也更高一些尤其当数据集本身和ImageNet分布比较接近时。但随机初始化也不是不能用它在某些domain shift场景下反而能逃出局部最优。建议你想验证CoOp核心机制时直接把CTX_INIT设为a photo of a跑一个benchmark再换成随机初始化跑一遍对比一下差异会理解得更深。3.2 看懂n_ctx、CSC、class_token_position这三个超参数决定了context向量的组织方式。n_ctx是上下文向量的个数换句话说就是可学习tokens的长度。官方实验里比较常用的是4、8、16。n_ctx太小模型没有足够的自由空间去拟合下游任务n_ctx太大训练参数量变大也可能稀释掉类别本身的语义信息。论文里的tendency是8或16在多数数据集上表现不错但没有一个值能通杀所有情况。CSC的全称是Class-Specific Context它决定每个类别是不是拥有自己独立的一组context向量。默认False也就是所有类别共享同一组context向量这在数学上可以理解为所有类别共用同一个“提示语”只是末尾的类别名不同。CSC设为True时每个类别各有一组独立的context向量参数量从n_ctx维变成n_cls乘以n_ctx维。这个设置通常能提升模型在域外数据上的泛化能力代价是训练更慢、更容易过拟合。代码实现里如果CSC为Trueforward内会把self.ctx扩展成[n_cls, n_ctx, ctx_dim]否则共享同一份参数。class_token_position严格来说是一个字符串配置决定类别名token插入的位置。常见的是end也就是前缀context 类别名 后缀论文里还有一个front选项类别名放在最前面。官方实验大部分用的end模式实际也是这个模式效果最稳。3.3 token_prefix和token_suffix里存的是什么这两组buffer是PromptLearner里容易把人绕晕的地方其实理解以后很简单。CLIP对文本的处理第一步是把文本token化序列开头固定是SOT标记也就是sequence start标记token id为49406结尾固定是EOT标记token id为49407。PromptLearner构造时会把每组prompt文本比如a photo of a dog.先token化再经过token_embedding转成完整的embedding序列然后做切片序列里的第一个token单独切出来就是token_prefix序列末尾1 n_ctx位置之后的所有token切出来就是token_suffix。也就是说token_prefix保存的是每一行prompt开头的那个SOT向量token_suffix保存的是类别名后面的一小段后缀向量中间空出来的n_ctx个位置留给后续拼接可学习的context向量。这样做的好处是每次只动态拼接context前面和后面的embedding都是预先算好的省掉了重复计算也减少了显存占用。构造阶段通过register_buffer保存这两组张量不参与梯度更新也不需要反复计算。3.4 一个关于dtype的细节CLIP模型默认用的是float16OpenAI发布的权重是这个精度。PromptLearner在初始化时会把所有张量统一转成和clip_model.dtype一致的类型否则后续拼接时容易出现dtype不匹配的报错。实际写代码时如果你想用float32做全精度微调需要把整个CLIP和PromptLearner的dtype都改成float32只改一边某些情况下会报错某些情况下会静默产生精度损失这种隐形问题排查起来很痛苦。4. 前向计算里的细节文本特征是怎么拼出来的4.1 一次完整的forward过程PromptLearner的forward方法逻辑不复杂但位置要看得细。以默认的end位置、非CSC模式为例代码大致是def forward(self): ctx self.ctx if ctx.dim() 2: ctx ctx.unsqueeze(0).expand(self.n_cls, -1, -1) prefix self.token_prefix suffix self.token_suffix prompts torch.cat([prefix, ctx, suffix], dim1) text_features self.text_encoder(prompts, self.tokenized_prompts) return text_features这里ctx的初始形状是[n_ctx, ctx_dim]unsqueeze(0)变成[1, n_ctx, ctx_dim]再用expand扩展到[n_cls, n_ctx, ctx_dim]。为什么非CSC模式下也要展开成n_cls份因为后面要和n_cls行文本的prefix、suffix对齐拼接。由于是共享的contextexpand复制出来的每一行内容相同但维度必须匹配才能执行cat。拼接后的prompts形状是[n_cls, seq_len, ctx_dim]其中seq_len 1 n_ctx suffix_len。这个张量直接传入重构后的TextEncoder和positional embedding相加进入Transformer堆叠。注意这里不是先走一遍token embedding再相加而是PromptLearner在外部已经把embedding环节做完了TextEncoder里的第一步就是位置编码。4.2 EOT位置的含义与特征抽取Transformer输出后代码会用tokenized_prompts.argmax(dim-1)取出每个序列里EOT token的位置。CLIP原文的做法是取EOT位置对应的特征作为整句话的表征因为自助语言模型里EOT位置的特征在训练时被设计成汇总前面所有token的信息就像是序列末尾的hidden state。CoOp沿用了这个约定不需要改动。这里有一个细节值得注意因为tokenized_prompts的长度是固定的实际文本中间会出现很多padding token也就是值为0的空位EOT通常在有效文本的末尾。argmax取的是最后一个非padding位置的下标这个操作的前提条件是EOT的token id大于其他padding token id在CLIP的tokenizer里成立所以不会取错位置。如果你换了其他tokenizer这套代码不一定能直接复用。4.3 为什么输出的是text_features而不是logitsPromptLearner的forward返回的是所有类别的text features而不是最终的分类logits。分类logits是在训练器里用image features和text features做矩阵乘法或者更准确地说是余弦相似度之后才得到的。这样设计的好处是模块职责清晰PromptLearner只负责生成“提示后的文本特征”相似度计算和损失计算留在外层。如果你想把CoOp迁移到其他任务比如检索或者图文匹配只需要复用PromptLearner生成文本特征的部分不需要改模型结构。我在自己的项目里就是把PromptLearner单独打包出来接入一个图文检索模型复用性很好。4.4 一个容易出错的地方类别名的预处理类别名在进入tokenizer之前要先做一次清洗。原始数据集里的class name可能带下划线比如golden_retriever代码里会先替换成空格变成golden retriever。这步很重要因为CLIP的tokenizer在遇到下划线时可能把它当成一个独立token导致category语义被割裂。还有一些数据集里的类名带数字编号或者特殊符号最好都提前去掉否则会影响后续文本特征的质量。5. 训练与损失传播哪些参数在更新哪些被冻住5.1 冻结策略的实现方式CoOp在训练阶段会冻结CLIP的整个backbone包括图像编码器和文本编码器只允许更新PromptLearner中的context参数。具体实现通常是这样先将CLIP模型参数全部设置为requires_grad_(False)然后单独将PromptLearner里的ctx参数设为requires_grad_(True)最后把ctx传给优化器。有个坑我要单独说一下部分PyTorch版本中如果你直接model.requires_grad_(False)然后再构造优化器同时只传入ctx参数这是没问题的但如果你错误地把所有text encoder参数都传入优化器训练会变得极其不稳定甚至迅速发散。因为CLIP在float16下大参数量的Transformer在低精度条件下做梯度更新对学习率极其敏感很容易直接nan。这也是CoOp官方代码冻住全部backbone的深层次原因之一为了稳定。5.2 相似度计算与交叉熵的写法训练器里的核心计算可以简化为image_features image_encoder(images) image_features image_features / image_features.norm(dim-1, keepdimTrue) text_features prompt_learner() text_features text_features / text_features.norm(dim-1, keepdimTrue) logits 100.0 * image_features text_features.T loss F.cross_entropy(logits, labels)这里的100.0是CLIP原文里带的logit scale不是随便设的。它相当于温度系数temperature 0.01在训练时用来放大相似度之间的差距。你可以试着改小这个值比如改成20会发现softmax输出变得更平缓训练收敛变慢最终准确率明显下降。另一个细节是因为batch内所有图像和所有类别的text features做全量相似度矩阵如果类别数量很大比如1000类text features矩阵一次就是[1000, 512]这个在显存上完全没压力因为文本分支本身只跑一次前向。但如果你在训练时每个step都重新计算全部类别的文本特征而类别特别多上下文长度又长文本Transformer的前向也并非完全可以忽略实际会占几百MB到1GB不等的显存。5.3 学习率的选择逻辑官方配置里PromptLearner的context参数通常使用比较大的学习率常见的是0.002或0.02而冻结参数的模型完全不受学习率影响。为什么可以设这么大因为待更新的参数只有几千到几万个参数量小梯度信号相对干净不容易震荡。我有一个经验如果你从模板初始化学习率偏大问题不大但如果你用随机初始化学习率最好从0.002往下调否则前期loss可能跳得厉害。另外优化器官方用的SGD这在提示学习场景下比Adam更常见部分实验里SGD的收敛效果和泛化效果都更好。如果你想用Adam建议学习率降到原来的十分之一左右。5.4 少样本的batch设置与数据增强CoOp主打的场景是少样本比如每个类别只给1、2、4、8、16张图。训练时通常用全部训练样本但每个step的batch size设置会影响BN等层面的行为好在CLIP的图像编码器用的LayerNorm不受batch size影响。数据增强方面的做法和普通图像分类任务类似随机裁剪、随机水平翻转测试时resize到224再中心裁剪。有一个实战建议如果每个类只有1张或2张样本训练集总量非常少这时候建议关掉RandomResizedCrop这种过强的数据增强改用固定Resize加轻微扰动否则模型会很难拟合loss长时间降不下去。我最初用16-shot跑StanfordCars时loss曲线一直起伏后来排查发现是数据增强太激进样本不够学换弱增强后整个训练稳定了很多。6. 实验效果与避坑经验盘点6.1 官方实验结论的直观印象CoOp论文里的核心实验结论可以概括成三点。第一在绝大多数少样本设定下CoOp显著优于手工模板提示尤其当样本量在4-shot到16-shot时提升幅度通常有1到3个点在细粒度分类数据集上差距更明显。第二CoOp和linear probe对比时优势同样明显因为linear probe只用了CLIP的图像特征而CoOp额外利用了文本语义空间的先验知识。第三CoOp在域外泛化上反而存在一定退化比如在ImageNet上训练后直接测试ImageNetV2或ImageNet-Sketch效果不如手工模板这也直接催生了后来的CoCoOp。回到代码层面你会发现这些结论并不难理解。CoOp是在特定数据集上优化context所以它学到的是偏向这个数据集的语义模式天然存在过拟合问题。这也是为什么CSC模式有时在域外表现更好因为它每个类有独立的context至少在类别之间保留了更多区分度。6.2 我复跑时踩过的几个问题第一个问题是不收敛。现象是loss一直在3.4左右也就是接近均匀分布的交叉熵值但准确率始终不涨。排查到最后发现是CLIP模型被加载成eval模式而PromptLearner却处于训练模式同时学习率设成了0.2高得离谱。这里提醒一下CLIP模型的dropout默认比一般模型少但如果数据量小dropout仍然会影响训练稳定性最好保持CLIP权重的一致配置不要手动额外加dropout。第二个问题是文本特征和图像特征的对齐轴搞反。CLIP原版代码中图像特征是从image_features text_features.T得到每张图对每个类别的logits但如果你从HuggingFace的transformers库加载CLIP它的输出顺序和OpenAI官方版本有差异导致train和test阶段特征轴不一致结果居然有一半类别准确率为0。这个问题尤其阴险因为loss看起来是下降的但那是模型在用一个错误的配对方式做拟合纯属在噪声上学出表面规律。第三个问题是类别顺序必须和label编码严格一致。Dassl框架的数据集类会按类别名字母排序构成classes列表如果你的自定义数据集不是按字母排序训练时构建text features的顺序就会和label对应不上准确率会低到异常。建议在训练前打印一遍text prompt列表人工对照一下确保每个类别名和标签索引对齐。第四个问题是显存占用。如果类别数特别多比如专家类别n_ctx设成32文本Transformer的一次前向并不便宜。官方实现是一次性把全部类别的prompts都送进TextEncoder如果显存不够可以用循环分批计算text features然后再拼接。但要注意batch size过大可能导致text features的数值不稳定特别是float16下建议对text features做layer norm或者保持默认不要频繁改精度。6.3 一份可以直接照抄的配置参考如果你只是想快速跑通CoOp在自己的数据集上的效果可以参考下面这些配置都是我在复跑中验证过比较稳定的组合配置项推荐值备注预训练模型ViT-B/32速度和效果最平衡适合调试输入分辨率224用336需要额外装载大分辨率权重n_ctx16默认值多数数据集都比较稳CTX_INITa photo of a语义起点好收敛快CSCFalse先跑共享版本再试CSC对比优化器SGD学习率0.002动量0.9训练epoch数50少样本下50轮基本收敛过多会过拟合batch size32如果显存不够降到16影响不大数据增强随机裁剪加水平翻转对少样本任务来说足够这些配置不是银弹但可以作为你调试的起点。以我自己在OxfordPets、Flowers102和Food101上的经验这套配置能稳定复现论文报告的大致水平。想再进一步提升可以试试直接用FP16训练并开启混合精度速度提升明显但对准确率的影响很小在少样本场景下甚至偶尔有微弱提升。最后聊几句个人体会CoOp这套代码读完之后我最大的感受是提示学习的代码实现并没有想象中复杂真正的难点反而在理解CLIP本身的文本处理流程以及想清楚哪些参数可以更新、哪些必须冻住。只要把SOT、EOT、token embedding、positional embedding这条链路摸顺了PromptLearner的代码其实读起来非常直白。如果你打算在CoOp方向上继续深入我建议下一步读一读CoCoOp的代码它的核心改动是引入了一个轻量级网络来动态生成context向量相当于把CoOp里的静态参数升级成条件生成。有了CoOp的底子再看CoCoOp会轻松很多。我当时就是这样顺着读下来的把两个仓库放在一起对比Prompt Learning这个方向的演进脉络会清晰很多。
返回列表