ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

CLIP 零样本图像分类 5 分钟跑通:每类 12 张图即可上线

CLIP 零样本图像分类 5 分钟跑通:每类 12 张图即可上线 CLIP 零样本图像分类 5 分钟跑通每类 12 张图即可上线【免费下载链接】CLIPCLIP (Contrastive Language-Image Pretraining), Predict the most relevant text snippet given an image项目地址: https://gitcode.com/GitHub_Trending/cl/CLIP人工标注一张图约 3 秒一个类别标 5000 张要占掉大半天而项目刚立项时往往只有十几张参考图。CLIPContrastive Language-Image Pretraining用对比学习成对训练图文把两者映射到同一向量空间靠余弦相似度判断相关性——用文字描述类别就能做零样本图像分类。30 秒自检你的场景该不该上 CLIP选型前先回答 4 个问题你的类别是 1几百个且会频繁新增是零标注还是每类只有十几张图希望新类别当天可用而不是长期微调是否在追上千类 top-1 极致精度或亚像素判别前 3 个答是、第 4 个答否CLIP 就是对的工具适合不适合替代方向1几百类频繁新增类别上千类 top-1 极致精度、监督精调零标注按文本分类、按文本找图像素级分割、目标定位、计数每类个位数样本起步类别边界极细、需亚像素判别新类别当天上线一次性交付、长期冻结一句话CLIP 擅长新、杂、少的分类精细几何与大规模细粒度识别仍是监督模型的主场。5 分钟上手从安装到第一次推理依赖只有 6 个torch、torchvision、ftfy、regex、tqdm、packaging完整清单见 requirements.txt。git clone https://gitcode.com/GitHub_Trending/cl/CLIP pip install -e CLIPimport torch import clip from PIL import Image device cuda if torch.cuda.is_available() else cpu # 首次运行自动下载约 335MB 权重到 ~/.cache/clipCPU 上自动回落 float32 model, preprocess clip.load(ViT-B/32, devicedevice) names [good solder joint, insufficient solder, solder bridge] # 统一包成完整英文短语裸类名缺少这是一张图的语义锚点 prompts [fa photo of a {n} solder pad on a circuit board for n in names] img preprocess(Image.open(pad_001.jpg)).unsqueeze(0).to(device) txt clip.tokenize(prompts, truncateTrue).to(device) # 文本张量默认在 CPU with torch.no_grad(): logits model(img, txt)[0] # 余弦相似度 × 可学习温度系数 print(logits.softmax(-1).cpu().numpy())输出是一组和为 1 的概率例如[0.03, 0.14, 0.83]最大值对应预测类别。这个数字能直接告诉你模型好不好使。想交互式验证完整流程可以跑 notebooks/Interacting_with_CLIP.ipynb。拆解核心能力跑通之后要拿下的 3 件事写出概率不失真的类别描述提示词质量对小样本分类的影响超过换模型。规则就三条所有类共用同一句式、只替换类名用完整短语而非裸词23 个模板做集成取平均templates [ a photo of a {} solder pad on a circuit board, a close-up of a {} solder joint, ] probs [] for t in templates: # 多模板集成各模板 softmax 概率取平均 txt clip.tokenize([t.format(n) for n in names], truncateTrue).to(device) with torch.no_grad(): probs.append(model(img, txt)[0].softmax(-1).cpu()) print(torch.stack(probs).mean(0).numpy())多模板集成通常能稳定抬高 top-1 几个点。现成模板在 data/prompts.md覆盖 26 个数据集的官方类名与句式。批量推理怎么配吞吐最高逐张调用吞吐远低于批量。关键两条图像走批处理文本做缓存# 类别集固定时文本嵌入算一次存下来之后每批只跑图像侧 text_feats model.encode_text(clip.tokenize(prompts, truncateTrue).to(device)) text_feats text_feats / text_feats.norm(dim-1, keepdimTrue) img_feats model.encode_image(img_batch) # img_batch16~32 张 img_feats img_feats / img_feats.norm(dim-1, keepdimTrue) probs (model.logit_scale.exp() * img_feats text_feats.T).softmax(-1)文本侧 1 次编码换图像侧 32 次编码叠加半精度后单张推理能压到 2030ms 量级。单张图怎么找出最相关的文本这是 CLIP 的本行给一张图输出候选文本里最相关的一条用 topk 取前 5text_feats model.encode_text(clip.tokenize(candidates, truncateTrue).to(device)) text_feats / text_feats.norm(dim-1, keepdimTrue) img_feats model.encode_image(img) / model.encode_image(img).norm(dim-1, keepdimTrue) sim (100 * img_feats text_feats.T).softmax(-1)[0] # 100 与 logit_scale 量级一致 values, idx sim.topk(5) print([candidates[i] for i in idx])零样本下 top-1 概率通常超过 60%仓库 CIFAR-100 示例中 snake 约 65%排序结果可直接使用。调优阶梯0 标注到 48 标注的 3 档第 1 档模板网格搜索。成本0 标注、几小时工时预期收益25 个点。best_tmpl, best_acc None, 0.0 for tmpl in candidates: # 准备 3~5 个模板候选 acc eval_on_valset(tmpl) # 在留出验证集上跑 if acc best_acc: best_tmpl, best_acc tmpl, acc第 2 档线性探针。成本每类 1050 张标注、CPU 几分钟训完预期收益515 个点。冻结模型离线提 512 维特征后拟合逻辑回归feats, ys [], [] with torch.no_grad(): for imgs, y in train_loader: feats.append(model.encode_image(imgs.to(device)).cpu()) ys.append(y) from sklearn.linear_model import LogisticRegression clf LogisticRegression(max_iter2000).fit( torch.cat(feats).numpy(), torch.cat(ys).numpy())第 3 档提示调优。成本GPU 加几个 epoch预期收益13 个点。仍冻结全部权重只在文本端学一组提示嵌入prompt torch.nn.Parameter(torch.randn(77, 512, devicedevice) * 0.02) opt torch.optim.AdamW([prompt], lr1e-4) # 77x512 约 3.9 万参数过拟合风险低先做完上一档、跑验证集确认瓶颈再进下一档——第 2 档能吃掉大部分差距第 3 档边际收益通常最小。模型怎么选ViT-B/32、ViT-B/16 与 RN50全部 9 个型号可用 clip/clip.py 里的clip.available_models()查看常用三个对比模型权重体积骨干结构取舍ViT-B/32约 335MB视觉 Transformer速度精度均衡默认首选ViT-B/16约 580MBViTpatch 16零样本精度 12 点推理更慢RN50约 530MBResNet-50CPU 上最快无卡环境优先三条性价比最高的动作半精度GPU 上model.half()显存与耗时约各降一半批处理一次 1632 张不要逐张调用缓存类别集不变时encode_text只算一次端到端案例SMT 焊盘四类筛查每类 12 张场景产线对焊盘分四类——正常、缺锡、锡桥、漏焊每类仅标注 12 张每类留出 8 张做验证集。提示词用双模板集成a photo of a {cls} solder pad on a circuit board a close-up of a {cls} solder joint方案验证集准确率锡桥类召回纯零样本双模板集成81.4%74.3%线性探针48 张标注94.8%96.4%提示调优48 张标注10 epoch95.2%97.2%部署形态特征提取 逻辑回归GPU 批 8 张 fp16约 28ms/张产线节拍内放得下。排障速查新手最容易卡的 5 个点RuntimeError: SHA256 checksum does not match下载中断或文件损坏删掉~/.cache/clip里的 .pt 重跑或把本地 checkpoint 路径直接传给clip.load。Expected all tensors to be on the same deviceclip.tokenize返回的文本张量默认在 CPUtxt clip.tokenize(...).to(device)。这个坑我见过太多次。中文类名概率分布怪异tokenizer 是按英文训练的 BPE中文会碎成不可识别 token用英文类名展示层再映射回中文标签。RuntimeError: Input ... is too long for context length 77描述超过 77 token 上限给clip.tokenize传truncateTrue。CPU 上出现 float16 警告clip.load在 CPU 会自动回落 float32无需手动处理嫌慢就换 RN50。上线前检查清单今天加载 ViT-B/32用 23 个类别跑零样本基线确认模板与语言没问题今天记录零样本准确率作为回归基线本周每类备齐至少 10 张标注图固定验证集本周跑线性探针与零样本对比决定是否继续下一档上线前图像、文本、模型三者设备与精度一致fp16 时文本侧同转上线前开启批处理与文本嵌入缓存实测单张延迟下一步就一件事把第二段代码里的names换成你的类名15 分钟内拿到第一个基线数字。【免费下载链接】CLIPCLIP (Contrastive Language-Image Pretraining), Predict the most relevant text snippet given an image项目地址: https://gitcode.com/GitHub_Trending/cl/CLIP创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表