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, padding=1), 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(dim=1) # [n_way, D] q_feat = encoder(query) # [n_way*q, D] dists = torch.cdist(q_feat, prototypes) ** 2 log_p = F.log_softmax(-dists, dim=1) 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个类、每类