元学习实战:2-way 5-shot示例解析与新手避坑指南

1次阅读
没有评论

共计 2377 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

背景痛点

在实际业务场景中,数据标注成本往往是制约模型应用的瓶颈。以医疗影像诊断为例,罕见病病例可能仅有数张标注样本,而传统深度学习需要数千张样本才能达到可用效果。小样本学习(Few-Shot Learning)的核心价值在于:让模型通过少量样本快速适应新任务。

传统迁移学习的局限性体现在:

  • 预训练特征偏向基础类别(如 ImageNet 的 1000 类)
  • 微调阶段容易过拟合极少量样本
  • 无法实现跨类别的知识迁移

技术对比

主流元学习算法可分为三类:

  1. 基于优化的方法 (如 MAML)
  2. 核心思想:学习一个易于快速适应的模型初始化参数
  3. 优点:适应性强,适合复杂任务
  4. 缺点:二阶导数计算成本高

  5. 基于度量的方法 (如 Prototypical Networks)

  6. 核心思想:在特征空间构建类别原型(prototype)
  7. 优点:计算高效,适合类别定义明确的任务
  8. 缺点:依赖特征空间的质量

  9. 基于记忆的方法 (如 MANN)

  10. 核心思想:通过外部存储模块保存历史经验
  11. 优点:适合长周期依赖任务
  12. 缺点:记忆模块难以稳定训练

对于新手入门,建议从 Prototypical Networks 开始,因其实现简单且效果稳定。

核心实现

任务定义

2-way 5-shot 表示:

  • 每个 episode 包含 2 个类别(way)
  • 每类提供 5 个支持样本(shot)
  • 查询集(query set)包含每类 15 个待分类样本

数学表达为:

$$
\begin{aligned}
S &= {(x_i,y_i)}{i=1}^{N \times K} \quad \text{(支持集)} \
Q &= {(x_j,y_j)}

\end{aligned}
$$}^{N \times M} \quad \text{(查询集)

其中 $N=2$ 为类别数,$K=5$ 为 shot 数,$M=15$ 为查询样本数。

Episode 训练机制

  1. 任务采样 :从数据集中随机抽取 N 个类别
  2. 样本划分 :每类随机选取 K 个样本作为支持集,其余作为查询集
  3. 原型计算 :对每个类别 $c$,计算其原型向量:
    $$
    p_c = \frac{1}{|S_c|} \sum_{(x_i,y_i) \in S_c} f_\theta(x_i)
    $$
  4. 距离度量 :使用欧式距离计算查询样本与各原型的距离
    $$
    d(f_\theta(x), p_c) = |f_\theta(x) – p_c|_2
    $$
  5. 损失计算 :最小化查询样本的负对数似然

代码示例

import torch
import torch.nn as nn
from torch.utils.data import Dataset

class PrototypicalNetwork(nn.Module):
    def __init__(self, encoder):
        super().__init__()
        self.encoder = encoder  # 共享的特征提取器

    def forward(self, support_x, support_y, query_x):
        """
        参数说明:
        support_x: [n_way * k_shot, C, H, W]
        support_y: [n_way * k_shot]
        query_x: [n_way * n_query, C, H, W]
        """
        # 特征提取
        support_features = self.encoder(support_x)  # [N*K, D]
        query_features = self.encoder(query_x)      # [N*Q, D]

        # 计算类别原型
        unique_labels = torch.unique(support_y)
        prototypes = []
        for label in unique_labels:
            # 获取当前类别的所有支持样本特征
            mask = support_y == label
            class_features = support_features[mask]
            # 计算原型 (均值向量)
            prototypes.append(class_features.mean(dim=0))
        prototypes = torch.stack(prototypes)  # [N, D]

        # 计算欧式距离
        dists = torch.cdist(query_features.unsqueeze(0),  
            prototypes.unsqueeze(0)
        ).squeeze(0)  # [N*Q, N]

        # 计算概率分布 (softmax over distances)
        logits = -dists
        return logits

实验分析

基准测试

在 miniImageNet(64 训练类 /16 验证类 /20 测试类)上的结果:

方法 5-way 1-shot 5-way 5-shot
Matching Networks 43.56% 55.31%
Prototypical Nets 49.42% 68.20%
MAML 48.70% 63.11%

硬件配置:NVIDIA V100 GPU,batch size=4,Adam 优化器

Shot 数量影响

元学习实战:2-way 5-shot 示例解析与新手避坑指南

关键观察:

  • 当 shot 数 <3 时,模型性能剧烈波动
  • shot 数 >10 后收益递减
  • 5-shot 是性价比最优的选择

避坑指南

类别不平衡处理

当支持集中各类样本数不一致时:

  • 对少数类进行特征增强(旋转 / 翻转)
  • 采用 Focal Loss 代替交叉熵
  • 原型计算时添加可学习的权重系数

特征空间坍塌

症状:所有样本特征聚集到同一区域

解决方案:

  1. 在训练目标中加入特征分散损失
    $$
    \mathcal{L}{div} = -\frac{1}{N}\sum
    $$}^N \log \frac{e^{|p_i|_2}}{\sum_j e^{|p_j|_2}
  2. 使用解耦的特征提取器(如:主干网络 + 投影头)
  3. 引入对比学习预训练阶段

延伸思考

  1. 跨域适应 :如何将在 miniImageNet 上训练的模型迁移到医疗影像领域?
  2. 增量学习 :当新类别不断加入时,如何避免灾难性遗忘?
  3. 主动学习 :如何智能选择最有价值的样本进行标注?

通过这个 2 -way 5-shot 的示例,我们验证了元学习在小样本场景下的有效性。实际部署时建议从 ProtoNets 开始验证,再逐步尝试更复杂的算法。

正文完
 0
评论(没有评论)