拓冰建站拓冰建站
首页 / 资讯中心 / 正文

从原理到PyTorch实现:零样本学习完整实战指南

简介这是一份围绕零样本学习的论文复现资源面向机器学习、计算机视觉方向的研习者与中高级开发者。作者参考一篇零样本识别相关论文实现了基础复现与基于PyTorch的原文复现两套代码配套语义空间矩阵、数据载入与KNN分类等脚本构成从数据处理到模型评估的完整实验路径。压缩包共11个文件以6个Python脚本为主体并附2个PDF文档含复现论文与说明、1个xls属性矩阵、1个ipynb交互式笔记及1个txt提示文件整体大小约38.98MB。由于原数据集较大资源包内已给出数据集获取与使用的相关说明下载后按Readme指引准备数据即可运行。该资源已有2850人学习适合想理解零样本学习兼容性建模机制、快速开展复现实验或参考其代码框架的读者。1. 零样本学习到底在解决什么问题先聊点实际的。传统监督学习要做图像分类你得先给每个类别准备几十上百张标注图模型看完这些图才能学会“这是什么”。但现实任务里总有那么几类东西样本极其稀缺甚至压根没有——比如稀有动物识别、新上架商品的分类、特定故障模式的诊断。这时候常规套路直接失效因为你没有带标签的训练数据喂给模型。零样本学习Zero-Shot LearningZSL就是冲着这个痛点来的。它的核心思路很直白既然没有目标类别的图像样本那就用类别本身的语义描述来搭桥。我给你看一张“斑马”的图模型从没见过斑马长什么样但它知道斑马有“条纹”“像马”“黑白配色”这些属性再结合训练时见过的马、老虎、奶牛等动物模型就能推理出“这玩意儿大概率是斑马”。整个过程不需要一张斑马的训练图像这就是“零样本”的含义。这个方向适合谁去研究如果你是搞计算机视觉的研究生、做模型算法的工程师或者对多模态学习感兴趣的技术爱好者都值得深入看一遍。它不光是论文里的花架子在工业界的实用价值也很高——比如电商平台对零销量新品做自动打标、安防系统识别新出现的异常行为、医疗影像对罕见病做初筛这些场景天然适合零样本学习落地。需要说明的是这篇文章会同时讲清楚论文思路和代码实现两个维度。前者帮你建立理论框架后者给你一个可以直接跑起来的PyTorch基线工程。建议你边读论文边对照代码效果会好很多。2. 论文核心思路拆解为什么语义信息能当“桥梁”2.1 从“视觉特征”到“语义特征”的映射逻辑零样本学习的论文虽然每年出一大堆但底层逻辑绕不开同一个框架。我先把这个框架掰开揉碎讲清楚。想象你有一个巨大的特征空间里面住着两类东西一类是视觉特征——就是CNN或者ViT从图像里提取出来的向量比如一张猫的图片经过ResNet编码后变成的2048维向量另一类是语义特征——就是每个类别用一句话或者一组属性描述出来的向量比如“猫”的语义向量可能是“有胡须、会喵喵叫、肉食性、体型小”这些属性的组合。零样本学习干的事情就是在视觉特征和语义特征之间学一个“对齐函数”。训练阶段你有狗、猫、鸟这些可见类的图像和它们的语义描述于是你可以逼着模型学会“一群视觉向量和一群语义向量应该某种程度地对应上”。到了测试阶段你给模型一张“狐狸”的图它没见过狐狸但可以通过已经学到的对齐函数把狐狸的视觉向量映射到语义空间里跟“狐狸”的语义向量算相似度谁得分高就判给谁。这个思路和人类的学习方式确实有点像。你没见过榴莲但别人告诉你“这是一种有刺的、闻起来臭吃起来香的热带水果”你真遇到了大概率能认出来。零样本学习就是把这种“基于属性推理”的能力形式化成数学。2.2 两类设定传统ZSL与广义ZSL这里有个非常重要的区分几乎所有零样本学习论文都会强调刚入门的人特别容易搞混。传统零样本学习Conventional ZSL的设定是测试时只从未见类里找答案。也就是说尽管模型在训练时见过可见类的样本但测试阶段出现的图像只可能属于未见类。这种设定在实践中偏理想化了。广义零样本学习Generalized ZSLGZSL才是更贴近真实世界的设定测试阶段图像既可能来自可见类也可能来自未见类模型要在两者混合的类别空间里做判断。这个难度一下子就上去了——模型天生偏向它见过的类很容易把未见类的图片强行分到某个可见类头上。拿代码实现来说这两种设定对应两套不同的评估逻辑。传统ZSL直接算没见过类上的Top-1准确率就行但GZSL需要分别计算可见类样本的准确率和未见类样本的准确率再算两者的调和平均。调和平均这个指标很关键它惩罚那种“把全部样本都判给可见类”的偷懒做法——如果你这么做未见类准确率就是0调和平均直接归零。2.3 类嵌入Class Embedding的构造方案零样本学习里有个关键组件叫类嵌入class embedding本质就是每个类别对应的语义向量。它的质量直接决定模型上限很多论文就在这个环节做文章。我给你盘点一下主流的构造方式你写代码的时候也要在这里做选择。第一种是属性标注法。这是最经典的做法人工为每个类别标注一组可判别属性比如AwA数据集里每个动物类都标注了85个属性是否有条纹、是否群居、体型大小等。这种方式解释性强、精度高但代价是人工标注成本极大而且属性集合一旦定死就难以扩展。第二种是词向量法。利用word2vec、GloVe或者fastText在大规模文本上预训练好的词向量直接把类别名称映射成向量。这个方案零成本但问题是单词语义太粗糙——“狐狸”和“狼”的词向量可能非常接近导致模型区分不开。第三种是文本描述法也是现在的主流方向。用BERT、GPT这类语言模型把一个类别的详细文字描述编码成向量比如使用CLIP的文本编码器。这种方案既不需要人工标注属性语义信息又比单个词丰富得多尤其是CLIP这类多模态模型出现之后文本描述法的效果有了质的飞跃。从代码实现的角度我建议你直接用CLIP的文本编码器来生成类别嵌入。理由很简单第一CLIP已经在4亿图文对上做过对齐训练拿到的语义向量和视觉向量的关系天然更紧密第二用HuggingFace的transformers库几行代码就能拿到向量不用自己训练语言模型第三CLIP的visual encoder也可以顺手当特征提取器用整个pipeline的代码风格会非常统一。3. 从零搭一个ZSL基线工程基于PyTorch的完整实现3.1 技术选型为什么不推荐自己训视觉特征写代码之前先说个我踩过的坑。最初做零样本学习实验时我想着“端到端训练一个自己的特征提取器应该更灵活”于是用ResNet在ImageNet上从头微调。结果训练时间翻了三倍效果反而不如直接用预训练特征。原因其实不复杂。零样本学习研究的核心矛盾在“视觉-语义对齐”而不是“视觉特征提取”。特征提取这件事大规模预训练模型已经做得足够好了你在这个环节投入算力纯属浪费。正确的做法是把视觉特征提取器当作一个冻结的“特征数据库”你只训练上层的对齐映射模块。所以我最终的选型方案是这样的视觉特征提取使用torchvision自带的ResNet101加载在ImageNet上预训练好的权重语义特征提取使用HuggingFace的CLIP文本编码器用类别描述文本生成嵌入对齐模型一个两层的MLP多层感知机将视觉特征投影到语义空间损失函数基于余弦相似度的对比损失数据集CUB鸟类数据集200类鸟其中150类作为可见类50类作为未见类选CUB是有原因的。这个数据集类别多、类间差异小都是鸟区别在细微颜色和花纹是零样本学习论文里最常用的基准之一你拿它跑出来的结果可以跟论文直接对比。3.2 特征提取与数据加载这一节直接写代码先处理视觉特征提取。我这里先把整个流程拆解清楚方便你照着搭。import torch import torch.nn as nn from torchvision import models, transforms from PIL import Image def build_visual_encoder(): 构建冻结的视觉特征提取器 model models.resnet101(pretrainedTrue) # 去掉最后的全连接层只保留特征提取部分 model nn.Sequential(*list(model.children())[:-1]) model.eval() # 冻结参数 for param in model.parameters(): param.requires_grad False return model def extract_visual_feature(image_path, model, device): 提取单张图片的视觉特征 transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) img Image.open(image_path).convert(RGB) img_tensor transform(img).unsqueeze(0).to(device) with torch.no_grad(): feature model(img_tensor) # shape: [1, 2048, 1, 1] - [1, 2048] return feature.view(feature.size(0), -1)有两点需要提醒你。ResNet101在这里会输出2048维的向量这个维度是冻结的后续MLP的输入维度必须跟它对齐。另外我特意把归一化的均值和标准差设成了ImageNet的标准值因为预训练权重是在这个分布下学出来的如果你换成自己随便设的值特征质量会明显下降。CLIP文本编码器生成类别嵌入的代码也顺手贴一下from transformers import CLIPProcessor, CLIPModel clip_model CLIPModel.from_pretrained(openai/clip-vit-base-patch32) clip_processor CLIPProcessor.from_pretrained(openai/clip-vit-base-patch32) def get_class_embedding(class_names, descriptionsNone): 根据类别名或描述文本生成语义向量 if descriptions is None: texts [fa photo of a {name} for name in class_names] else: texts descriptions inputs clip_processor(texttexts, return_tensorspt, paddingTrue) with torch.no_grad(): embeddings clip_model.get_text_features(**inputs) # 做L2归一化方便后续算余弦相似度 embeddings embeddings / embeddings.norm(dim-1, keepdimTrue) return embeddings3.3 对齐模型与损失函数视觉到语义空间的投影对齐模型是整个代码里最核心的部分。它的输入是2048维的视觉向量输出是512维的语义向量和CLIP文本嵌入对齐。class AlignmentNetwork(nn.Module): 视觉特征到语义特征的对齐网络 def __init__(self, visual_dim2048, semantic_dim512, hidden_dim1024): super().__init__() self.net nn.Sequential( nn.Linear(visual_dim, hidden_dim), nn.BatchNorm1d(hidden_dim), nn.ReLU(inplaceTrue), nn.Dropout(0.5), nn.Linear(hidden_dim, semantic_dim) ) def forward(self, x): return self.net(x)两层结构加Dropout就足够了不要堆太多层。零样本学习的训练数据本身就有限模型太深非常容易过拟合。我自己试过四层的版本训练集上的loss降得跟过山车一样但测试集的表现反而不如这个简单的两层版本。损失函数这里直接使用余弦相似度。它天然适合度量两个向量的方向一致性而且对向量的绝对值不敏感在embedding对齐任务里比欧氏距离的表现稳定得多。def compute_loss(visual_feat, semantic_embed, temperature0.07): visual_feat: [batch_size, semantic_dim]经过对齐网络映射后的视觉特征 semantic_embed: [batch_size, semantic_dim]当前batch中样本所属类别的语义向量 visual_feat nn.functional.normalize(visual_feat, dim-1) semantic_embed nn.functional.normalize(semantic_embed, dim-1) # 计算余弦相似度矩阵 logits visual_feat semantic_embed.t() / temperature # 正样本是batch中对角线位置 labels torch.arange(logits.size(0)).to(logits.device) loss nn.functional.cross_entropy(logits, labels) return loss有个小细节值得留意temperature参数的取值会影响训练难度我习惯设成0.07这是多模态对比学习里一个比较通用的经验值。取值太小会让logits之间的差距过大梯度容易不稳取值太大会让所有相似度都挤在一起模型学不出区分度。3.4 完整训练流程逐步跑通你的第一个ZSL模型接下来把训练和评估的完整流程串起来。我会按实际运行的顺序走一遍每一步都注明输入输出。第一步准备数据划分。CUB数据集里的200类鸟用前150类做训练可见类后50类做测试未见类。这个划分方式在论文里很常见属于标准协议。第二步实例化各组件并初始化训练器from torch.utils.data import DataLoader, Dataset import os import random class CUBSimpleDataset(Dataset): 简化的CUB数据集加载器 def __init__(self, image_paths, labels, class_names): self.image_paths image_paths self.labels labels self.class_names class_names def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image_path self.image_paths[idx] visual_feat extract_visual_feature(image_path, visual_encoder, device) label self.labels[idx] return visual_feat.squeeze(0), label # 构建文本描述 descriptions [fa photo of a {name} for name in unseen_class_names] # 生成全部类别的语义向量 all_class_embeddings get_class_embedding(all_class_names, descriptions)注意这里有个性能优化点由于视觉特征提取器是冻结的我们可以提前把所有训练图像的视觉特征都提取好存起来训练时直接读特征而不是重跑ResNet。这一步能把训练时间缩短一个数量级实际体验非常明显。第三步训练循环。常规的PyTorch流程batch_size我建议设64Adam优化器初始学习率1e-3训练30个epoch就足够收敛了。device torch.device(cuda if torch.cuda.is_available() else cpu) align_net AlignmentNetwork().to(device) optimizer torch.optim.Adam(align_net.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30) for epoch in range(30): align_net.train() total_loss 0.0 for batch_visual_feat, batch_labels in train_loader: batch_visual_feat batch_visual_feat.to(device) # 取当前batch样本对应的语义向量 batch_semantic all_class_embeddings[batch_labels].to(device) projected align_net(batch_visual_feat) loss compute_loss(projected, batch_semantic) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() scheduler.step() print(fEpoch {epoch1}/{30}, Loss: {total_loss/len(train_loader):.4f})第四步评估。分别测试传统ZSL和GZSL两种设置下的准确率def evaluate_zsl(align_net, all_class_embeddings, test_loader, seen_classesNone): 传统ZSL只在未见类中做分类 广义ZSL在可见类和未见类的并集中做分类 align_net.eval() correct 0 total 0 # 决定在哪个类别集合里搜索 if seen_classes is None: # 传统ZSL只用未见类的语义向量 search_embeddings all_class_embeddings[unseen_label_indices] valid_labels list(range(len(unseen_label_indices))) else: # 广义ZSL全部类别参与 search_embeddings all_class_embeddings valid_labels list(range(all_class_embeddings.size(0))) with torch.no_grad(): for batch_visual_feat, batch_labels in test_loader: batch_visual_feat batch_visual_feat.to(device) projected align_net(batch_visual_feat) projected nn.functional.normalize(projected, dim-1) search_embeddings nn.functional.normalize(search_embeddings, dim-1) similarity projected search_embeddings.t() _, preds similarity.topk(1, dim1) preds preds.squeeze(1) for i in range(len(batch_labels)): if preds[i].item() valid_labels.index(batch_labels[i].item()): correct 1 total 1 return correct / max(total, 1)这套工程跑下来在CUB数据集上传统ZSL的Top-1准确率一般能到55%到65%之间广义ZSL的调和平均在40%到50%之间。这个水平跟2020年前后的论文结果基本持平作为基线工程完全够用。如果你用CLIP的视觉编码器替换ResNet101分数还有不小的提升空间。4. 关键细节踩坑记录那些论文里不会写的教训4.1 数据的“泄漏”问题比想象中隐蔽零样本学习最忌讳的就是数据泄漏。但“泄漏”不光是说你误把测试类别的图像放进训练集——还有更隐蔽的情况。CUB数据集一共有11788张图片每张图片都归属于一个类别。如果你在做类别划分时把同一种鸟分成两个子类一部分用于训练一部分用于测试模型就相当于“变相见过”测试类别的样本。论文里管这个叫“unbiased”划分是特别容易踩的坑。还有一个更隐蔽的泄漏点出现在类嵌入的构造环节。如果用CLIP的文本编码器生成类别嵌入而CLIP的训练数据里恰好有CUB测试类别的图片和文本描述语义向量就已经“见过”这些类了。这个泄漏理论上会推高测试指标但很难完全规避——即使是论文作者也基本默认了这一层“软泄漏”因为完全剔除不现实。我的建议是做对比实验时所有方法都使用相同的类嵌入来源这样系统性偏差是公平的结论依然可信。4.2 Hubness现象与归一化的选择用余弦相似度做最近邻检索时有个统计学现象会让结果悄悄变差这个现象叫Hubness。简单来说在高维空间里会有某些“流行点”频繁地成为很多查询向量的最近邻——它们就像是空间里的社交达虫跟谁都“像”。放在零样本学习里这意味着某些语义向量会被反复预测成答案导致单类准确率极不均衡有几个类的准确率特别高另外几个类几乎从来预测不对。缓解Hubness的一个常用技巧是在相似度计算时做“跨域归一化”比如先对相似度矩阵做行归一化再取最大值或者使用排名聚合策略。另外一个跟归一化直接相关的坑是视觉投影向量和语义向量的尺度如果不一致会破坏余弦相似度的判别能力。我的习惯是在计算相似度前两边都做一次L2归一化这样逻辑上更干净也避免了不同特征尺度带来的干扰。别小看这行代码有时候它能带来两到三个百分点的提升。4.3 GZSL下类别偏置的应对策略广义零样本学习最大的挑战是模型天然偏置向可见类。原因在于训练时模型只见过可见类的对齐关系语义空间中可见类区域的“吸引力”天然强于未见类区域。所以测试时哪怕一张图确实是未见类的模型也经常把它判到某个可见类头上。有个简单的缓解办法是“校准”——对可见类的相似度分数打一个折扣比如乘一个0.7的系数人为降低可见类的吸引力。这个系数在验证集上搜索最优值即可。从代码层面你可以在相似度矩阵上做这样一个操作# 对可见类的logits做折扣 similarity[:, seen_label_indices] * 0.7 cf max(0, similarity.max() * 0.1) similarity similarity - cf * (similarity cf).float()第二行代码来自一篇经典的GZSL论文思路是把过高的相似度分数往下压一压缓解“分数虚高”的问题。我自己实验下来这两行代码能让GZSL的调和平均从40%提升到46%到48%左右。不是大改动效果却很明显。5. 代码调试中遇到的四个典型问题5.1 模型不收敛或Loss异常振荡如果你发现loss死活降不下去先检查一下是不是温度参数设置得不合理以及特征是否做了归一化。很多第一次上手零样本学习的同学直接把原始图像输入网络就开始训练loss自然乱跳。另外还要确认优化器的学习率是否过大我实测Adam加1e-3对两层MLP是合适的超过1e-2基本就会出问题。比较推荐的排查顺序是先看看视觉特征是不是正常分布——打印一下特征的均值和方差正常的预训练特征均值应该接近0、方差在1的数量级附近再检查语义向量是否做了归一化最后检查温度参数。5.2 评估时标签索引对不上传统ZSL评估时要特别注意一个坑你用的模型输出的是“语义向量列表的索引”而不是真实的类名标签。如果训练集包含150类、测试集只有50类而你的测试集标签是沿用原始CUB的200类编号直接拿原始标签跟预测索引比较结果一定是错的。我一般这样处理建立一个映射字典把原始类别标签映射到“当前测试子集的相对索引”评估时两个索引体系保持一致再算准确率。这个错误很隐蔽通常表现是准确率极低接近随机猜测排查时先核对这一点。5.3 内存不足或训练速度过慢如果不做特征缓存ResNet101跑一遍整张CUB数据集是相当慢的。建议像我前面说的先把特征全部提取好存成npy文件训练时直接加载速度能提升十倍以上。如果数据量更大比如ImageNet级别的建议用DataLoader的persistent_workers配合pin_memory也能提速不少。5.4 预训练模型加载失败torchvision加载ResNet101需要联网下载权重国内网络环境时容易超时。解决方案是手动下载权重文件放到本地再用model.load_state_dict(torch.load(resnet101.pth))加载。CLIP模型也是类似的思路提前用HuggingFace的snapshot_download把模型拉到本地然后把from_pretrained的路径指到本地目录即可。6. 零样本学习的扩展方向下一步可以怎么玩基础的零样本学习工程跑通之后沿着论文的方向可以往几个方向扩展。我搭完基线后第一个尝试的方向是引入“生成式”思路——既然要把视觉特征映射到语义空间为什么不直接反过来用条件生成模型把语义向量“翻译”成视觉特征呢这样就能通过生成样本把零样本问题变成传统的监督学习问题。这个思路对应f-CLSWGAN那类工作效果在某些数据集上比直接映射要好。第二个值得关注的方向是transformer架构带来的变化。CLIP出现后视觉和语义空间天然对齐得更好很多论文开始直接用CLIP做零样本分类完全不训练任何额外的对齐模块——这就是现在所谓的“开放词汇”识别。在我的实操对比里直接冻结CLIP做zero-shot分类CUB上的Top-1准确率能做到70%以上比ResNet加MLP的基线高一截。第三个方向是多模态数据的融合。现在很多工作开始把音频、文本、知识图谱信息都揉进类嵌入里。比如做细粒度鸟类识别时除了用“a photo of a xxx”这种描述再把鸟的叫声文字化描述、栖息地信息也编码进语义向量模型能抓住的判别信息就更多了。这个方向最终的目标是解决“开集识别”问题——模型不仅要知道这张图属于已定义的哪个类还要能判断它根本不属于任何已知类。这是零样本学习和开放世界视觉结合的前沿方向也是我个人觉得未来两三年会持续产出高质量论文的方向。最后分享一个实操体会零样本学习这个领域代码门槛真的不高——你只需要把数据集加载、特征提取、对齐网络、评估链路四块搞清楚一篇论文的实验部分基本就能复现。真正花时间的地方在于理解为什么这么设计、在哪个环节做改进能带来有效提升。建议你先把我这篇里的基线工程跑通再拿一两篇经典论文比如关于GZSL校准、transformer时代的多模态对齐的代码做对比实验很快就能建立起自己的直觉。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门