ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

原型网络实战:少样本学习中的PyTorch实现与调参避坑

原型网络实战:少样本学习中的PyTorch实现与调参避坑 简介这份PyTorch实现以原型网络为核心面向从事少样本学习研究与应用的机器学习开发者原型网络采用类原型距离度量完成分类无需额外微调即可处理小样本分类任务适合初学者快速入门该方向并搭建实验基线。压缩包内共有十二个文件主要包含七个Python源码脚本分别实现模型结构、原型损失函数、训练流程、批次采样、Omniglot数据集加载、参数解析等模块另附两张图片用于说明网络结构或训练结果以及自述文件与许可证文件整体体量仅约135KB结构精简、模块划分清晰。目前已有九百二十人学习下载。读者可借由这套代码获得完整可运行的少样本分类基线既能直接启动训练流程并观察分类效果也能在源码基础上调整数据和网络模块便于复现论文实验或进行后续算法改进。1. 少样本学习的落地难题为什么原型网络值得上手金融风控里的新欺诈模式、医学影像里的罕见病灶、工业质检里的新型缺陷这类场景的共同痛点是正常样本堆积如山目标样本只有几张。少样本学习Few-shot Learning要解决的就是“每个类别只给几张标注模型还得能用起来”的问题。原型网络Prototypical Networks是其中思路最直接、复现成本最低的一支——它把每个类别的支撑样本嵌入到一个向量空间取均值作为“原型”查询样本离哪个原型近就判给哪类。整个模型用PyTorch写下来不过两三百行不需要分布式训练也不需要二次求导。这篇文章按“原理 → 最小实现 → 关键参数 → 踩坑 → 验证”的顺序给你一条能在一周内跑通并复现少样本基准的路。2. 原型网络在少样本学习里的定位支撑集、原型向量与最近邻分类2.1 少样本任务的数学表述N-way K-shot 与查询集在进入代码之前先把任务定义锁死。少样本学习里一个最基本的概念叫 episode一个 episode 就是一次独立的小型分类任务。它由两部分组成支撑集support set和查询集query set。支撑集里包含 N 个类别每个类别 K 个已标注样本这就是常说的 N-way K-shot查询集里则是从同样 N 个类别里取出的另一批未标注样本用来评估模型在这个小任务上的分类表现。举个具体例子5-way 5-shot 意味着每个训练回合随机挑 5 个类每类给 5 张支撑样本再每类拿 5 张当查询样本。模型只能看支撑集来推断查询集的标签。这里面有个反直觉的地方K 越小任务越难1-shot 时每类只有一张支撑样本连类内方差都估计不出来模型只能硬记那张图的特征。少样本学习和传统分类的另一个本质区别是类别集合在训练和测试时是错开的。训练时见过猫和狗测试时却要区分杯子和桌子。所以模型不能“记住类别”必须学会“按支撑集现场对比”。这也解释了为什么原型网络的核心不是分类头而是嵌入函数——它只负责把图像映射成向量分类完全靠支撑集算出来的原型和查询向量之间的距离。2.2 原型网络的核心假设类内均值就是嵌入空间里的类代表原型网络的基本逻辑非常朴素。设嵌入函数为 fθ它把输入图像映射成 d 维向量。对于第 k 类它的原型 c_k 就是该类所有支撑样本嵌入的均值c_k (1 / K) * Σ fθ(x_i)查询样本 x 的嵌入是 z模型计算 z 到每个原型 c_k 的欧氏距离再用 softmax 转成类别概率p(yk|x) exp(-||z - c_k||²) / Σ_j exp(-||z - c_j||²)注意欧氏距离前面有个负号距离越小概率越大这和普通 softmax 里“得分越高概率越大”正好相反。实现时通常在距离矩阵前加一个负号再传给交叉熵损失。这个设计成立的前提是同一个类别在嵌入空间里近似服从单峰分布峰的中心就是这个类的均值。换句话说模型训练的目标就是把每类样本在嵌入空间里聚成一团同时把不同类别的团尽量分开。这个假设比孪生网络的“同类距离近、异类距离远”更具体也比 MAML 那种“学会快速初始化”更轻量。损失函数直接对查询样本的预测结果做交叉熵梯度回传。可以这样理解支撑集负责构造原型查询集负责纠偏嵌入函数。原型网络没有显式计算类间边界所有边界都由欧氏距离在嵌入空间里隐式决定。2.3 与孪生网络、MAML 的选型对比什么时候选原型网络少样本学习方向常见方案有三个孪生网络Siamese Network、模型无关元学习MAML和原型网络。表格对比它们的核心差异维度孪生网络MAML原型网络训练思路判断两张图是否同类学习一个易于微调的初始化参数学习嵌入函数让类内均值成为有效原型支撑集使用方式两两配对只看相似/不相似每个任务的支撑集上做几步梯度更新直接对嵌入取均值二次求导不需要需要显存开销大不需要1-shot 表现依赖配对构造质量波动大表现强但训练极不稳定表现稳定是 1-shot 基准里的强基线代码复杂度较低较高涉及内循环外循环最低两三行就能完成核心逻辑典型适用场景人脸验证、签名比对测试时类别分布与训练差异大的场景图像分类、文本分类类别原型相对稳定选型建议很直接如果你只有一张消费级 GPU又需要快速跑出基准结果原型网络是第一选择。MAML 在难度更高的跨分布任务上上限更高但训练过程对学习率、内循环步数极其敏感翻车概率大。孪生网络适合二分类验证任务但把它推广到多分类少样本任务时需要额外设计配对策略和损失权重不如原型网络干净。原型网络也有明显的边界它对嵌入空间的分布假设过于理想。如果数据里的类内分布是多峰的——比如同一种产品有多个外观批次——均值会把所有峰值拉平原型落在峰之间精度必然下降。这个边界后面会在避坑章节里具体说。3. 在 PyTorch 里跑通原型网络的最小实现采样器、卷积编码器与训练循环3.1 环境准备PyTorch 环境搭建与 Python 版本对应原型网络依赖的库非常少torch、torchvision、numpy、tqdm。我的建议是用 Anaconda 建独立环境避免把系统自带的 Python 环境弄乱。配置 PyTorch 环境时最容易出问题的点就是 Python 版本与 PyTorch 版本对应关系。目前稳定分支 PyTorch 对 Python 3.8 到 3.12 都支持但 CUDA 版本与显卡驱动必须分开对待。conda create -n proto python3.9 -y conda activate proto pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118这段命令里Python 3.9 是兼容性和稳定性都最保险的选择torch 和 torchvision 必须同时安装否则后续调用 torchvision.datasets 或 transforms 时会碰到版本咬合错误。如果你的机器没有 NVIDIA GPU把 index-url 去掉直接 pip install torch 得到的 CPU 版本也能跑通本章代码只是训练速度慢到难以接受。Linux 系统安装 Python 后如果遇到 conda 命令找不到检查一下 shell 的 PATH 是否包含 anaconda3/bin。Windows 用户建议走 WSL 或直接装 CUDA 版 PyTorch混合环境容易在数据加载阶段报段错误。装完后用下面一段命令验证 CUDA 是否可用python -c import torch; print(torch.__version__, torch.cuda.is_available())看到 True 且输出 PyTorch 版本号环境就算通了。注意显卡驱动版本过低会报 CUDA driver version is insufficient这种情况要么升级驱动要么装更低版本的 cu111 对应包。3.2 构造 Episode 采样器让每个 batch 成为一次少样本任务原型网络的训练数据不是普通 DataLoader 按 batch 吐样本而是每个 batch 必须是一个完整的 episode。我习惯把采样器写成一个迭代器类输入全部训练样本的标签输出支撑集和查询集的索引。import torch import numpy as np from collections import defaultdict class EpisodeSampler: def __init__(self, labels, n_way, k_shot, k_query, episodes_per_epoch): self.n_way n_way self.k_shot k_shot self.k_query k_query self.episodes_per_epoch episodes_per_epoch self.label_to_indices defaultdict(list) for idx, label in enumerate(labels): self.label_to_indices[label].append(idx) self.labels list(self.label_to_indices.keys()) def __iter__(self): for _ in range(self.episodes_per_epoch): episode_classes np.random.choice(self.labels, self.n_way, replaceFalse) support_indices [] query_indices [] for cls in episode_classes: indices np.array(self.label_to_indices[cls]) np.random.shuffle(indices) support_indices.extend(indices[:self.k_shot].tolist()) query_indices.extend(indices[self.k_shot:self.k_shot self.k_query].tolist()) yield support_indices, query_indices这个采样器有几个参数直接决定任务难度。n_way 越大分类边界越复杂k_shot 越小原型越不稳定。我一般先用 5-way 5-shot 做初跑因为模型在这个配置下能较快收敛方便验证代码流水线有没有错。等到跑通后再改成 5-way 1-shot 压测模型在极端条件下的真实能力。还有一个隐藏的坑支撑集和查询集绝对不能有交集否则模型会在评估阶段“作弊”。上面的代码在类别内 shuffle 后先取前 k_shot 个给支撑集再取后面 k_query 个给查询集天然保证了无交集。如果你的数据集每个类别样本数少于 k_shot k_query采样器会直接报索引越界这就是第一个常见的翻车点后面避坑章节细说。3.3 定义卷积编码器原型网络的特征提取器嵌入函数是原型网络里唯一可学习的部分。这里给出一个常见的四层卷积编码器输入 Omniglot 的 28×28 灰度图输出 64 维向量。这个结构在多个少样本基准上都能达到可复现的基线水平。import torch.nn as nn import torch.nn.functional as F class ConvEncoder(nn.Module): def __init__(self, in_channels1, hidden_dim64, embedding_dim64): super().__init__() self.conv1 nn.Conv2d(in_channels, hidden_dim, 3, padding1) self.bn1 nn.BatchNorm2d(hidden_dim) self.conv2 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1) self.bn2 nn.BatchNorm2d(hidden_dim) self.conv3 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1) self.bn3 nn.BatchNorm2d(hidden_dim) self.conv4 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1) self.bn4 nn.BatchNorm2d(hidden_dim) def forward(self, x): x F.relu(self.bn1(self.conv1(x))) x F.max_pool2d(x, 2) x F.relu(self.bn2(self.conv2(x))) x F.max_pool2d(x, 2) x F.relu(self.bn3(self.conv3(x))) x F.max_pool2d(x, 2) x F.relu(self.bn4(self.conv4(x))) x F.adaptive_avg_pool2d(x, (1, 1)) return x.view(x.size(0), -1)这个编码器的设计有两点值得注意。第一每层卷积后面都接 BatchNormReLUBatchNorm 在小 batch 下会带来一定波动但它能让训练稳定很多。第二最后用 adaptive_avg_pool2d 把特征图压成 1×1任何分辨率的输入都能统一输出 64 维向量换数据集时不需要改模型结构。把中间层的 hidden_dim 从 64 改成 128模型容量提升但训练时间接近翻倍。在 28×28 的 Omniglot 上64 维嵌入已经够用但如果你换到 miniImageNet 那种 84×84 的彩色图建议把第一层 in_channels 改成 3embedding_dim 提到 128否则特征容量会成为精度瓶颈。3.4 训练与验证循环支撑集前向、原型计算与交叉熵回传核心训练逻辑只需要四个步骤支撑集前向得到嵌入按类取均值得到原型查询集前向得到嵌入计算距离并求交叉熵损失。下面的代码把原型计算和欧氏距离封装成独立函数方便训练和验证复用。def compute_prototypes(support_features, support_labels, n_way): prototypes [] for i in range(n_way): class_mask (support_labels i) class_features support_features[class_mask] prototype class_features.mean(dim0) prototypes.append(prototype) return torch.stack(prototypes) def euclidean_distance(query_features, prototypes): n query_features.size(0) m prototypes.size(0) query_squared query_features.pow(2).sum(dim1, keepdimTrue).expand(n, m) proto_squared prototypes.pow(2).sum(dim1).unsqueeze(0).expand(n, m) cross_term torch.matmul(query_features, prototypes.t()) dist query_squared proto_squared - 2 * cross_term return dist.sqrt() def compute_loss(query_features, query_labels, prototypes): dist euclidean_distance(query_features, prototypes) logits -dist return F.cross_entropy(logits, query_labels)这里的关键是用展开公式计算欧氏距离避免显式写双层循环。query_squared 是每个查询样本的向量模长proto_squared 是每个原型的模长cross_term 是查询向量与原型向量的点积矩阵最后组合出完整的距离矩阵。显存占用与 n_way × k_query 成正比在小数据集上很宽松。训练循环里要格外注意支撑集和查询集必须进入同一个模型前向但支撑集的梯度只通过“原型”间接影响模型参数。PyTorch 的自动求导会自动处理这条链不需要手动冻结参数。一个训练回合的代码大致是for support_idx, query_idx in train_sampler: support_data batch_data[support_idx].to(device) support_labels batch_labels[support_idx].to(device) query_data batch_data[query_idx].to(device) query_labels batch_labels[query_idx].to(device) support_feat encoder(support_data) query_feat encoder(query_data) prototypes compute_prototypes(support_feat, support_labels, n_way) loss compute_loss(query_feat, query_labels, prototypes) optimizer.zero_grad() loss.backward() optimizer.step()训练时的 batch_data 是从完整数据集里一次性取出的目的是让采样器能按索引自由组合支撑集和查询集。数据量达到几万张时把全部数据塞进显存不现实可以改成每次先随机抽一个子集再喂给采样器。验证阶段记得用 torch.no_grad() 包裹否则验证过程会累积计算图导致显存溢出。4. 影响精度的三个关键参数N-way、距离度量与特征归一化4.1 N-way K-shot 怎么配5-way 1-shot 到 5-way 5-shot 的差别N-way 和 K-shot 不只是实验配置也是精度的直接调节阀。K-shot 每多一张支撑样本原型估计的方差就小一截精度提升非常明显。常见经验数字在 Omniglot 上用上面的四层卷积编码器5-way 1-shot 的验证精度大约 96% 上下5-way 5-shot 能推到 99% 左右而 20-way 1-shot 会掉到 90% 以下。这很好理解——类别越多嵌入空间里需要区分的原型也越多彼此“打架”的概率变大。配置难度原型稳定性常见基线参考5-way 1-shot较高差仅依赖单张96%±0.55-way 5-shot中等好99%±0.220-way 1-shot高极差90%±0.820-way 5-shot偏难一般97%±0.4这个表说的是同规模同数据下的相对差异具体数字会随编码器容量和输入分辨率波动。训练时把 n_way 加大比如从 5 提到 10模型被迫在同一个任务里区分更多决策边界得到的嵌入通常泛化能力更强。但代价是每个 episode 的支撑样本数量线性上升训练一个 epoch 的时间也跟着变长。我的习惯是训练阶段用 10-way 5-shot评估阶段用标准的 5-way 1-shot。这样训练难度高于测试难度模型在标准指标上通常能占到便宜但又不会像 maml 那样对任务难度极其敏感。4.2 欧氏距离还是余弦相似度高维嵌入下的选择原始论文用的是欧氏距离但很多人在复现时发现把编码器输出做 L2 归一化后再算欧氏距离等价于余弦相似度的某种带温度形式的变体。关键是欧氏距离在高维空间里会退化成模长主导如果两个嵌入的模长差距很大距离计算几乎只反映模长差异方向信息被淹没。我把两种方案都试过在 Omniglot 上欧氏距离和余弦相似度差异不大但换到图像噪声更强的数据集时L2 归一化加欧氏距离往往更稳。因为归一化强制把所有样本放在同一个超球面上网络只能靠“方向”来区分类别不再投机取巧地靠拉长模长来刷损失。实现方式非常轻量def l2_normalize(features, dim1): return F.normalize(features, p2, dimdim)前向之后、计算原型之前对支撑集和查询集特征都做一次归一化。注意顺序——必须先归一化再取均值不能取完均值再归一化。因为原型是多个向量的平均如果先平均再归一化等价于对方向平均做了一个球面投影如果先归一化再平均则每个支撑样本在原型计算中的权重都被拉平了。这个细微差别在 1-shot 场景下影响不大但在 5-shot 场景下会让原型偏移几个百分点。4.3 温度参数与 L2 归一化把特征压进一个球面归一化之后所有距离都落在 [0, 2] 区间softmax 输入的分辨率大幅度压缩。这会导致一个副作用logits 的绝对值太小softmax 输出过于平滑训练收敛变慢。解决办法是引入温度参数 T把距离除以 T 再求负号def compute_loss_temperature(query_features, query_labels, prototypes, temperature1.0): dist euclidean_distance(query_features, prototypes) logits -dist / temperature return F.cross_entropy(logits, query_labels)temperature 小于 1 会把 logits 放大相当于让 softmax 分布更尖锐模型被迫对决策边界更自信temperature 大于 1 则让分布更平滑。在 L2 归一化后的设置里我一般把 T 设在 0.5 到 0.8 之间。T 太小时训练开始阶段梯度变化很剧烈容易在早期就陷入错误的局部最优T 太大又会导致梯度太小训练半天精度原地不动。这个参数可以和 learning rate 联动调节如果发现 loss 曲线锯齿状明显先试着调大一点温度再看。温度参数不是原型网络的标配但从实践角度看它是投入产出比最高的一个调参位。很多“照着论文代码跑不出来”的抱怨最后都发现是缺少了这层缩放。另外还要注意查询集和支撑集共享同一个特征提取器推理阶段的温度必须与训练阶段一致否则概率分布会被扭曲影响精度评估的公平性。5. 原型网络避坑指南五个常见问题从现象到根因5.1 现象训练 loss 下降但验证精度不涨训练集 loss 一路走低每个 epoch 都能看到明显的下降但跑到验证集上一看精度在 60% 附近来回震就是上不去。很多人会下意识去调模型结构或学习率其实问题通常出在数据流水线上。验证的时候支撑集和查询集采样方式与训练不一致例如训练时支撑集每类 5 张验证时支撑集每类 1 张原型质量的差距瞬间暴露。原因还有一类验证集的类别数比训练集少太多。训练时用 10-way验证时用 5-way分类难度降了但嵌入函数是朝着 10-way 训练的它花了很多容量去区分那些在验证阶段根本不出现的边角类别。解决方式是预先定义好实验协议SR 训练与验证的 n_way、k_shot 必须一致最多允许训练 n_way 大于测试 n_way 这一种不对称设置并且要明确记录在实验笔记里。5.2 现象1-shot 时性能断崖式下跌5-shot 验证精度 98%切到 1-shot 直接掉到 70%这种断崖很容易让新手怀疑代码写错了。排除代码问题后最可能的原因是支撑样本的嵌入输出没有做归一化导致单张样本的模长噪声直接主导了距离计算。5-shot 时均值会把模长波动拉平1-shot 时没有任何平均机制模长离谱的支撑样本会让原型严重偏移。解决分两步。第一步在特征进入距离计算前做 L2 归一化第二步把支撑集和查询集的输入预处理保持一致特别是灰度图归一化的均值和方差不能训练用一个统计量、验证用另一个。还有一个小技巧1-shot 训练时可以在采样器里加入人工数据增强对支撑样本做随机旋转和位移后再取嵌入相当于用单张样本构造出多个虚拟视角缓解支撑样本信息不足的问题但这会增加训练时间只作为最后的救急手段。5.3 现象评估精度比训练精度还高看起来像是好事但往往说明评估设置太宽松指标不可信。最常见的根因是训练时 episode 里的 n_way 比评估时大比如训练 20-way、评估 5-way20-way 的难度远高于 5-way评估精度自然虚高。另外还有一个隐蔽陷阱如果某个类别的总样本数恰好等于 k_shot k_query那么这个类别在训练和评估的不同 epoch 里可能反复出现在支撑集和查询集的不同位置上数据泄漏边界变得模糊。解决方式是做一个严格的实验日志把训练 n_way、评估 n_way、每类样本数、支撑集查询集切分种子全部记录下来。评估精度高于训练精度本身不用慌张但要先确认不是泄漏造成的假象。最常见的回归测试是固定随机种子重跑评估三次看精度标准差。如果标准差超过 1.5%说明采样器噪声太大结果不具备可比性。5.4 现象不同类别的原型互相挤在一起训练结束后把原型向量投影到二维可视化发现不同类别的原型点混成一团几乎分不开。这背后是嵌入函数没有学到类别可分的特征。一个常见原因是损失只在查询样本上计算支撑样本本身不直接参与损失——如果 k_shot 很小原型对梯度的影响有限模型可能只把查询样本拉向正确的原型却没有动力把原型彼此推开。解决方法是引入一个辅助的类间约束项最简单的方式是在原型上直接加一个正则损失拉大不同原型之间距离def prototype_dispersion_loss(prototypes): normalized F.normalize(prototypes, dim1) similarity torch.matmul(normalized, normalized.t()) mask torch.eye(prototypes.size(0), deviceprototypes.device).bool() off_diag similarity[~mask] return off_diag.sum() / off_diag.numel()这个辅助损失把原型两两之间的余弦相似度当作惩罚项相似度越高惩罚越大。一般把权重设在 0.1 到 0.3 之间。加完之后训练 loss 会比原来大一点点但原型之间的分离度会明显改善。5.5 现象显存占用随 episode 增大翻倍支撑集和查询集的样本数一增大显存直接爆掉。这是因为训练循环把支撑集和查询集拼在同一个 batch 里前向整张计算图同时保留了支撑嵌入和查询嵌入。如果是 20-way 5-shot支撑集 100 张加查询集 100 张一次前向就要承担 200 张图的激活值。常用解决方法是把支撑集和查询集分成两次前向查询集前向时显式解耦梯度链with torch.no_grad(): support_feat encoder(support_data) prototypes compute_prototypes(support_feat, support_labels, n_way) query_feat encoder(query_data)这样支撑集前向不保留梯度显存占用大约减半代价是支撑特征的梯度不再传递给编码器嵌入函数只能通过查询样本学习。对于 5-shot 以上的配置这个限制影响不大对于 1-shot 配置建议保留支撑集梯度因为每类只有一张支撑样本它的梯度信息非常宝贵。另一个并行方案是改用梯度检查点但实现成本偏高不如直接调整 batch 大小和 episode 规模。6. 验证原型网络学到了什么t-SNE、跨数据集测试与错样本复盘6.1 用 t-SNE 检查原型与查询样本的边缘分布精度指标只能告诉你“对不对”不能告诉你“为什么对”。一个我很依赖的做法是把测试集里随机几个 episode 的查询样本嵌入和原型嵌入一起丢给 t-SNE观察查询样本是否真的围绕对应原型聚集。如果某个类别的查询样本明显偏向另一个类别的原型说明嵌入空间里这两个类混淆了需要检查是不是训练时类别采样不平衡。from sklearn.manifold import TSNE import matplotlib.pyplot as plt all_features torch.cat([query_feat, prototypes], dim0).cpu().numpy() tsne TSNE(n_components2, perplexity20, n_iter500).fit_transform(all_features)6.2 跨数据集交叉评估检验泛化而不是过拟合单一数据集上的精度可能有欺骗性。更可信的做法是拿在 Omniglot 上训练好的编码器直接在另一个数据集的 episode 上跑前向不做任何微调看精度还剩多少。这种跨数据集评估才符合少样本学习的初衷——模型应该学会的是“按支撑集做对比”的能力而不是记住某个数据集独有的特征分布。6.3 把错误样本画出来别只看精度最后一个习惯是把预测错误样本和对应原型图打印出来。错误样本往往集中在图像模糊、遮挡或背景杂乱的例子上而不是模型“没学会”。每一轮实验后我都会留三分钟看图复盘找出错误样本是不是普遍带有某种共同特征。这个习惯在低精度阶段远比调参有用。希望这个方案能帮你从跑通到跑出可信指标少走几趟弯路。本文还有配套的精品资源点击获取
返回列表