元学习实战:如何用2-way 5-shot示例解决小样本分类难题

1次阅读
没有评论

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

image.webp

背景痛点

在机器学习领域,数据匮乏是一个普遍存在的挑战。特别是在医疗影像诊断、工业缺陷检测等场景中,获取大量标注样本往往成本高昂甚至不可行。传统深度学习方法在这些小样本场景中表现不佳,主要原因包括:

元学习实战:如何用 2 -way 5-shot 示例解决小样本分类难题

  • 模型参数过多,容易在小样本上过拟合
  • 缺乏从少量样本中快速学习新概念的能力
  • 难以捕捉类别间的细粒度差异

技术方案

元学习 (Meta-Learning) 为解决小样本学习问题提供了新思路。在众多元学习方法中,我们重点分析三种典型方法:

  1. MAML(Model-Agnostic Meta-Learning):通过优化模型初始参数,使其能快速适应新任务
  2. 匹配网络(Matching Networks):使用注意力机制计算样本间相似度
  3. 原型网络(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°)
  • 颜色抖动(轻微调整亮度 / 对比度)

训练策略

  1. 使用 Adam 优化器,初始学习率 3e-4
  2. 每 1000 个 episode 降低学习率(乘以 0.5)
  3. 梯度累积:当 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

类别增量学习

当新增类别时,为避免灾难性遗忘,可采取:

  1. 保留少量旧类别样本作为 replay buffer
  2. 在新任务训练时混合旧类别样本
  3. 使用弹性权重固化 (EWC) 正则化

避坑指南

样本不均衡

  • 对样本少的类别进行过采样
  • 在距离计算时引入类别权重
  • 使用 focal loss 调整类别重要性

距离度量选择

  • 欧式距离:适用于特征空间各向同性
  • 余弦相似度:对特征幅度不敏感
  • 实践中可尝试可学习的距离度量

Loss 震荡

可能原因及解决方案:

  1. 学习率过高 → 降低学习率或使用 warmup
  2. 样本噪声 → 检查数据质量
  3. 批次多样性不足 → 增加 n_way 或 q_query

总结与展望

2-way 5-shot 的原型网络为小样本分类提供了简洁有效的解决方案。未来值得探索的方向包括:

  • 如何将领域适应技术与元学习结合,解决跨域 few-shot 问题
  • 探索自监督预训练对元学习的促进作用
  • 设计更高效的原型更新机制,适应在线学习场景
正文完
 0
评论(没有评论)