AI小样本学习:从元学习到基础模型时代的Few-Shot实战

AI小样本学习:从元学习到基础模型时代的Few-Shot实战
引言标注数据贵、长尾类别多是几乎所有AI落地项目的通病。医疗影像里一个罕见病种可能只有几十张片子工业质检里新型缺陷出现时往往只有个位数样本客服意图分类每周都在加新类目。传统监督学习在这些场景下要么过拟合要么干脆学不动。小样本学习Few-Shot Learning要解决的正是这个问题每个类别只有1到5个标注样本时模型仍能给出可用的精度。这篇文章梳理从元学习到基础模型两条技术路线的核心思路并给出一个可以直接跑起来的Few-Shot分类实战方案。小样本学习难在哪深度模型的参数量动辄上百万而梯度下降需要足够多的样本来约束解空间。5个样本对应几百万参数解空间几乎不受约束模型把训练样本死记硬背就能拿到满分但一换样本就崩这就是典型的过拟合。更本质的问题是监督学习的归纳偏置几乎全部来自数据本身数据少意味着偏置弱。小样本学习的所有方法本质上都是在想办法把额外知识注入学习过程——要么来自其他任务元学习要么来自预训练迁移与基础模型要么来自人为设计的结构度量空间、记忆模块。理解这一点比记住某个具体算法更重要。元学习学会如何学习元学习Meta-Learning的经典设定是episodic training训练时不直接学一个分类器而是学习如何在N-way K-shot的小任务上快速适应。训练集被组织成成千上万个模拟小任务模型在这些任务上学会快速学习的能力测试时面对全新类别的小任务就能举一反三。主流方法可以分成三大家族。基于度量的方法代表是Matching Network和Prototypical Network原型网络。思路很直白学一个embedding函数让同类样本在特征空间里聚拢分类时直接比较查询样本与各类原型支持集特征均值的距离。原型网络用欧氏距离的softmax做分类简单、稳定、易实现是工程首选。基于优化的方法代表是MAML。它不学度量而是学一组好的初始化参数使得这组参数在新任务上只需一两步梯度更新就能收敛。MAML需要计算二阶梯度训练成本高但理论上可以套到任何基于梯度下降的模型上包括强化学习和回归。基于记忆和模型的方法用外部记忆模块或专门设计的更新器如Meta-Learner LSTM来存储和调用跨任务知识。思想漂亮但工程上用得最少。基础模型时代Few-Shot换了一种活法GPT-3之后小样本学习出现了一条完全不同的路线不训练直接Prompt。大规模预训练语言模型在预训练时见过海量任务形态把任务描述和几个示例写进上下文In-Context Learning模型就能现学现卖。这条路线的意义在于把Few-Shot从训练一个模型变成了调用一个模型。视觉领域同样有CLIP这样的预训练模型把类别名称写成文本Prompt用文本编码器生成原型图像编码器做匹配零样本就能分类再给几个样本微调一下Prompt向量CoOp的做法效果还能再涨一截。需要清醒认识的是In-Context Learning的天花板受限于预训练数据分布。领域偏移大比如专业医疗术语、工业缺陷图像时纯Prompt往往不如一个针对性训练的小模型。两条路线不是替代关系而是互补。实战原型网络搭建Few-Shot分类器下面用PyTorch实现一个最小可用的原型网络。核心逻辑不到50行用一个CNN把图像编码到特征空间计算各类原型按欧氏距离分类。import torch import torch.nn as nn import torch.nn.functional as F class Encoder(nn.Module): 4层CNN把28x28图像编码到64维特征 def __init__(self): super().__init__() def block(cin, cout): return nn.Sequential( nn.Conv2d(cin, cout, 3, padding1), nn.BatchNorm2d(cout), nn.ReLU(), nn.MaxPool2d(2)) self.net nn.Sequential( block(1, 64), block(64, 64), block(64, 64), block(64, 64)) def forward(self, x): return self.net(x).view(x.size(0), -1) # [B, 64] def prototypical_loss(encoder, support, query, n_way, k_shot): support: [n_way*k_shot, C,H,W], query: [n_way*q, C,H,W] s_feat encoder(support).view(n_way, k_shot, -1) prototypes s_feat.mean(dim1) # [n_way, D] q_feat encoder(query) # [n_way*q, D] dists torch.cdist(q_feat, prototypes) ** 2 log_p F.log_softmax(-dists, dim1) labels torch.arange(n_way).repeat_interleave(query.size(0) // n_way) labels labels.to(query.device) loss F.nll_loss(log_p, labels) acc (log_p.argmax(1) labels).float().mean() return loss, acc # 训练循环episodic每个episode采样n_way个类、每类