Few-shot学习实战:用PyTorch在5个样本内训练高精度分类模型
从零到一用PyTorch在极少量样本上构建高精度分类器的实战指南想象一下你正面对一个极具挑战性的项目客户提供了一批珍贵的医疗影像切片但每个病变类别只有寥寥几张标注图片或者你需要为一条新的工业生产线开发视觉质检系统但产线上能提供的缺陷样本屈指可数。在传统机器学习范式下数据科学家们可能会感到束手无策——没有海量数据模型似乎无从学起。然而这正是Few-shot Learning少样本学习大显身手的舞台。它不再依赖“数据暴力”而是转向“智慧迁移”让模型学会“举一反三”。今天我们就抛开理论空谈直接深入代码层面手把手带你用PyTorch打造一个在5个样本以内就能训练出可用分类模型的实战方案。无论你是面临数据稀缺困境的算法工程师还是对前沿机器学习应用充满好奇的开发者这篇文章都将为你提供一套清晰、可落地的工具箱。1. 理解核心Few-shot学习为何能“无中生有”在开始敲代码之前我们必须先打破一个思维定式模型性能一定与训练数据量成正比。Few-shot学习的核心思想在于它假设模型在接触新任务时并非从一张白纸开始。相反它已经通过预训练在大规模通用数据集如ImageNet上学会了如何“看”世界——识别边缘、纹理、形状等基础视觉特征。Few-shot学习要做的就是巧妙地调整和利用这些先验知识使其快速适应只有极少样本的新类别。这背后的哲学更像是人类的“类比学习”。当你第一次见到某种稀有鸟类时即使只看过一张照片你也能凭借对“鸟类”的通用认知有喙、有羽毛、会飞在未来认出它。Few-shot模型也是如此。其技术路径主要围绕三个关键点展开度量学习 核心是让模型学会一个“好的”特征空间。在这个空间里同一类别的样本彼此靠近不同类别的样本相互远离。这样对于一个新样本只需计算它与各类支持集那少量的训练样本在特征空间中的距离就能进行分类。孪生网络、原型网络是其中的经典代表。元学习 又称“学会学习”。其训练目标不是直接在某个具体任务上获得高精度而是让模型获得一种“快速适应能力”。训练过程模拟了Few-shot场景在大量不同的“小任务”上进行训练每个任务都有自己的微型训练集和测试集。通过这种方式模型内化了一套参数更新规则当遇到全新的、样本极少的小任务时能通过极少的几步梯度更新就达到良好性能。MAML是这一范式的里程碑。基于预训练模型的快速微调 这是目前工程上最常用、也往往最有效的策略。我们直接使用在大型数据集上预训练好的模型如ResNet、ViT作为特征提取器然后仅用一个简单的分类器如线性层去适配新类别。关键在于我们要冻结特征提取器的大部分层只让最后几层或分类器进行微调同时配合强力的正则化手段防止在极少样本上过拟合。为了更直观地理解这三种主流Few-shot学习范式的区别与适用场景我们可以参考下面的对比表格范式核心思想优点缺点典型应用场景度量学习学习一个通用的相似性度量空间直观推理速度快无需针对新任务更新模型参数对特征提取器的质量依赖高跨域任务可能表现不佳人脸验证、图像检索、类别稳定的细粒度分类元学习让模型学会“如何快速学习”理论优美适应新任务的速度可能极快训练复杂不稳定需要大量元训练任务计算成本高研究前沿、需要快速适应一系列相似但不同任务的场景快速微调利用强大预训练模型进行极小范围调整实现简单稳定性高能直接利用SOTA预训练模型性能受预训练模型与新任务领域差异影响大工业界首选医疗影像、工业质检等数据稀缺但与预训练数据有一定相关性的领域提示对于大多数实际应用尤其是从零开始的实践我强烈建议你从**“基于预训练模型的快速微调”** 入手。它结合了效果、稳定性和实现简易度是我们后续实战部分的核心。2. 实战环境搭建与数据准备理论清晰后我们立刻进入实战环节。首先确保你的环境已经就绪。我们将使用PyTorch作为核心框架。# 推荐使用conda创建虚拟环境 conda create -n fewshot python3.9 conda activate fewshot # 安装核心依赖 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 # 根据你的CUDA版本调整 pip install pytorch-lightning # 可选但能极大简化训练循环代码 pip install scikit-learn pandas matplotlib tqdm接下来是数据。Few-shot学习的数据组织方式与传统监督学习不同。我们通常采用“N-way K-shot”的设置来构建任务。N-way: 每个任务中包含N个不同的类别。K-shot: 每个类别提供K个带标签的样本作为支持集用于模型学习或适应。Query集: 同一批类别下的另外一些样本用于评估模型在该任务上的性能。假设我们有一个包含100个类别的数据集MiniDataset我们要进行5-way 1-shot的任务采样。下面是一个简单的PyTorchDataset类示例用于生成这样的任务import torch from torch.utils.data import Dataset, DataLoader import random from torchvision import transforms class FewShotTaskDataset(Dataset): 一个用于生成N-way K-shot任务的简单数据集包装器。 def __init__(self, base_dataset, n_way5, k_shot1, query_per_class5, transformNone): Args: base_dataset: 原始数据集应能按类别索引数据。 n_way: 任务类别数。 k_shot: 每个类别的支持集样本数。 query_per_class: 每个类别的查询集样本数。 transform: 数据增强变换。 self.base_dataset base_dataset self.n_way n_way self.k_shot k_shot self.q_per_class query_per_class self.transform transform # 假设base_dataset可以通过某种方式获取类别列表和样本索引 # 这里仅为示例你需要根据实际数据集结构调整 self.class_to_indices self._build_class_index() def _build_class_index(self): # 实现逻辑遍历base_dataset将同一类别的样本索引归到一起 # 返回格式{class_id: [index1, index2, ...]} # 此处为伪代码 class_idx {} for idx, (_, label) in enumerate(self.base_dataset): class_idx.setdefault(label, []).append(idx) return class_idx def __len__(self): # 定义可以生成多少个不同的任务例如基于类别组合数 return 1000 # 一个较大的数表示可以无限生成 def __getitem__(self, idx): # 随机采样一个任务 selected_classes random.sample(list(self.class_to_indices.keys()), self.n_way) support_set [] query_set [] for class_id in selected_classes: all_indices self.class_to_indices[class_id] # 随机选择支持集和查询集样本确保不重叠 selected random.sample(all_indices, self.k_shot self.q_per_class) support_indices selected[:self.k_shot] query_indices selected[self.k_shot:] for s_idx in support_indices: img, _ self.base_dataset[s_idx] if self.transform: img self.transform(img) support_set.append((img, class_id)) # 注意这里class_id需要被重新映射为0到N-1 for q_idx in query_indices: img, _ self.base_dataset[q_idx] if self.transform: img self.transform(img) query_set.append((img, class_id)) # 将支持集和查询集的图像和标签分别堆叠并重新映射标签为任务内标签[0, N-1] # ... (具体的堆叠和标签映射代码) # 返回格式 (support_images, support_labels), (query_images, query_labels) return (s_imgs, s_labels), (q_imgs, q_labels) # 示例使用CIFAR-100作为基础数据集 from torchvision.datasets import CIFAR100 base_data CIFAR100(root./data, trainTrue, downloadTrue) task_dataset FewShotTaskDataset(base_data, n_way5, k_shot1, query_per_class5, transformtransforms.ToTensor()) task_loader DataLoader(task_dataset, batch_size4, num_workers2) # batch_size 这里指每次加载几个任务注意上述FewShotTaskDataset是一个高度简化的示例框架。在实际应用中你需要根据具体数据集结构如Omniglot、miniImageNet或你的自定义数据来完善_build_class_index和__getitem__方法。元学习库如learn2learn或torchmeta提供了更成熟的任务采样器。3. 核心策略一基于预训练模型的快速微调这是我们的主力方案。我们选择一个在ImageNet上预训练好的ResNet-18将其最后的全连接层替换掉然后进行微调。import torch.nn as nn import torch.optim as optim from torchvision import models import torch.nn.functional as F class FewShotClassifier(nn.Module): def __init__(self, backbone_nameresnet18, n_way5, feature_dim512): super().__init__() # 加载预训练骨干网络 if backbone_name resnet18: self.backbone models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # 移除原始的分类头 self.feature_extractor nn.Sequential(*list(self.backbone.children())[:-1]) # 假设特征维度是512 (ResNet-18的最后一层特征图池化后) else: raise ValueError(fUnsupported backbone: {backbone_name}) # 冻结骨干网络的大部分层只微调最后1-2个block for param in self.feature_extractor.parameters(): param.requires_grad False # 解冻最后两个BasicBlock (resnet18的layer4) for param in self.backbone.layer4.parameters(): param.requires_grad True # 为新的N-way分类任务添加一个简单的分类头 self.classifier nn.Linear(feature_dim, n_way) def forward(self, x): # 提取特征 features self.feature_extractor(x) features features.view(features.size(0), -1) # 展平 # 分类 logits self.classifier(features) return logits # 训练循环的关键部分 def fast_adapt_train(model, support_imgs, support_labels, query_imgs, query_labels, optimizer, criterion, inner_steps10): 在一个任务上进行快速微调内循环和评估。 model.train() # 内循环在支持集上微调几步 for step in range(inner_steps): optimizer.zero_grad() logits model(support_imgs) loss criterion(logits, support_labels) loss.backward() optimizer.step() # 内循环后在查询集上评估 model.eval() with torch.no_grad(): query_logits model(query_imgs) query_loss criterion(query_logits, query_labels) query_pred query_logits.argmax(dim1) query_acc (query_pred query_labels).float().mean() return query_loss, query_acc # 主训练循环外循环 def main_training_loop(): n_way 5 k_shot 1 model FewShotClassifier(n_wayn_way).cuda() # 注意优化器只对需要梯度的参数解冻的层和分类器进行优化 optimizer optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr1e-3) criterion nn.CrossEntropyLoss() for epoch in range(num_epochs): for (s_imgs, s_labels), (q_imgs, q_labels) in task_loader: s_imgs, s_labels s_imgs.cuda(), s_labels.cuda() q_imgs, q_labels q_imgs.cuda(), q_labels.cuda() # 每个任务batch中的一个任务独立进行快速微调 # 这里简化处理实际中可能需要为每个任务临时克隆模型或使用二阶优化如MAML # 对于简单微调我们也可以在所有任务的支持集混合数据上做一步微调然后在查询集评估 # 以下是简化版将batch中所有任务的支持集合并进行一次梯度更新 optimizer.zero_grad() support_logits model(s_imgs) loss criterion(support_logits, s_labels) loss.backward() optimizer.step() # 评估当前模型在查询集上的表现可选用于监控 with torch.no_grad(): query_logits model(q_imgs) query_acc (query_logits.argmax(dim1) q_labels).float().mean() print(fEpoch {epoch}, Loss: {loss.item():.4f}, Query Acc: {query_acc.item():.4f})这个方案的核心在于谨慎的解冻策略和强有力的正则化。我们只微调网络最深层的、最具任务特异性的部分同时保持底层通用特征的稳定。在实际操作中你还需要加入以下关键技巧来对抗过拟合数据增强的极致利用 对于仅有的几个样本我们必须通过增强“创造”出多样性。除了标准的随机裁剪、翻转还可以尝试RandAugment或AutoAugment 自动搜索的增强策略组合。MixUp或CutMix 在图像或特征层面混合样本创造虚拟训练数据。针对领域的增强例如在医疗影像中模拟不同的造影剂浓度在工业质检中模拟不同的光照和角度。标签平滑 将硬标签0或1转换为软标签如0.9和0.1防止模型对少数样本过于自信。Dropout与权重衰减 即使在微调阶段也保持一定比例的Dropout并使用较大的权重衰减值。4. 核心策略二度量学习与原型网络实现如果你需要一种无需为每个新任务更新模型参数的方法度量学习是更好的选择。原型网络是其中最直观的一种。它的思想非常简单为每个类别计算一个“原型”该类所有支持样本特征向量的均值然后通过计算查询样本特征与各个原型的距离如欧氏距离来进行分类。import torch import torch.nn as nn class PrototypicalNetwork(nn.Module): def __init__(self, backbone): super().__init__() self.backbone backbone # 一个预训练的特征提取器 def forward(self, support_x, support_y, query_x): Args: support_x: [num_support, C, H, W] support_y: [num_support] query_x: [num_query, C, H, W] Returns: query_logits: [num_query, n_way] n_way len(torch.unique(support_y)) # 提取特征 support_features self.backbone(support_x) # [num_support, feature_dim] query_features self.backbone(query_x) # [num_query, feature_dim] # 计算每个类别的原型均值 prototypes [] for class_id in range(n_way): # 选出当前类别的所有支持样本特征 mask (support_y class_id) class_features support_features[mask] prototype class_features.mean(dim0) # [feature_dim] prototypes.append(prototype) prototypes torch.stack(prototypes, dim0) # [n_way, feature_dim] # 计算查询样本特征与所有原型的欧氏距离的平方 # 扩展维度以便广播计算 # query_features: [num_query, feature_dim] - [num_query, 1, feature_dim] # prototypes: [n_way, feature_dim] - [1, n_way, feature_dim] dists torch.cdist(query_features.unsqueeze(1), prototypes.unsqueeze(0), p2).squeeze(1) # [num_query, n_way] # 将距离转换为概率负距离距离越小概率越大 logits -dists return logits # 训练原型网络 def train_protonet(model, train_loader, optimizer, n_way, k_shot): model.train() for batch_idx, ((s_imgs, s_labels), (q_imgs, q_labels)) in enumerate(train_loader): s_imgs, s_labels, q_imgs, q_labels s_imgs.cuda(), s_labels.cuda(), q_imgs.cuda(), q_labels.cuda() optimizer.zero_grad() logits model(s_imgs, s_labels, q_imgs) # 前向传播计算logits loss F.cross_entropy(logits, q_labels) loss.backward() optimizer.step() if batch_idx % 50 0: acc (logits.argmax(dim1) q_labels).float().mean() print(fBatch {batch_idx}, Loss: {loss.item():.4f}, Acc: {acc.item():.4f})原型网络的训练目标就是让同类样本的特征在嵌入空间中尽可能聚集。一旦训练完成对于全新的N-way K-shot任务你只需要将支持集样本输入网络得到特征计算原型然后就可以直接对查询集进行分类无需任何梯度更新。这种“一次前向传播搞定”的特性使其在需要快速推理的场景下非常有吸引力。5. 进阶技巧与避坑指南在实际项目中仅仅套用上述模型往往不够。下面这些技巧和注意事项是我在多个少样本项目实践中总结出来的“血泪经验”。数据层面是决胜关键。当样本量小于5时每一个样本都价值连城。人工数据清洗与标注复核 极少的样本里如果混入一个错误标注或低质量样本对模型的伤害是毁灭性的。务必投入精力确保支持集样本是清晰、典型、无歧义的。智能数据增强流水线 不要只用简单的RandomHorizontalFlip。构建一个包含几何变换旋转、缩放、透视、颜色抖动亮度、对比度、饱和度、添加噪声高斯噪声、椒盐噪声以及CutMix、RandAugment的增强组合。可以使用albumentations库来构建强大的增强管道。利用无标注数据如果存在 如果领域内存在大量无标签数据可以结合自监督学习如SimCLR、MoCo先为特征提取器进行域内预训练这能显著提升特征质量。模型与训练技巧选择合适的预训练模型 在医疗影像领域在RadImageNet大型医学影像数据集上预训练的模型通常比在ImageNet上预训练的模型迁移效果更好。工业质检中在ImageNet基础上用大量正常品图像进行自监督预训练也很有帮助。学习率与优化器 使用较小的学习率如1e-4到1e-5进行微调。考虑使用AdamW优化器并配合CosineAnnealingLR学习率调度器让学习率平滑下降。早停法 由于数据量极小模型很容易在几个epoch内就过拟合。密切监控在一个保留的验证任务集上的性能一旦性能开始下降立即停止训练。集成学习 训练多个不同的模型例如使用不同的预训练骨干、不同的数据增强种子、不同的模型初始化然后对它们的预测进行平均或投票。这在极低数据量下是提升稳定性和性能的有效手段。评估与上线可靠的评估协议 不要只在一个5-way 1-shot任务上测试。应该报告在数百个随机采样的任务上的平均准确率和95%置信区间。这能更真实地反映模型的泛化能力。设计反馈循环 上线后系统很可能会对预测结果不确定置信度低。将这些“困难样本”记录下来交由专家进行标注并定期将其加入支持集重新微调模型。这是一个让系统在应用中持续进化的关键机制。最后我想分享一个在工业瑕疵检测中的真实体会我们曾为一个客户构建一个检测新型划痕的模型初始样本只有3张正面光照下的图片。单纯使用原型网络效果不佳。后来我们做了三件事第一用这3张图生成了超过200张包含不同光照角度、模拟不同深浅的增强图像第二选择了一个在金属表面缺陷公开数据集上微调过的ResNet作为骨干第三采用了上述的快速微调策略并加入了较强的Dropout。最终模型在产线上的检出率达到了95%以上。这个案例告诉我在Few-shot学习中对有限数据的“精耕细作”和“对症下药”的模型调整比追求复杂的算法魔术更为重要。当你手里只有几颗种子时思考如何为每一颗创造最好的生长环境远比幻想一片森林来得实际。