共计 4422 个字符,预计需要花费 12 分钟才能阅读完成。
背景与痛点
少样本学习(Few-shot Learning)是机器学习中的一个重要研究方向,旨在让模型能够在极少量样本(通常每类只有 1 - 5 个样本)的情况下快速学习新概念。在实际应用中,我们经常会遇到数据稀缺的情况,这时候传统的深度学习模型往往表现不佳。

主要痛点包括:
- 样本效率低:传统深度学习需要大量标注数据
- 模型泛化能力差:在小样本上训练容易过拟合
- 计算成本高:每次遇到新任务都需要重新训练
- 适应性差:难以快速适应新的类别或场景
技术选型对比
解决少样本学习问题主要有以下几种技术路线:
- 元学习(Meta-learning)
- 优点:可以学习 ” 如何学习 ”,在新任务上快速适应
-
缺点:训练过程复杂,需要设计特殊的训练流程
-
迁移学习(Transfer Learning)
- 优点:利用预训练模型的特征提取能力
-
缺点:对领域差异敏感,微调需要谨慎
-
度量学习(Metric Learning)
- 优点:通过比较样本间距离进行分类
-
缺点:需要设计合适的度量函数
-
数据增强(Data Augmentation)
- 优点:简单直接,易于实现
- 缺点:生成样本质量影响模型性能
在实际应用中,往往会结合使用这些方法。下面我们将重点介绍基于元学习的 Prototypical Networks 方法。
核心实现(PyTorch)
数据加载
from torch.utils.data import Dataset
import numpy as np
import torch
class FewShotDataset(Dataset):
"""
少样本学习数据集
Args:
data: 输入数据 [n_samples, n_features]
labels: 对应标签 [n_samples]
n_way: 每次训练的类别数
k_shot: 每个类的支持样本数
q_query: 每个类的查询样本数
"""
def __init__(self, data, labels, n_way=5, k_shot=1, q_query=5):
self.data = data
self.labels = labels
self.n_way = n_way
self.k_shot = k_shot
self.q_query = q_query
# 获取所有类别及其索引
self.classes = np.unique(labels)
self.class_to_indices = {c: np.where(labels == c)[0] for c in self.classes
}
def __len__(self):
return 100 # 固定 episode 数量
def __getitem__(self, _):
# 随机选择 n_way 个类别
selected_classes = np.random.choice(self.classes, self.n_way, replace=False)
support = []
query = []
for c in selected_classes:
# 获取当前类的所有样本索引
indices = self.class_to_indices[c]
# 随机选择 k_shot + q_query 个样本
selected = np.random.choice(indices, self.k_shot + self.q_query, replace=False)
# 前 k_shot 个作为支持集,其余作为查询集
support.extend(selected[:self.k_shot])
query.extend(selected[self.k_shot:])
# 转换为 tensor
support_data = torch.stack([self.data[i] for i in support])
query_data = torch.stack([self.data[i] for i in query])
# 创建支持集和查询集的标签 (0 到 n_way-1)
support_labels = torch.LongTensor([np.where(selected_classes == self.labels[i])[0][0]
for i in support
])
query_labels = torch.LongTensor([np.where(selected_classes == self.labels[i])[0][0]
for i in query
])
return support_data, support_labels, query_data, query_labels
模型架构
import torch.nn as nn
import torch.nn.functional as F
class ProtoNet(nn.Module):
"""Prototypical Networks 实现"""
def __init__(self, encoder):
super(ProtoNet, self).__init__()
self.encoder = encoder
def forward(self, support, support_labels, query):
"""
前向传播
Args:
support: 支持集数据 [n_way * k_shot, feature_dim]
support_labels: 支持集标签 [n_way * k_shot]
query: 查询集数据 [n_way * q_query, feature_dim]
Returns:
logits: 查询样本的 logits [n_way * q_query, n_way]
"""
# 编码支持集和查询集
support_emb = self.encoder(support)
query_emb = self.encoder(query)
# 计算每个类的原型 (均值向量)
n_way = len(torch.unique(support_labels))
prototypes = torch.stack([support_emb[support_labels == c].mean(0)
for c in range(n_way)
])
# 计算查询样本与每个原型的距离 (负平方欧式距离)
dists = -torch.cdist(query_emb, prototypes, p=2) ** 2
return dists
训练流程
from torch.utils.data import DataLoader
def train_protonet(model, train_data, val_data, epochs=100, lr=0.001):
"""
训练 ProtoNet 模型
Args:
model: ProtoNet 实例
train_data: 训练数据集
val_data: 验证数据集
epochs: 训练轮数
lr: 学习率
"""
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
criterion = nn.CrossEntropyLoss()
train_loader = DataLoader(train_data, batch_size=1, shuffle=True)
val_loader = DataLoader(val_data, batch_size=1, shuffle=False)
for epoch in range(epochs):
model.train()
train_loss = 0
train_acc = 0
for support, support_labels, query, query_labels in train_loader:
optimizer.zero_grad()
# 前向传播
logits = model(support, support_labels, query)
# 计算损失和准确率
loss = criterion(logits, query_labels)
acc = (logits.argmax(1) == query_labels).float().mean()
# 反向传播
loss.backward()
optimizer.step()
train_loss += loss.item()
train_acc += acc.item()
# 验证
model.eval()
val_loss = 0
val_acc = 0
with torch.no_grad():
for support, support_labels, query, query_labels in val_loader:
logits = model(support, support_labels, query)
loss = criterion(logits, query_labels)
acc = (logits.argmax(1) == query_labels).float().mean()
val_loss += loss.item()
val_acc += acc.item()
# 打印训练和验证指标
print(f"Epoch {epoch+1}/{epochs}:")
print(f"Train Loss: {train_loss/len(train_loader):.4f}")
print(f"Train Acc: {train_acc/len(train_loader):.4f}")
print(f"Val Loss: {val_loss/len(val_loader):.4f}")
print(f"Val Acc: {val_acc/len(val_loader):.4f}")
性能优化技巧
- 特征提取器选择
- 使用在大规模数据集上预训练的特征提取器(如 ResNet)
-
根据任务需求调整特征提取器的复杂度
-
数据增强
- 对少量样本进行合理的数据增强
-
可以使用 Mixup、CutMix 等高级增强方法
-
损失函数改进
- 除了交叉熵损失,可以加入对比损失、中心损失等
-
使用标签平滑(Label Smoothing)减少过拟合
-
训练策略
- 调整学习率调度(如余弦退火)
-
使用更大的 n_way 和 k_shot 进行训练
-
模型集成
- 结合多个不同的少样本学习方法
- 对不同模型的预测结果进行集成
生产环境建议
- 部署注意事项
- 模型大小和计算效率的平衡
- 支持新类别增量更新的能力
-
在线学习的实现方式
-
常见问题解决方案
- 类别不平衡问题:使用类别加权
- 样本质量差:增加数据清洗步骤
-
领域偏移:定期更新模型
-
监控指标
- 准确率、召回率等传统指标
- 样本效率(学习曲线)
- 推理时间
动手实验
为了验证所学知识,我们设计一个简单的实验:
- 准备一个小的图像数据集(如 mini-ImageNet 的子集)
- 实现 Prototypical Networks 模型
- 训练模型并观察 5 -way 1-shot 和 5 -way 5-shot 的性能差异
- 尝试不同的特征提取器(如 CNN、ResNet 等)
- 对比有无数据增强的效果
通过这个实验,你可以直观地理解少样本学习的挑战和解决方案。
总结
少样本学习是一个实用且充满挑战的研究方向。本文介绍了 Prototypical Networks 的实现方法,并分享了性能优化和生产部署的经验。在实际应用中,我们需要根据具体场景选择合适的技术路线,并持续优化模型性能。希望这篇文章能帮助你快速上手少样本学习,构建高效能的 Agent 模型。
