资讯详情 Prototypical Networks原理与PyTorch实现:少样本学习的实战指南
📅 2026/10/11 0:50:53
简介面向少样本学习研究者的 PyTorch 实现资源完整复现了原型网络这一经典方法适用于图像分类、小数据量场景下的模型训练与算法对比。压缩包共 12 个文件整体约 135KB以 7 个 Python 源码文件为主干分别负责数据集加载、批次采样、原型编码、损失计算和训练过程另有 2 张 PNG 图示、1 份 Markdown 说明、许可证文件与 gitignore 配置。源码模块划分清晰包含 Omniglot 数据读取、原型网络模型、原型损失函数、训练脚本等部分可直接在 Omniglot 数据集上运行训练与评估也可替换数据接口应用到自定义分类数据集。随包文档和图示对网络结构、原型距离与更新流程做了辅助说明便于学习阅读。目前已有 920 人学习下载适合具备一定 PyTorch 基础、希望掌握少样本学习实现细节的机器学习开发者参考。1. Prototypical Networks 是什么为什么少样本任务先跑它手头每类只有三五张标注图分类器怎么训都像在碰运气这时候 Prototypical Networks 是 few-shot learning 里最值得先试的方案。它用一个嵌入网络把样本映射到特征空间按类别算出“原型”再用欧氏距离判断查询样本离哪个类最近。整个流程在 PyTorch 里实现起来非常直白没有对抗训练也没有复杂的元学习外层优化。适合两类人一类是刚入坑少样本学习、想先把数据到模型的 pipeline 跑通的工程师另一类是手里有少量标注数据、想找一个比直接微调更稳的训练策略的人。下面从原理讲到可复现代码再把这些年容易踩的边界坑挑出来说清楚。2. 原型网络的核心逻辑嵌入空间、类原型与欧氏距离2.1 分类头为什么在少样本场景下失效先看普通图像分类器的结构一个卷积主干后面接一个全连接层输出每个类别的 logits。训练时卷积主干由所有类别共享但全连接层里每一类都有一组独立权重。想让“第 10 类”的权重学准就得有足够多“第 10 类”的样本参与训练。而在 few-shot learning 的标准设定里每个类别只有 1 到 5 个支撑样本support set可用。用五六个样本去拟合一组高维分类权重结果就是权重几乎完全欠定训练集上能压到零误差换一批查询样本立刻翻车。原型网络换了一个思路不再保存任何类别专属参数而是在每个任务episode里动态计算类别中心。给定一个嵌入函数f_phi第 k 类的原型就是该类别所有支撑样本嵌入向量的均值。查询样本进来之后算它与每个原型之间的距离距离最近的类就是预测结果。当支撑样本只有 1 个时它退化成最近质心分类器nearest centroid classifier。这个设计让分类器的“容量”不再依赖类别数而是只依赖嵌入空间的质量所以每类样本再少也能跑起来。这套设计在 pytorch 实战里的实现成本很低关键是你不需要在模型里预留num_classes个参数。模型的最后一层不再是一个 Linear而是“均值计算 距离计算”这两个操作都是动态的换一 arms 任务、换一批类别都不需要改网络结构。2.2 原型怎么算欧氏距离为什么够用原型网络的数学表达很简短。设第 k 类的支撑集为S_k嵌入函数为f_phi那么原型是c_k 1 / |S_k| * sum(f_phi(x)), x in S_k给定一个查询样本 x先算它的嵌入f_phi(x)再算它与每个原型 c_k 的欧氏距离d_k || f_phi(x) - c_k ||^2把负距离放进 softmax得到 x 属于第 k 类的概率。训练时最小化交叉熵损失loss -log p(y k | x) logsumexp(d_all) - d_k注意这里没有额外引入可学习参数唯一的可学习部分就是嵌入网络f_phi。这也解释了为什么它适合作为少样本学习的 baseline你把注意力全放在“怎么把样本嵌入到好的特征空间”上而不是放在“怎么调分类头”。欧氏距离和求均值是一对天然组合。均值的定义本身就是最小化类内欧氏距离平方和的点也就是说如果你要用一个向量代表一个类均值在欧氏距离下是最优代表。这一点让原型网络在理论上比余弦距离更自洽。实际使用中余弦距离等价于对向量做 L2 归一化后再算欧氏距离很多时候效果差别不大但欧氏距离可以直接用torch.cdist一行算完梯度也更干净所以默认先用欧氏距离没有毛病。2.3 Prototypical 与 Matching Networks、Siamese 的差别少样本学习里常见的方法还有 Matching Networks 和 Siamese Networks。三者都基于“嵌入 距离比较”但细节差异决定了它们适合的场景不同。方法分类依据额外可学习参数更适合的场景Prototypical Networks查询样本到类原型的欧氏距离无类别多、每类支撑样本少追求简单稳定Matching Networks查询样本到每个支撑样本的注意力加权和无支撑样本内部差异大需要逐样本比较Siamese Networks样本对之间的二分类相似度相似度阈值人脸验证、签名验证这种两两对比任务实际选型时我一般这样判断如果业务目标是“给定一个查询样本在 N 个类里选一个”优先跑原型网络如果业务目标是“判断两张图是不是同一类”Siamese 更直接如果支撑集里某些样本特别有代表性、某些样本很噪声Matching 的注意力机制会更鲁棒。不过真到落地阶段原型网络通常作为第一个 baseline因为它几乎没有需要调的黑匣子效果不够再用更复杂的方法去替换。3. 用 PyTorch 搭建最小实现模型定义与参数初值3.1 一个能出 loss 的最小模型定义环境不需要多复杂装好 PyTorch 和 torchvision 就能开始。下面这个模型是我常用的最小版本四个卷积 block 做嵌入输出一个固定维度的向量。代码里有个细节最后一个 block 不做 MaxPool而是用 AdaptiveAvgPool2d(1) 把空间维度压成 1x1这样换输入分辨率也不会崩。import torch import torch.nn as nn def conv_block(in_ch, out_ch, do_poolTrue): layers [ nn.Conv2d(in_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ] if do_pool: layers.append(nn.MaxPool2d(2)) return nn.Sequential(*layers) class ProtoNet(nn.Module): def __init__(self, in_ch3, hidden64): super().__init__() self.encoder nn.Sequential( conv_block(in_ch, hidden), conv_block(hidden, hidden), conv_block(hidden, hidden), conv_block(hidden, hidden, do_poolFalse), ) self.avgpool nn.AdaptiveAvgPool2d(1) def forward(self, x): # 输入 x: [batch, in_ch, H, W] h self.encoder(x) # 输出形状: [batch, hidden, 1, 1]view 成 [batch, hidden] return self.avgpool(h).view(x.size(0), -1)这段代码的逻辑是输入任意分辨率的图像经过四个卷积 block 后得到[batch, hidden, H, W]再用 AdaptiveAvgPool2d 强制压到[batch, hidden, 1, 1]。biasFalse是因为后面紧跟 BatchNorm2dBatchNorm 自带可学习的平移参数卷积再加 bias 属于冗余。hidden64是少样本任务里比较稳的嵌入维度太小表达不够太大在小数据集上容易过拟合。前向方法里还需要一个处理 episode 的函数。episode 是少样本训练的基本单位把一个支撑集和一个查询集一起喂进去class ProtoNet(nn.Module): # ... 上面 encoder 部分省略 def forward_episode(self, support, query, n_shot): # support: [n_way * n_shot, in_ch, H, W] # query: [n_way * n_query, in_ch, H, W] support_emb self.forward(support) # [n_way * n_shot, hidden] query_emb self.forward(query) # [n_way * n_query, hidden] n_way support.size(0) // n_shot # 按类别把支撑样本分组每组 n_shot 个 support_emb support_emb.view(n_way, n_shot, -1) prototypes support_emb.mean(dim1) # [n_way, hidden] # 查询样本与所有原型的欧氏距离 dist torch.cdist(query_emb, prototypes) # [n_way * n_query, n_way] return -disttorch.cdist(query_emb, prototypes)一次算出所有查询样本和所有原型之间的欧氏距离结果是[n_way * n_query, n_way]的矩阵。返回负距离作为 logits后续接CrossEntropyLoss即可。注意n_way是从支撑集第一维除以n_shot推断出来的所以调用时n_shot必须传对否则分组全乱。3.2 关键参数初值与调参方向少样本训练里参数设对一半结果就稳了一半。下面这张表是常见的初值和建议范围。参数初值或建议范围说明n_way5标准评估、按业务类别数每个 episode 里的类别数越大越难n_shot1 或 5每个类支撑样本数1-shot 是最难的设定n_query5 到 15每个类查询样本数影响梯度稳定性输入尺寸84x84miniImageNet、32x32小图固定一个尺寸不要训练验证各用各的hidden64嵌入维度小数据 64 足够更大不一定更好优化器Adam lr1e-3收敛快SGD 配 lr1e-2 更稳但更慢调参时我个人会先固定 5-way 5-shot 跑通再看业务里实际类别数和可用的标注数量。如果业务里每个类只有 2 张图那就把 2-shot 作为核心配置n_query 设成 5 左右。way 数越大任务越难如果模型在 5-way 上都过拟合就不要急着上 20-way。3.3 骨干网络选型先用小网络打底很多入门教程一上来就用 ResNet18但在少样本场景下这是把双刃剑。深层网络容量大在小样本上过拟合更快而且预训练权重如果和你的数据分布差异大迁移效果反而差。常见做法是先跑通上面的四层卷积网络把它作为 baseline确认 pipeline 没问题之后再换 ResNet18 或 ViT。换主干时只需要保证 forward 输出一个固定维度的向量其他部分不用动。如果你要处理的是灰度图比如 Omniglot 那种把in_ch1即可。如果图片分辨率不是 84x84也不用改代码AdaptiveAvgPool2d 会保证输出维度一致但要注意分辨率太小会丢失细节建议所有输入统一缩放不要训练用大图、推理用小图。4. Episode 训练循环采样器、loss 与张量形状变化4.1 先从数据集中采样一个 episodefew-shot learning 的训练和普通分类完全不同它不按 batch 从全量数据里随机抽样本而是按 episode 组织任务。每个 episode 从训练集里随机选 n_way 个类每个类再随机选 n_shot n_query 个样本前 n_shot 个作为支撑集后 n_query 个作为查询集。import random def build_label_to_indices(dataset): dataset 里每个元素是 (x, label)返回 label - 下标列表 的映射 label_to_idx {} for i, (_, label) in enumerate(dataset): label_to_idx.setdefault(label, []).append(i) return label_to_idx def sample_episode(label_to_idx, n_way, n_shot, n_query): classes random.sample(sorted(label_to_idx.keys()), n_way) support_idx, query_idx [], [] for cls in classes: # 每个类随机抽 n_shot n_query 个样本不重复 idx random.sample(label_to_idx[cls], n_shot n_query) support_idx idx[:n_shot] query_idx idx[n_shot:] return support_idx, query_idx, classes这里必须用random.sample而不是random.choices前者是不放回抽样保证同一个 episode 里支撑集和查询集不会出现同一条样本。如果某个类的样本数少于 n_shot n_query会在random.sample直接抛异常所以构建数据集时最好过滤掉样本过少的类或者把 n_shot 调小。这个采样顺序很关键它决定了后续 target 的排列方式后面训练循环里会看到。build_label_to_indices兼容torchvision.datasets.ImageFolder如果你的数据是目录结构直接用它构建映射即可。训练时我从不把数据集打乱后重新分配 label而是让 label 保持原始语义这样采样时类别不会互相污染。4.2 训练循环里每个张量的形状变化拿到支撑集和查询集的下标后把它们转成张量喂给forward_episode再算交叉熵。下面是一个完整的训练循环import torch import torch.nn.functional as F n_way, n_shot, n_query 5, 5, 5 episodes_per_epoch 100 num_epochs 30 model ProtoNet(in_ch3, hidden64).cuda() optimizer torch.optim.Adam(model.parameters(), lr1e-3) label_to_idx build_label_to_indices(train_dataset) for epoch in range(num_epochs): total_loss 0.0 for _ in range(episodes_per_epoch): s_idx, q_idx, _ sample_episode(label_to_idx, n_way, n_shot, n_query) support torch.stack([train_dataset[i][0] for i in s_idx]).cuda() query torch.stack([train_dataset[i][0] for i in q_idx]).cuda() # target 排列与 query_idx 的顺序一致 # 前 n_query 个属于 classes[0]接着 n_query 个属于 classes[1]依此类推 target torch.arange(n_way).repeat_interleave(n_query).cuda() scores model.forward_episode(support, query, n_shot) loss F.cross_entropy(scores, target) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() print(fepoch {epoch} loss {total_loss / episodes_per_epoch:.4f})逐行解释关键部分。support的原始形状是[n_way * n_shot, C, H, W]也就是 25 张图经过forward_episode后变成[25, 64]。query同样从[n_way * n_query, C, H, W]变成[25, 64]。target用torch.arange(n_way).repeat_interleave(n_query)生成结果是[0,0,0,0,0, 1,1,1,1,1, ...]和查询集的排列顺序严格对应。如果sample_episode里类别顺序是随机抽出来的 classes 列表那么查询集前五个样本必然属于第一个类所以 target 从 0 开始递增这个对齐关系是整个训练不出 bug 的核心。loss 用的是CrossEntropyLoss它内部自带 softmax输入scores是负距离矩阵。你也可以改成F.log_softmax(scores, dim1)配合NLLLoss效果等价但直接用CrossEntropyLoss少一步。反向传播时梯度只经过嵌入网络prototype 是通过均值算出来的不需要也不应该有梯度。4.3 每个 epoch 重新采样还是固定 episode训练阶段我一般每个 epoch 都重新采样 episode这样模型每个 epoch 看到的任务都不一样泛化更好。代价是 loss 曲线波动比较大因为不同 episode 的难度不一样这是正常现象不用看到 loss 反弹就慌。验证阶段正好相反一定要固定一批 episode不要每次重新随机采样。常见做法是训练结束后用固定随机种子生成 20 个 episode把支撑集、查询集和 target 全部缓存下来每次评估都在同一批任务上跑。如果不固定同一个模型两次评估的准确率可能差好几个点你会以为模型坏了其实只是采样方差在捣乱。关于固定验证 episode 的具体实现我放在第 6 章展开这里先记住一个原则训练随机、验证固定。这个原则也适用于对比实验所有候选模型必须在同一批验证任务上比较否则结论不可信。5. 避坑与排查少样本训练里最容易翻车的 5 个细节5.1 Loss 不降且接近 ln(n_way)先查标签对齐现象模型训练了十几轮loss 稳定在 1.6 附近5-way 时 ln(5)≈1.61准确率几乎没有上升。原因scores的列顺序和target的类别顺序没有对齐。比如支撑集的类别顺序是随机抽的但查询集的 target 还是用的固定顺序模型学到的是一个错误映射loss 降不下去因为它无法同时满足两个互相矛盾的监督信号。解决在第一个 epoch 里打印scores[:5]和target[:5]确认前 5 行确实对应第 0 类。更简单的方法是打印sample_episode返回的classes列表再对照查询集前 5 张图的实际 label确保类顺序一致。这个坑几乎每个刚写原型网络的人都会踩一次排查方式就是打印、打印、再打印。5.2 训练准确率高而验证低way/shot 配置不一致现象训练集上 query accuracy 到了 95%换到验证集只剩 45%明显不合理。原因训练和验证用了不同的 n_way 或 n_shot。比如训练用 5-way 5-shot验证却用了 5-way 1-shot任务变难准确率当然跳水。另一种可能是验证 episode 每次都重新随机采样导致评估结果方差极大看起来像没训好。解决把训练和验证的参数完全统一。先固定 5-way 5-shot 跑通再去测 1-shot 等更难的设置。验证 episode 缓存固定消除采样方差。如果不统一基础配置后面的调参全是在噪声里找信号。5.3 Loss 出现 NaN 或距离矩阵溢出特征范数失控现象训练到中途 loss 突然变成 NaN或者 loss 在正常和 NaN 之间反复横跳。原因嵌入网络输出的特征向量范数过大欧氏距离的平方在torch.cdist里数值溢出。深层网络在少样本训练中很容易把特征模长推到很大的值尤其是没有做归一化时。解决在forward_episode里对嵌入向量做 L2 归一化让范数限制在 1 附近距离上限也就被限制住梯度更稳定。另一个有效手段是调低学习率Adam 用 1e-3 如果崩了降到 5e-4 再试。如果换了数据分布优先查输入有没有归一化到相同量纲原始像素和标准化像素训练出来的特征范数差异很大。5.4 原型计算按错维度模型没报错效果全废现象程序不报错loss 也下降但验证 accuracy 一直在一半以下怎么调学习率都没用。原因support_emb.view(n_way, n_shot, -1).mean(dim1)里的dim写错。写成dim0会跨类别求均值把所有类的原型混成一个模型直接失去分类能力。这个 bug 很隐蔽因为 loss 仍然会下降只是模型学到的东西毫无意义。解决在forward_episode里加一行调试代码打印prototypes.shape确认是[n_way, hidden]而不是[n_shot, hidden]或[n_way * n_shot, hidden]。更稳妥的做法是写一个单元测试随机生成 support 和 query跑一次前向手动验证支持集第 0 类均值等于第一个 prototype。5.5 预训练主干反而过拟合特征分布不匹配现象加载 ImageNet 预训练的 ResNet18 后loss 比小网络降得更快但验证集上表现反而更差。原因预训练特征是在大规模自然图像上学到的如果你的业务数据是特定领域那些特征不一定适配你的域反而会把 domain-specific 的噪声带进来。尤其在少样本场景下模型容量太大迁移特征对小样本任务的贡献可能是负的。解决先跑四层小卷积网络确认 baseline 稳定之后再尝试预训练主干并且要冻结 BatchNorm 的统计参数用model.eval()模式做推理。我的血泪经验是小样本任务里简单网络加合理正则往往比复杂预训练主干更稳。如果你在 PyTorch 里做对比实验一定要把预训练模型的优化器参数单独设置不要和新增层用同一个学习率。6. 跑通之后怎么验证固定 episode、评估指标与可视化6.1 固定一批验证 episode别再让结果随机漂少样本评估最大的敌人是采样方差。跑通训练循环之后先做一件让结果可复现的事用固定随机种子预生成 20 个验证 episode保存下标序列每次评估都用同一批数据。import numpy as np import torch def evaluate(model, eval_episodes, n_way, n_shot, n_query): model.eval() accs [] with torch.no_grad(): for s_idx, q_idx in eval_episodes: support torch.stack([train_dataset[i][0] for i in s_idx]).cuda() query torch.stack([train_dataset[i][0] for i in q_idx]).cuda() target torch.arange(n_way).repeat_interleave(n_query).cuda() scores model.forward_episode(support, query, n_shot) pred scores.argmax(dim1) acc (pred target).float().mean().item() accs.append(acc) return np.mean(accs), np.std(accs)报告结果时报 mean 和 std不要只报一个平均值。比如“acc 79.2% ± 2.1%”比“acc 79.2%”可信得多。如果 std 超过 3 到 4 个点说明评估任务太少或者模型本身不稳定先增加验证 episode 数量再看模型是否真有改进。6.2 三个指标加一张散点图确认原型空间是聚拢的只看准确率不够还要看原型空间的几何。我一般会补三个检查第一个是 query 的平均置信度如果精度不错但置信度低说明模型“犹豫但碰巧对了”第二个是类原型两两之间的距离如果某些原型间距特别近那两个类在嵌入空间里几乎重叠容易混淆第三个是每个类 query 样本和本类原型的平均距离这个值远大于该类的类内支撑样本平均距离时说明泛化有问题。可视化也很有用把嵌入向量用 PCA 压到二维平面每个类的原型画成大圆点查询样本画成小圆点正常情况是每个类一坨查询点紧贴原型。如果某个类的查询点散落全场说明这个类没学好回去检查这个类的支撑样本质量。t-SNE 可视化更漂亮但慢小样本场景里先看 PCA 足够定位问题。如果要把模型部署出去注意原型网络的前向依赖 episode 的拼接方式转 ONNX 时建议把整个forward_episode包装成一个模块导出不要在部署代码里用 Python 原生循环拼接张量否则动辄被速度卡脖子。我自己最早做这个方向时最常翻车的不是模型设计而是验证不固定、评估参数乱调结果同一个模型跑两次差好几个点还以为哪里改坏了。后来把验证 episode 全部缓存下来所有改动都在这同一批任务上比较问题定位快了不止一倍。希望帮到你。本文还有配套的精品资源点击获取