共计 2892 个字符,预计需要花费 8 分钟才能阅读完成。
背景痛点
在机器学习领域,数据匮乏是一个普遍存在的挑战。特别是在医疗影像诊断、工业缺陷检测等场景中,获取大量标注样本往往成本高昂甚至不可行。传统深度学习方法在这些小样本场景中表现不佳,主要原因包括:

- 模型参数过多,容易在小样本上过拟合
- 缺乏从少量样本中快速学习新概念的能力
- 难以捕捉类别间的细粒度差异
技术方案
元学习 (Meta-Learning) 为解决小样本学习问题提供了新思路。在众多元学习方法中,我们重点分析三种典型方法:
- MAML(Model-Agnostic Meta-Learning):通过优化模型初始参数,使其能快速适应新任务
- 匹配网络(Matching Networks):使用注意力机制计算样本间相似度
- 原型网络(Prototypical Networks):为每个类别计算原型表示,基于距离进行分类
其中,原型网络因其简单高效,特别适合 2 -way 5-shot 任务。2-way 表示每个 episode 包含 2 个类别,5-shot 表示每个类别提供 5 个支持样本。数学表示为:
- 支持集 $S={(x_i,y_i)}_{i=1}^{N\times K}$,其中 N 为类别数,K 为每类样本数
- 类别 c 的原型 $p_c=\frac{1}{|S_c|}\sum_{(x_i,y_i)\in S_c}f_\phi(x_i)$
- 查询样本 x 的类别概率 $p(y=c|x)=\frac{\exp(-d(f_\phi(x),p_c))}{\sum_{c’}\exp(-d(f_\phi(x),p_{c’}))}$
PyTorch 实现
Episode 生成器
class EpisodeSampler:
def __init__(self, dataset, n_way, k_shot, q_query):
self.dataset = dataset
self.n_way = n_way
self.k_shot = k_shot
self.q_query = q_query
def __iter__(self):
# 随机选择 n_way 个类别
classes = np.random.choice(len(self.dataset.classes),
self.n_way,
replace=False
)
# 为每个类别采样 k_shot+q_query 个样本
support, query = [], []
for c in classes:
indices = np.random.choice(len(self.dataset.class_to_idx[c]),
self.k_shot + self.q_query,
replace=False
)
support.extend(indices[:self.k_shot])
query.extend(indices[self.k_shot:])
yield torch.stack(support), torch.stack(query)
特征提取器
class CNNBackbone(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(nn.Conv2d(3, 64, 3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(64, 64, 3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(64, 64, 3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(64, 64, 3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(2)
)
def forward(self, x):
return self.net(x).view(x.size(0), -1)
原型网络
class PrototypicalNetwork(nn.Module):
def __init__(self, backbone):
super().__init__()
self.backbone = backbone
def forward(self, support_x, support_y, query_x):
# 提取特征
support_z = self.backbone(support_x)
query_z = self.backbone(query_x)
# 计算原型
prototypes = []
for c in torch.unique(support_y):
mask = (support_y == c)
prototype = support_z[mask].mean(0)
prototypes.append(prototype)
prototypes = torch.stack(prototypes)
# 计算距离
dists = torch.cdist(query_z, prototypes)
logits = -dists
return logits
实现细节
数据预处理
- 标准化:使用 ImageNet 均值和标准差进行归一化
- 数据增强:对小样本任务尤为重要,推荐使用:
- 随机水平翻转
- 小角度旋转(±15°)
- 颜色抖动(轻微调整亮度 / 对比度)
训练策略
- 使用 Adam 优化器,初始学习率 3e-4
- 每 1000 个 episode 降低学习率(乘以 0.5)
- 梯度累积:当 GPU 内存不足时,可累积多个 episode 的梯度再更新
可视化原型
def plot_prototypes(prototypes, labels):
# t-SNE 降维
tsne = TSNE(n_components=2)
points = tsne.fit_transform(prototypes)
# 绘制散点图
plt.figure(figsize=(10,8))
for i, (x,y) in enumerate(points):
plt.scatter(x, y, label=labels[i])
plt.legend()
plt.show()
生产考量
计算资源
| Backbone | GPU 内存(MB) | 单 episode 耗时(ms) |
|---|---|---|
| Conv4 | 1200 | 15 |
| ResNet18 | 3800 | 45 |
类别增量学习
当新增类别时,为避免灾难性遗忘,可采取:
- 保留少量旧类别样本作为 replay buffer
- 在新任务训练时混合旧类别样本
- 使用弹性权重固化 (EWC) 正则化
避坑指南
样本不均衡
- 对样本少的类别进行过采样
- 在距离计算时引入类别权重
- 使用 focal loss 调整类别重要性
距离度量选择
- 欧式距离:适用于特征空间各向同性
- 余弦相似度:对特征幅度不敏感
- 实践中可尝试可学习的距离度量
Loss 震荡
可能原因及解决方案:
- 学习率过高 → 降低学习率或使用 warmup
- 样本噪声 → 检查数据质量
- 批次多样性不足 → 增加 n_way 或 q_query
总结与展望
2-way 5-shot 的原型网络为小样本分类提供了简洁有效的解决方案。未来值得探索的方向包括:
- 如何将领域适应技术与元学习结合,解决跨域 few-shot 问题
- 探索自监督预训练对元学习的促进作用
- 设计更高效的原型更新机制,适应在线学习场景
正文完
发表至: 未分类
四天前
