Prototypical Networks实战5步搞定小样本分类附PyTorch代码当你的训练数据只有每个类别寥寥几张图片时传统深度学习方法往往会束手无策。这正是小样本学习大显身手的场景——而Prototypical Networks原型网络作为其中最优雅的解决方案之一仅需5个核心步骤就能构建出可用的分类系统。今天我们就用PyTorch从零开始拆解这个看似高深实则简单的算法。1. 理解原型网络的核心思想想象你第一次见到长颈鹿时即使只看过几张照片下次再见到不同姿势的长颈鹿也能认出来。人类这种从小样本学习的能力正是原型网络试图在机器学习中复现的。其核心可以用三个关键词概括原型(Prototype)每个类别的平均形象就像你脑海中典型长颈鹿的概念度量空间(Metric Space)所有数据都被映射到一个空间在这里同类样本紧密聚集距离分类新样本通过计算与各类原型的距离来决定类别归属与需要成对比较的孪生网络不同原型网络的计算效率显著提升。我们来看一个直观的例子方法计算复杂度新增类别成本传统分类器O(n)需重新训练孪生网络O(n²)无需调整原型网络O(n)即时适应# 伪代码展示原型网络分类逻辑 def classify(query, prototypes): distances [euclidean_distance(query, p) for p in prototypes] return softmax(-distances) # 距离越近概率越高2. 准备小样本数据集实战中我们使用Omniglot——包含50种文字系统的1623个字符每个字符仅有20个手写样本。这个数据集完美模拟了现实中的小样本场景from torchvision.datasets import Omniglot from torchvision.transforms import Compose, Resize, ToTensor transform Compose([Resize(28), ToTensor()]) dataset Omniglot(root./data, downloadTrue, transformtransform)关键步骤是构建episode——元学习特有的数据组织方式。每个episode包含N-way随机选择N个类别如5类K-shot每类取K个样本作为支持集如每类5张图Q-query每类取Q个查询样本用于计算损失如每类15张图def sample_episode(dataset, n_way5, k_shot5, q_query15): classes random.sample(dataset._character_images.keys(), n_way) support, query [], [] for cls in classes: imgs random.sample(dataset._character_images[cls], k_shot q_query) support.extend(imgs[:k_shot]) # 支持集样本 query.extend(imgs[k_shot:]) # 查询集样本 return torch.stack(support), torch.stack(query)提示实际应用中你可能需要自定义数据集类来封装这些逻辑确保数据加载的高效性3. 构建嵌入网络原型网络的核心是一个将原始输入映射到度量空间的嵌入函数fφ。这个网络不需要特殊设计一个简单的CNN就能胜任import torch.nn as nn class EmbeddingNet(nn.Module): def __init__(self): super().__init__() self.net nn.Sequential( nn.Conv2d(1, 64, 3), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Flatten() ) def forward(self, x): return self.net(x)为什么这个结构有效通过连续的卷积和池化操作网络逐步学习到低级特征边缘、纹理中级特征部件、形状高级语义特征整体结构关键细节最后一层不使用全连接保持空间结构不变这对后续的距离计算至关重要。4. 实现原型计算与分类现在来到算法的核心部分——原型计算和分类过程。这个阶段需要精确实现三个数学操作原型计算$c_k \frac{1}{|S_k|} \sum_{x_i \in S_k} f_\phi(x_i)$距离度量通常使用欧氏距离$d(z, c_k) ||z - c_k||^2$概率分布$p(yk|z) \frac{\exp(-d(z, c_k))}{\sum_{k} \exp(-d(z, c_{k}))}$PyTorch实现如下def compute_prototypes(support, support_labels, n_way): 计算每个类别的原型向量 prototypes [] for k in range(n_way): # 选出当前类别的所有支持样本 mask (support_labels k) class_embeddings support[mask] prototypes.append(class_embeddings.mean(dim0)) return torch.stack(prototypes) def prototypical_loss(query, query_labels, prototypes, n_way): 计算原型网络的损失 distances torch.cdist(query, prototypes) # 计算查询样本与所有原型的距离 log_p -distances log_p_y log_p.gather(1, query_labels.unsqueeze(1)) loss -log_p_y.mean() acc (log_p.argmax(dim1) query_labels).float().mean() return loss, acc注意距离计算有多种选择欧氏距离最常用但余弦相似度在某些场景可能表现更好5. 训练策略与技巧原型网络的训练有其特殊性需要特别注意以下要点训练循环设计每个episode随机采样新的类别组合计算当前episode的原型评估查询集上的损失反向传播更新嵌入网络参数def train_epoch(model, optimizer, dataset, n_way5, k_shot5, q_query15): model.train() total_loss, total_acc 0, 0 for _ in range(100): # 每个epoch 100个episode support, query sample_episode(dataset, n_way, k_shot, q_query) support_emb model(support) query_emb model(query) prototypes compute_prototypes(support_emb, support_labels, n_way) loss, acc prototypical_loss(query_emb, query_labels, prototypes, n_way) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() total_acc acc.item() return total_loss / 100, total_acc / 100提升性能的实用技巧学习率调度使用ReduceLROnPlateau根据验证损失调整学习率嵌入维度最后一层特征维度通常设为64-256之间数据增强对小样本尤为重要可尝试随机旋转(-15°, 15°)轻微透视变换弹性变形# 增强后的transform from torchvision.transforms import RandomAffine, ElasticTransform train_transform Compose([ RandomAffine(degrees15, translate(0.1,0.1), scale(0.9,1.1)), ElasticTransform(alpha20.0), Resize(28), ToTensor() ])在实际项目中我发现两个常见陷阱值得注意原型漂移当支持样本过少时原型可能无法准确代表类别。解决方法是在测试时使用全部支持样本计算原型。维度灾难嵌入空间维度太高会导致距离度量失效。通过实验找到适合你数据集的维度很关键。
Prototypical Networks实战:5步搞定小样本分类(附PyTorch代码)
Prototypical Networks实战5步搞定小样本分类附PyTorch代码当你的训练数据只有每个类别寥寥几张图片时传统深度学习方法往往会束手无策。这正是小样本学习大显身手的场景——而Prototypical Networks原型网络作为其中最优雅的解决方案之一仅需5个核心步骤就能构建出可用的分类系统。今天我们就用PyTorch从零开始拆解这个看似高深实则简单的算法。1. 理解原型网络的核心思想想象你第一次见到长颈鹿时即使只看过几张照片下次再见到不同姿势的长颈鹿也能认出来。人类这种从小样本学习的能力正是原型网络试图在机器学习中复现的。其核心可以用三个关键词概括原型(Prototype)每个类别的平均形象就像你脑海中典型长颈鹿的概念度量空间(Metric Space)所有数据都被映射到一个空间在这里同类样本紧密聚集距离分类新样本通过计算与各类原型的距离来决定类别归属与需要成对比较的孪生网络不同原型网络的计算效率显著提升。我们来看一个直观的例子方法计算复杂度新增类别成本传统分类器O(n)需重新训练孪生网络O(n²)无需调整原型网络O(n)即时适应# 伪代码展示原型网络分类逻辑 def classify(query, prototypes): distances [euclidean_distance(query, p) for p in prototypes] return softmax(-distances) # 距离越近概率越高2. 准备小样本数据集实战中我们使用Omniglot——包含50种文字系统的1623个字符每个字符仅有20个手写样本。这个数据集完美模拟了现实中的小样本场景from torchvision.datasets import Omniglot from torchvision.transforms import Compose, Resize, ToTensor transform Compose([Resize(28), ToTensor()]) dataset Omniglot(root./data, downloadTrue, transformtransform)关键步骤是构建episode——元学习特有的数据组织方式。每个episode包含N-way随机选择N个类别如5类K-shot每类取K个样本作为支持集如每类5张图Q-query每类取Q个查询样本用于计算损失如每类15张图def sample_episode(dataset, n_way5, k_shot5, q_query15): classes random.sample(dataset._character_images.keys(), n_way) support, query [], [] for cls in classes: imgs random.sample(dataset._character_images[cls], k_shot q_query) support.extend(imgs[:k_shot]) # 支持集样本 query.extend(imgs[k_shot:]) # 查询集样本 return torch.stack(support), torch.stack(query)提示实际应用中你可能需要自定义数据集类来封装这些逻辑确保数据加载的高效性3. 构建嵌入网络原型网络的核心是一个将原始输入映射到度量空间的嵌入函数fφ。这个网络不需要特殊设计一个简单的CNN就能胜任import torch.nn as nn class EmbeddingNet(nn.Module): def __init__(self): super().__init__() self.net nn.Sequential( nn.Conv2d(1, 64, 3), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Flatten() ) def forward(self, x): return self.net(x)为什么这个结构有效通过连续的卷积和池化操作网络逐步学习到低级特征边缘、纹理中级特征部件、形状高级语义特征整体结构关键细节最后一层不使用全连接保持空间结构不变这对后续的距离计算至关重要。4. 实现原型计算与分类现在来到算法的核心部分——原型计算和分类过程。这个阶段需要精确实现三个数学操作原型计算$c_k \frac{1}{|S_k|} \sum_{x_i \in S_k} f_\phi(x_i)$距离度量通常使用欧氏距离$d(z, c_k) ||z - c_k||^2$概率分布$p(yk|z) \frac{\exp(-d(z, c_k))}{\sum_{k} \exp(-d(z, c_{k}))}$PyTorch实现如下def compute_prototypes(support, support_labels, n_way): 计算每个类别的原型向量 prototypes [] for k in range(n_way): # 选出当前类别的所有支持样本 mask (support_labels k) class_embeddings support[mask] prototypes.append(class_embeddings.mean(dim0)) return torch.stack(prototypes) def prototypical_loss(query, query_labels, prototypes, n_way): 计算原型网络的损失 distances torch.cdist(query, prototypes) # 计算查询样本与所有原型的距离 log_p -distances log_p_y log_p.gather(1, query_labels.unsqueeze(1)) loss -log_p_y.mean() acc (log_p.argmax(dim1) query_labels).float().mean() return loss, acc注意距离计算有多种选择欧氏距离最常用但余弦相似度在某些场景可能表现更好5. 训练策略与技巧原型网络的训练有其特殊性需要特别注意以下要点训练循环设计每个episode随机采样新的类别组合计算当前episode的原型评估查询集上的损失反向传播更新嵌入网络参数def train_epoch(model, optimizer, dataset, n_way5, k_shot5, q_query15): model.train() total_loss, total_acc 0, 0 for _ in range(100): # 每个epoch 100个episode support, query sample_episode(dataset, n_way, k_shot, q_query) support_emb model(support) query_emb model(query) prototypes compute_prototypes(support_emb, support_labels, n_way) loss, acc prototypical_loss(query_emb, query_labels, prototypes, n_way) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() total_acc acc.item() return total_loss / 100, total_acc / 100提升性能的实用技巧学习率调度使用ReduceLROnPlateau根据验证损失调整学习率嵌入维度最后一层特征维度通常设为64-256之间数据增强对小样本尤为重要可尝试随机旋转(-15°, 15°)轻微透视变换弹性变形# 增强后的transform from torchvision.transforms import RandomAffine, ElasticTransform train_transform Compose([ RandomAffine(degrees15, translate(0.1,0.1), scale(0.9,1.1)), ElasticTransform(alpha20.0), Resize(28), ToTensor() ])在实际项目中我发现两个常见陷阱值得注意原型漂移当支持样本过少时原型可能无法准确代表类别。解决方法是在测试时使用全部支持样本计算原型。维度灾难嵌入空间维度太高会导致距离度量失效。通过实验找到适合你数据集的维度很关键。