
简介基于PyTorch实现的原型网络Prototypical Networks代码包面向少样本学习研究者和机器学习开发者可在Omniglot数据集上快速复现经典少样本分类方法帮助读者理解小样本场景下的模型设计思路。包内共12个文件包含7个Python脚本、2张示意图、1份说明文档、1份开源协议和1个gitignore配置Python脚本覆盖数据批采样、原型损失计算、模型结构定义与完整训练流程模块划分清晰适合二次改造或嵌入其他项目。整个压缩包约135KB体积轻量、使用便捷目前已有920人学习下载。通过阅读源码和示意图可以系统掌握嵌入映射、原型生成、距离度量以及训练迭代等关键细节在Omniglot等数据集上进行对比验证无论是入门少样本学习还是复现论文实验这份代码都能作为可靠的基础实现参考。1. 原型网络小样本学习里最值得先复现的基线小样本分类Few-shot Learning的任务是让模型在每类只有 5 张甚至 1 张标注图时仍能对新类别做出可靠预测。Prototypical Networks 是这条赛道上绕不开的经典方案核心思想极其朴素把每个类别的支撑样本映射到特征空间后取均值作为「类原型」查询样本只需判断自己离哪个类原型最近。这个思路让它在 miniImageNet 等基准上达到了当时顶尖的水平而且训练稳定、显存占用低、代码量小——相比 MAML 那种二阶梯度方案原型网络几乎没有什么玄学调参非常适合作为你入局小样本学习的第一块跳板。本文就顺着这个标题从网络结构、episode 采样、训练评估到部署落地把一套可以直接跑通的 PyTorch 实现拆给你新手能照步骤复现熟手也能从中看到距离度量、特征提取器选择这些更深的门道。2. 先理解原型网络在算什么类均值、距离度量与 episode 机制2.1 原型网络的核心公式与直觉原型网络的出发点可以概括为一句话用类内样本的嵌入均值代表这个类。假设支撑集里有 N 个类别每个类别 K 个样本即 N-way K-shot先经过特征提取网络 f_φ 得到嵌入向量那么类别 c 的原型向量就是p_c (1 / |S_c|) * Σ f_φ(x_i)查询样本 x 的预测概率则来自 softmax 形式的距离度量p(y c | x) softmax(-d(f_φ(x), p_c))其中 d 通常取欧氏距离的平方。这里有两个值得说明的设计选择第一为什么用均值而不是像匹配网络那样对支撑集做注意力加权均值操作等价于对类内嵌入做高斯均值假设在 N-way K-shot 的设定下这个偏置能带来方差上的优势尤其当 K 很小时均值比注意力机制更能抵抗支撑样本本身的噪声。第二为什么用欧氏距离而不是余弦相似度论文里给过一个很关键的结论当特征提取器是线性映射时欧氏距离 原型均值对应着「高斯判别分析」的等价形式而余弦距离没有这种统计可解释性。你在动手改距离度量时这条结论是判断改法是否靠谱的基准线。2.2 Episode 采样训练数据组织方式决定了模型上限原型网络不在普通 batch 里训练而是使用episode任务机制。一个 episode 就是一个 N-way K-shot 的模拟任务先从训练集类别里随机抽 N 个类每个类抽 K 个支撑样本 Q 个查询样本模型在这一轮只对这几个类做分类。这样做的目的是让训练时的任务分布和测试时保持一致——测试时你面对的就是一个从未见过的 N-way 分类任务。常见的设置是 5-way 5-shot支撑集每类 5 张查询集每类 15 张左右。我在实际项目里还有一个经验episode 的采样尽量保证支撑集和查询集的图像来自同一批类别但拍摄条件或背景要有差异如果能做到的话这样模型会更早学会忽略背景干扰而不是记住类别专属的纹理。对大多数公开数据集来说直接用随机划分即可不需要刻意构造这种「域偏移」。2.3 从 metric learning 视角理解原型网络把小样本问题拆开看它实际上是「嵌入学习」和「度量选择」两件事的叠加。原型网络选择的路径是先用大规模数据预训练一个通用的特征提取器然后在 episode 训练中微调编码器使类内距离压缩、类间距离拉开。这个过程本质上就是 metric learning只不过它没有显式的对比损失而是靠分类交叉熵推动特征空间变形。理解这一点对工程落地很关键如果你要处理的数据集和预训练模型来源差异很大比如用 ImageNet 权重做医学切片迁移效果会大打折扣。常见做法是保留预训练权重的前几层只微调后几层或者干脆用更强的骨干网络替换。后面第 4 章我会给出骨干网络替换的具体操作和参数建议。3. 搭建最小可复现的 Prototypical Network模型结构与 episode 数据流3.1 工程目录与依赖准备开始写代码前先确认你的环境。这个项目只需要 PyTorch 和 torchvision不需要额外的第三方库。我的建议是用 conda 新建一个独立环境避免和本地其他项目的依赖版本冲突。安装时有一个高频坑GPU 版本的 PyTorch 需要先确认 CUDA 版本命令行里执行nvidia-smi看顶部 CUDA Version然后到 PyTorch 官网选对应的安装命令不要直接用pip install torch默认装 CPU 版。# 创建环境并安装依赖以 CUDA 12.1 为例 conda create -n proto python3.9 conda activate proto conda install pytorch torchvision pytorch-cuda12.1 -c pytorch -c nvidia如果你的机器没有独立显卡CPU 版也能跑通完整流程只是训练速度慢不少。后面所有代码我都按「无特殊情况无需改动」来写数据准备环节只依赖 torchvision 自带的数据集或标准文件夹结构。3.2 特征提取器封装卷积骨干与嵌入层原型网络的特征提取器可以用任何卷积网络但为了先跑通流程我建议从一个轻量级四层卷积网络开始。这个结构也是原型网络论文里使用的配置每层 64 个 3×3 卷积核每层后接 BatchNorm、ReLU 和 2×2 最大池化。它对 84×84 的输入图像会自动输出 64×5×5 的特征图展平后得到 1600 维嵌入。import torch import torch.nn as nn import torch.nn.functional as F class Conv64(nn.Module): 4层卷积特征提取器输出维度由in_channels和input_size共同决定 def __init__(self, in_channels3, hidden_dim64, input_size84): super().__init__() self.features nn.Sequential( nn.Conv2d(in_channels, hidden_dim, 3, padding1), nn.BatchNorm2d(hidden_dim), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # 84 - 42 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1), nn.BatchNorm2d(hidden_dim), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # 42 - 21 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1), nn.BatchNorm2d(hidden_dim), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # 21 - 10向下取整 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1), nn.BatchNorm2d(hidden_dim), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # 10 - 5 ) self.out_dim hidden_dim * 5 * 5 if input_size 84 else None def forward(self, x): return self.features(x).view(x.size(0), -1)这段代码有两个必须注意的地方。第一inplaceTrue的 ReLU 可以节省少量显存但如果你后续要在特征图上做梯度回传或修改操作记得去掉 inplace——这是一个容易让人排查半天的隐性坑。第二out_dim我在这里写死了 84×84 输入的推算值如果你把输入分辨率改了这个维度要同步改否则分类层的矩阵乘法会直接报维度错误。换成 ResNet 等预训练骨干时你在forward里加一个全局平均池化就能把输出维度固定下来。3.3 Episode 数据加载器N-way K-shot 的核心逻辑PyTorch 自带的DataLoader不能直接满足 episode 采样需求需要自己封装。这里我给你一个可直接复用的EpisodeSampler思路根据标签索引把同一个类别的样本索引归拢到一个字典里每次迭代时随机选 N 个类再从这 N 个类里各抽 K Q 个样本。这样返回的(x_query, y_query)中y_query已经是相对这个 episode 的序号而不是原始标签。class EpisodeSampler: 按N-way K-shot构建一个episode的支撑集与查询集 def __init__(self, labels, n_way5, k_shot5, k_query15): self.labels torch.as_tensor(labels) self.class_to_indices {} for idx, lab in enumerate(self.labels.tolist()): self.class_to_indices.setdefault(lab, []).append(idx) def get_episode(self): classes torch.randperm(len(self.class_to_indices))[:self.n_way] support_x, support_y [], [] query_x, query_y [], [] for new_cls_id, orig_cls in enumerate(classes.tolist()): indices torch.as_tensor(self.class_to_indices[orig_cls]) perm torch.randperm(len(indices)) support_idx indices[perm[:self.k_shot]] query_idx indices[perm[self.k_shot:self.k_shot self.k_query]] support_x.append(support_idx) support_y.append(torch.full((self.k_shot,), new_cls_id)) query_x.append(query_idx) query_y.append(torch.full((self.k_query,), new_cls_id)) return (torch.cat(support_x), torch.cat(support_y)), \ (torch.cat(query_x), torch.cat(query_y))写这个采样器时有三个细节值得遵守。第一同一类的支撑和查询样本从同一个randperm结果里切分避免重复采样造成「同类样本同时出现在支撑集和查询集」的数据泄漏——这是小样本任务里最容易翻车的地方后面避坑章节会再展开。第二每轮 episode 都重新randperm保证类别组合的多样性。第三k_query不宜太大否则单次 episode 的显存占用会线性增长常规设置为 15在支撑集只有 5 类 × 5 张 25 张的情况下每次前向只需处理 25 75 100 张图消费级显卡毫无压力。接下来是数据加载的主体流程。以你本地的文件夹格式数据为例图片组织方式为train/类别名/*.jpg我推荐先用torchvision.datasets.ImageFolder读一次把标签映射关系缓存下来再交给上面的采样器。常见的坑是忘记设置图像归一化用 ImageNet 均值方差或者忘了 resize导致预训练骨干的效果断崖式下降。from torchvision import datasets, transforms from torch.utils.data import DataLoader train_transform transforms.Compose([ transforms.Resize((84, 84)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) base_dataset datasets.ImageFolder(rootdata/train, transformtrain_transform) sampler EpisodeSampler(base_dataset.targets, n_way5, k_shot5, k_query15) # 每个epoch迭代多个episode for ep in range(200): s_idx, q_idx sampler.get_episode() # 后续按索引抓取图像和标签组装成张量即可ImageFolder会自动按文件夹名的字母序分配标签 0N-1这个映射关系在base_dataset.classes里能查到用于测试阶段还原真实类别名。我的建议是你把采样器封装成标准的IterableDataset这样就能直接交给 PyTorch 的DataLoader用num_workers开启多进程加载能显著缩短大批量 episode 的数据读取时间。4. 训练与评估闭环损失函数、反向传播与 few-shot 测试协议4.1 原型计算与损失函数实现拿到一个 episode 的支撑集嵌入和查询集嵌入后计算流程分三步对支撑集嵌入按类别求均值得到原型计算查询嵌入到所有原型的距离用距离取负做 log_softmax 得到交叉熵损失。这里距离函数我用的是欧氏距离的平方注意是平方而不是开方后的结果——因为这对应高斯判别分析里的马氏距离特例开方反而会扰动梯度尺度。def prototypical_loss(proto, query_emb, query_y): 原型网络的负对数似然损失返回损失值与预测准确率 dists torch.cdist(query_emb, proto) ** 2 # (n_query, n_way) log_probs F.log_softmax(-dists, dim1) loss F.nll_loss(log_probs, query_y) pred log_probs.argmax(dim1) acc (pred query_y).float().mean() return loss, acctorch.cdist是计算所有查询嵌入到所有原型距离的最直接方式底层有优化比手写两层循环快得多。如果你想替换距离度量比如改成余弦距离那需要先对嵌入做 L2 归一化再计算1 - cosine_similarity此时注意原型的计算应该用归一化前的均值还是归一化后的均值这个差异会直接影响精度建议自己跑对比实验确认。大多数公开基准里欧氏距离比余弦相似度高 12 个百分点这也是论文结论之外的实证共识。4.2 一个完整的训练循环整个训练循环写出来非常短这正是原型网络的魅力。每轮迭代采样一个 episode前向得到原型和预测算损失和准确率反向后更新参数。下面这段代码可以直接作为训练的起点其中 Adam 优化器和 1e-3 的学习率是论文和社区验证过的合理默认值。model Conv64(in_channels3) optimizer torch.optim.Adam(model.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size20, gamma0.5) for epoch in range(100): model.train() running_loss, running_acc 0.0, 0.0 for _ in range(100): support_x, support_y, query_x, query_y sample_episode(...) support_emb model(support_x) query_emb model(query_x) proto torch.stack([support_emb[support_y c].mean(0) for c in range(n_way)]) loss, acc prototypical_loss(proto, query_emb, query_y) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() running_acc acc.item() scheduler.step() print(fEpoch {epoch}: loss{running_loss/100:.4f}, acc{running_acc/100:.4f})这段代码里三个参数值得你反复调试。第一lr1e-3在小型 Conv64 网络上通常最优但如果换成 ResNet-18 这类预训练骨干我会把学习率降到 1e-4 附近防止大步长破坏预训练权重。第二StepLR的step_size和gamma不必照搬如果你的训练总轮数只有 50 轮每 10 轮降一半会更合理。第三sample_episode里支撑集和查询集的图像要经过完全相同的预处理不能训练时一个有归一化、验证时另一个忘了加——这是我见过最频繁的低级错误之一。4.3 测试协议搭一个严谨的 few-shot 评估循环测试时的小样本协议和训练很相似但有三个差别支撑集样本是固定的不再随机抽查询集来自未见过的测试类评估要在多个随机种子下重复取平均。严谨的做法是把测试类别里每类随机抽 K 个当作支撑集剩下的作为查询集循环 1000 次不同随机划分取平均准确率作为最终指标同时记录 95% 置信区间。def evaluate(model, test_dataset, n_way5, k_shot5, k_query15, n_episodes1000): model.eval() accs [] with torch.no_grad(): for _ in range(n_episodes): s_idx, q_idx sampler.get_episode() # 注意这里要用测试集的EpisodeSampler support_x, support_y build_tensors(test_dataset, s_idx) query_x, query_y build_tensors(test_dataset, q_idx) support_emb model(support_x) query_emb model(query_x) proto torch.stack([support_emb[support_y c].mean(0) for c in range(n_way)]) _, acc prototypical_loss(proto, query_emb, query_y) accs.append(acc.item()) mean_acc torch.tensor(accs).mean().item() interval 1.96 * torch.tensor(accs).std().item() / (n_episodes ** 0.5) print(fFew-shot acc: {mean_acc:.2f}% ± {interval:.2f}%) return mean_acc评估时最容易忽视的是「支撑集和查询集不可重叠」这条铁律。测试阶段用get_episode时支撑分区的索引和查询分区必须来自同一个类索引列表但切分位置不能重复。如果你在测试时使用和训练时相同的随机源而没有做这个区分语义上就等价于把答案提前给模型看了一眼得到的精度虚高 1020 个百分点这是小样本学习论文里反复强调的「作弊评估」。另外遇到测试类别数少于n_way的情况EpisodeSampler会直接报错你需要先确认数据集类别数量、再做min(n_way, len(classes))的保护处理。5. 五个高频踩坑记录从显存溢出到数据泄漏5.1 传递索引张量时 CPU/GPU 设备不一致导致训练中断现象训练跑到第二个 epoch 就报RuntimeError: expected device cuda:0 but got device cpu而且报错位置在原型计算那行的support_y c布尔索引上。原因support_y是在EpisodeSampler里用 CPU 张量构造的模型和数据被搬到了 GPU但标签索引用它来筛选嵌入时PyTorch 自动要求所有参与索引的张量在同一设备上。解决在sample_episode返回张量后统一调用.to(device)或者把support_y和query_y直接构造在device上。我更推荐前者因为标签本来就不占显存CPU 上维护反而更灵活。5.2 数据泄漏同类图像同时出现在支撑集和查询集现象训练时准确率飙升到 95% 以上但测试集上只有 60%差距大得不合理。原因采样器先对全类索引做randperm然后从perm[:k_shot]和perm[k_shot:k_shotk_query]取两个分区。如果k_shot k_query大于该类样本总数就会导致支撑集和查询集索引重叠。很多公开数据集的类内样本量恰好接近阈值不经意间就踩中了。解决采样前显式检查len(indices) k_shot k_query不足的在抽到该类时跳过并重新采样。更稳妥的做法是按类别分别randperm后再拼接。5.3 批量归一化在小 batch 下的「最后一坑」现象支撑集只有 5 类 × 5 张 25 张时BatchNorm2d的统计量不稳定训练波动大甚至出现 acc 来回跳的情况。原因BN 层在 batch 维度统计均值方差batch 只有 25 时估计噪音很大。小样本任务里这个问题被放大了。解决把骨干网络中的BatchNorm2d换成GroupNorm(num_groups8)或者InstanceNorm2d。在不改变网络整体结构的前提下这个替换通常能带来 13 个百分点的稳定提升而且对 batch 大小的敏感度大幅下降。这已经成了小样本社区默认的工程改法。5.4 固定num_workers过大会导致数据加载死锁现象训练时偶尔卡死GPU 利用率掉到 0终端无报错。原因IterableDataset配合DataLoader多进程时采样器状态在多个 worker 之间互相干扰典型症状是每个 worker 重复采样或阻塞等待。解决如果用了IterableDataset且内部维护状态num_workers设为 0 最稳妥想要加速就把采样逻辑改成无状态函数或直接在__iter__内重新初始化。就本项目的计算量而言瓶颈其实在 GPU 前向数据加载很少成为性能瓶颈不必在此过度优化。5.5 替换预训练骨干时输出维度对不上导致前向失败现象把Conv64换成了torchvision.models.resnet18(pretrainedTrue)之后model(support_x)直接报维度不匹配或最后分类层输入尺寸错误。原因ResNet 的forward输出是 512 维但代码里out_dim仍按64 * 5 * 5 1600设定。解决替换骨干时在forward末尾加一个F.adaptive_avg_pool2d(1)把特征图压成 1×1再view(-1, 512)这样输出维度永远和网络结构解耦。我封装的写法如下class ResNetBackbone(nn.Module): def __init__(self): super().__init__() from torchvision import models self.base models.resnet18(pretrainedTrue) self.base.fc nn.Identity() # 去掉最后的全连接层 def forward(self, x): return self.base(x)ResNet 的全局平均池化已经包含在内部输出正好是 512 维无需手动再池化。如果你用 ViT 之类的 Transformer 骨干同样要关注[CLS] token或池化头的输出位置不同实现的输出约定差别很大。6. 进阶技巧从可视化到跨数据集验证的完整检查清单模型训好后不要急着收工我看大多数踩坑都发生在「训练曲线还行、但不知道模型到底学到了什么」的阶段。这里给你一套我在实际项目中常用的验证与进阶流程。第一步是嵌入空间可视化。取测试集某个 episode 的嵌入用 t-SNE 或 UMAP 降到二维后散点图检查支撑样本的嵌入是否聚成 N 个清晰的簇查询样本是否围绕各自的类原型分布。如果聚类结构混乱、类间重叠严重说明特征提取器训练不充分优先调学习率或增加训练 epoch如果聚类清晰但查询点分散在簇边缘说明距离度量可能太敏感可以试一下对嵌入做 L2 归一化后再算欧氏距离。第二步是跨数据集的泛化实验这是检验模型是否过拟合到训练集类别的试金石。建议准备一个和训练集分布差异明显的第二数据集比如从自然图像换到医学影像把测试协议原样跑一遍。如果精度掉幅超过 15 个百分点说明骨干网络的底层特征过于偏向源域你应该冻结前几层只微调后几层或者增加数据增强随机裁剪、色彩抖动、Cutout 等来逼迫模型学到更通用的形状特征。第三步是工程化收敛。当你确认精度达标后把模型导出为 TorchScript 或 ONNX 格式用于在线推理或边缘设备部署。导出时记得嵌入特征归一化和距离计算也纳入导出范围否则部署推理时手写欧氏距离计算容易出错。推荐的导出方式是直接用torch.jit.trace因为整个前向流程没有控制流分支trace 足够稳定。我的习惯是在每个新数据集上至少跑三组随机种子取均值再下结论小样本任务的方差比常规监督学习大得多单次运行的结果只能当参考。训练时也把每个 episode 的支撑类列表打印出来检查过一轮确保类别分布符合预期。这个方向你能投入的潜力很大从替换骨干、调距离度量到引入数据增强每一步改动都能用一套统一的评估协议量化希望你从这个最小实现开始走通自己的 few-shot 流程也希望这篇笔记能成为你调试时的一本速查手册。本文还有配套的精品资源点击获取