Skip to content

元学习

本页介绍元学习(Meta-Learning):让模型学会”如何学习”——在大量小任务上训练出一种快速适应新任务的能力。它是迁移学习的进阶分支,也是少样本学习(Few-Shot Learning,即每类只有少量样本就能学习的设定)的核心方法,曾被视为通向通用人工智能的关键路径之一。

元学习就像培养一个会学习的学生,而不是只会刷一种题的刷题机器:

  • 传统监督学习 = 填鸭式训练:给 10 万张猫狗图,硬塞出一个猫狗分类器。换到识别鸟类就基本不会,要从头再喂 10 万张。
  • 元学习 = 训练”学习本能”:给它 1000 种不同的分类任务(每种任务只有 5 张样本),让它反复练习”快速上手新任务”。学完后,遇到第 1001 种新任务时,只需看 5 张样本就能分类——它学会的是 适应过程本身。

最经典的元学习场景叫 K-shot N-way 分类:给 K 张/类 × N 类的样本(如 5-shot 3-way = 每类 5 张共 3 类 = 15 张),看这几张就能给新样本分类。人类天生具备这种能力——看两张陌生水果照片就能认出第三张,而传统 CNN 做不到。

元学习有三大流派:

  • 基于优化的元学习(MAML):寻找一组”敏感”的初始参数——只需一两步梯度下降就能适应任意新任务。代表算法 MAML(Model-Agnostic Meta-Learning,模型不可知元学习)。
  • 基于度量的元学习:学一个特征空间,让同类样本靠近、异类样本远离;新样本与各类原型比距离即可。代表算法 Prototypical Net(原型网络)、Matching Net(匹配网络)。
  • 基于模型的元学习:用 RNN 或记忆增强网络(Memory-Augmented Network,即带外部存储矩阵的网络)显式存储任务经验,新任务时读取记忆做预测。代表算法 Memory-Augmented Neural Network。

直觉上理解:MAML 找的不是”某个任务的最优解”,而是”距离所有任务最优解都最近的一个出发点”——从这点出发,不管去哪个任务,一两步就到。

MAML 的数学核心是二阶优化——对”学习过程”本身求导。

设模型参数为 θ\theta,一个新任务 TiT_i 上的”适应”过程是:用任务 TiT_i 的支持集(support set,少样本)做几步梯度下降:

θi′=θ−α⋅∇θLTi(θ)\theta_i' = \theta - \alpha \cdot \nabla_\theta L_{T_i}(\theta)

这是一阶更新——θ\theta 经过一两步梯度下降变成 θi′\theta_i',就是适应后的参数。α\alpha 是内循环学习率(通常 0.01),LTiL_{T_i} 是模型在任务 TiT_i 上的损失函数。

MAML 要优化的 元目标 是:适应后的参数 θi′\theta_i' 在该任务查询集(query set,用于评估的样本)上的损失尽可能小:

min⁡θ∑iLTi(θi′)=∑iLTi ⁣(θ−α⋅∇θLTi(θ))\min_\theta \sum_i L_{T_i}(\theta_i') = \sum_i L_{T_i}\!\left(\theta - \alpha \cdot \nabla_\theta L_{T_i}(\theta)\right)

注意这个目标里 θ\theta 既出现在外层,也出现在内层梯度中——对 θ\theta 求导要穿过梯度运算,因此涉及二阶导(Hessian 矩阵,即梯度的梯度)。具体展开:

∇θLTi(θi′)=LTi′(θi′)⋅(I−α⋅∇θ2LTi(θ))\nabla_\theta L_{T_i}(\theta_i') = L'_{T_i}(\theta_i') \cdot \left(I - \alpha \cdot \nabla^2_\theta L_{T_i}(\theta)\right)

实际工程中常用一阶近似(First-Order MAML, FOMAML)省略 Hessian 项,效果接近但训练快得多。

OpenAI 提出的 Reptile 算法进一步简化了 MAML:它在每个任务上做多步 SGD(而非 1 步),然后用”初始参数和适应后参数的差值”作为更新方向:

θ←θ+β⋅(θi′−θ)\theta \leftarrow \theta + \beta \cdot (\theta_i' - \theta)

Reptile 的直觉是:在多步 SGD 后,参数会朝着该任务的最优区域移动;把所有任务的”移动方向”平均起来,就找到了一个接近所有任务最优的出发点。这比 FOMAML 更简单(不需要计算内循环梯度),且效果相当。

ANIL(Almost No Inner Loop) 是 MAML 的工程优化:内循环中只微调最后一层分类头,backbone(特征提取器)保持不变。因为 MAML 学到的元知识大部分集中在特征提取器里,最后一层分类头才是每个任务真正需要适应的部分。ANIL 训练速度快一个数量级,精度损失极小。

原型网络的思想极简:每个类别的”原型” = 该类所有支持集样本特征的平均向量:

cn=1K∑x∈Snfθ(x)c_n = \frac{1}{K} \sum_{x \in S_n} f_\theta(x)

其中 f_theta 是特征提取器(通常是一个 4 层或 6 层 CNN),S_n 是第 n 类的 K 个支持样本,c_n 就是第 n 类的”中心点”。对新样本 x 做分类时,计算它与各类原型的距离(通常用平方欧氏距离或余弦相似度),softmax 归一化即得预测概率:

p(y=n∣x)=exp⁡(−dist(fθ(x),cn))∑mexp⁡(−dist(fθ(x),cm))p(y = n \mid x) = \frac{\exp(-\text{dist}(f_\theta(x), c_n))}{\sum_m \exp(-\text{dist}(f_\theta(x), c_m))}

训练时在大量 N-way K-shot 任务上反复优化交叉熵损失,特征提取器逐渐学会”把同类聚到一起”的通用表示。

距离度量的选择:原型网络论文发现平方欧氏距离(squared Euclidean distance)效果最好,余弦相似度次之。这是因为欧氏距离在高维空间中对方向的敏感度低于对绝对位置的敏感度,更符合”同一类样本在特征空间中形成一个簇”的直觉。余弦相似度则忽略向量长度,在某些跨域场景中更鲁棒。

Matching Network(匹配网络)是少样本学习的奠基工作之一。它用注意力机制做分类:

p(y∣x)=∑ia(x,xi)⋅yip(y \mid x) = \sum_i a(x, x_i) \cdot y_i

其中 a(x, x_i) 是查询样本 x 与第 i 个支持样本 x_i 之间的注意力权重(用 softmax 归一化的余弦相似度)。与原型网络不同的是,匹配网络对每个支持样本单独做注意力,而不是取类平均——这让它在 1-shot 场景下表现更好。

  • 元训练(Meta-Training):从大量任务中(每个任务含支持集 + 查询集)学习元参数(MAML 的初始 theta 或原型网络的特征提取器)。
  • 元测试(Meta-Testing):在全新任务上,用支持集做少量适应(MAML 跑几步 SGD;原型网络算几个原型向量),在查询集上评估。

任务的构造方式通常是 episodic training(情景训练)——每个训练 step 模拟一次少样本任务:从训练集采样 N 个类、每类 K 张做支持集,再采样若干张做查询集。这和传统监督学习的”按 batch 喂数据”截然不同——每个 batch 是一个完整的 mini-task。

一个关键原则是 类别不重叠:元训练用到的类别(如动物类别 1-200)和元测试用到的类别(如动物类别 201-300)必须完全不同。否则模型只是”记住了”特定类别的特征,而不是学会了”适应新类别”的能力。

PyTorch:Prototypical Network 的核心前向逻辑

Section titled “PyTorch:Prototypical Network 的核心前向逻辑”
import torch
import torch.nn as nn
class ProtoNet(nn.Module):
def __init__(self, feat_dim=64):
super().__init__()
self.encoder = nn.Sequential( # 简化版特征提取器
nn.Flatten(), nn.Linear(28*28, 256), nn.ReLU(), nn.Linear(256, feat_dim))
def forward(self, support_x, support_y, query_x):
# support_x: [N_way * K_shot, C, H, W], support_y: [N_way * K_shot]
z_s = self.encoder(support_x) # 支持集特征
z_q = self.encoder(query_x) # 查询集特征
# 计算每个类的原型:同类样本特征取平均
prototypes = []
for c in support_y.unique():
prototypes.append(z_s[support_y == c].mean(0)) # 每类原型向量
prototypes = torch.stack(prototypes) # [N_way, feat_dim]
# 查询样本到各类原型的负距离 = 相似度 logits
dist = torch.cdist(z_q, prototypes) # 欧氏距离矩阵
return -dist # 距离越小 logit 越大

PyTorch:MAML 的核心训练循环(简化版)

Section titled “PyTorch:MAML 的核心训练循环(简化版)”
import torch
import torch.nn as nn
import torch.nn.functional as F
import higher # pip install higher,用于高阶微分
def maml_train_step(model, tasks, meta_lr=1e-3, inner_lr=0.01, inner_steps=1):
"""
MAML 单步元训练。
tasks: 一个 batch 的任务,每个任务含 (support_x, support_y, query_x, query_y)
"""
meta_optimizer = torch.optim.Adam(model.parameters(), lr=meta_lr)
meta_loss = 0.0
for support_x, support_y, query_x, query_y in tasks:
# 用 higher 创建可微分的内循环副本
with higher.innerloop_ctx(model, torch.optim.SGD(
model.parameters(), lr=inner_lr)) as (fmodel, diffopt):
# 内循环:在支持集上做 K 步梯度下降
for _ in range(inner_steps):
support_logits = fmodel(support_x)
support_loss = F.cross_entropy(support_logits, support_y)
diffopt.step(support_loss)
# 外循环:在查询集上计算元损失
query_logits = fmodel(query_x)
query_loss = F.cross_entropy(query_logits, query_y)
meta_loss += query_loss
meta_loss = meta_loss / len(tasks)
meta_optimizer.zero_grad()
meta_loss.backward() # 自动计算二阶导(穿过 diffopt.step)
meta_optimizer.step() # 更新元参数 theta
return meta_loss.item()
  • N-way K-shot 的选择:训练时通常用比测试更高的 N(如训练 20-way、测试 5-way),这样训练出的模型在测试时任务更简单,表现更好。这叫”训练难度大于测试难度”策略。
  • episodic batch size:每次元更新采样 4-16 个任务(而非 1 个),降低元梯度的方差。但太大显存压力大(MAML 要对每个任务存计算图)。
  • 类别采样均衡:确保每个 episodic task 中各类样本数严格相等,否则原型计算会偏向样本多的类。
  • 内循环学习率(alpha):通常 0.01(SGD),不宜太大(否则适应过拟合)。FOMAML/ANIL 可适当增大。
  • 外循环学习率(meta_lr):通常 1e-3(Adam),配合 warmup(前几个 epoch 线性升温)防止元训练初期发散。
  • 元训练总 epoch:通常 50000-60000 个 episodic step,配合元验证做 early stopping。
  • 梯度裁剪:MAML 的二阶梯度容易爆炸,建议裁剪到 max_norm=1.0。
  • 任务 dropout:随机跳过部分内循环步,增加适应性鲁棒性。
  • label smoothing(标签平滑):在查询集损失中加 label smoothing,让原型网络的 softmax 不过度自信,泛化更好。
  • 先用大规模预训练:现代最佳实践是先用自监督(如 SimCLR、MoCo)或大规模监督(如 ImageNet-21k)预训练一个强 backbone(如 ResNet-18/34 或 ViT-S),再在其上做元学习——比从头做元学习效果好 10-20%。详见自监督学习与迁移学习。
  • 冻结 vs 微调 backbone:用预训练 backbone 时,只做元学习最后一层(ANIL 模式),效率最高且效果接近全量 MAML。
  • 原型网络、Matching Network 训练简单(无内循环)、推理快、对小数据鲁棒,少样本分类的首选 baseline 往往是它们而非 MAML。
  • 原型网络的推理只是几个矩阵乘法(算原型 + 算距离),无需梯度下降,比 MAML 快几个数量级。

大语言模型(LLM)的兴起深刻改变了元学习的研究范式:

  • In-Context Learning(上下文学习)= 隐式元学习:GPT-3 等大语言模型只需几个示例写在 prompt 里就能完成新任务——不更新任何参数,纯靠注意力机制”理解”示例并泛化。这本质上是一种隐式的元学习:模型在预训练时见过海量不同任务,学到了”根据上下文快速适应”的能力。详见语言模型演进。
  • In-Context Learning vs 传统元学习:传统元学习(MAML/ProtoNet)通过梯度更新或度量计算来适应,需要显式的”适应”步骤;而 in-context learning 把适应过程融入了单次前向传播。两者的共同点是:都在大量任务上训练,学会的是”适应能力”而非”某个任务的知识”。

测试时训练与测试时适应(TTT / TTA)

Section titled “测试时训练与测试时适应(TTT / TTA)”

Test-Time Training(TTT) 和 Test-Time Adaptation(TTA) 是元学习思想的新应用:模型在推理时(而非训练时)利用测试样本自身的信息做快速适应。例如,对一张测试图片做自监督旋转预测,用其梯度更新部分参数,再做正式分类。这不需要额外标注,相当于”推理时做一次元学习的内循环”。代表性工作如 TENT、MEMO 等已在 2024 年广泛用于域适应(domain adaptation,即从训练域迁移到分布不同的测试域)。

2024-2025 年的研究趋势是让元学习从”小模型 + 小数据”转向”基础模型 + 高效适配”:

  • MetaICL:在大量任务上做 in-context learning 的元训练,让 LLM 的 few-shot 能力进一步提升。
  • LoRA 元学习:为大量任务学习一组好的 LoRA(Low-Rank Adaptation,低秩适配)初始化,新任务只需几个样本就能快速适配出高质量的 LoRA 权重。
  • Hypernetwork 元学习:用一个 hypernetwork(超网络,即生成其他网络权重的网络)根据任务描述直接生成适配参数,比梯度适应快得多。

随着 CLIP、LLaVA 等多模态模型的普及,少样本学习扩展到多模态场景:

  • CLIP-based few-shot:用 CLIP 的图像-文本联合嵌入空间做少样本分类,通常超过传统 ProtoNet。CLIP 预训练的视觉编码器已经在 4 亿对图文上学到了极通用的表示,少量样本就能”激活”新类别。
  • Few-shot 视频/3D 识别:将元学习扩展到视频动作识别(每个动作类别只有几段视频样本)和 3D 点云分类(每类几个物体扫描),2025 年已有不少工作将这些任务降维到 CLIP 的特征空间中处理。

跨域少样本学习(Cross-Domain Few-Shot, CDFS)

Section titled “跨域少样本学习(Cross-Domain Few-Shot, CDFS)”

传统少样本学习假设训练和测试的图像来自同一域(如都是自然图像),但真实场景中测试域可能完全不同(如从自然图像到医学影像、卫星图、草图的跨域)。2024-2025 的研究热点是 跨域少样本学习:训练域是 ImageNet 的自然图像,测试域是医学影像或遥感图像——核心挑战是弥合域差距。代表性方法包括域不变特征学习和基于 CLIP 的跨域迁移。

元学习与强化学习结合(Meta-RL)让智能体能快速学会新任务。2025 年 Meta-RL 在机器人操作(如新物体抓取)和游戏 AI(如星际争霸新策略适应)上取得进展。详见强化学习里程碑。

  • 药物分子属性预测:每种新药的分子数据极少(合成一批才几十个),但药物家族成百上千。元学习让模型在大量药物家族上学到”快速适应新家族”的能力,新药只需少量样本即可预测活性——MIT 的 Jenkins 团队用 MAML 在抗生素发现上取得突破。
  • 医疗罕见病识别:罕见病影像数据天然稀少(每种罕见病全球可能只有几百例),元学习把常见病的视觉特征迁移到罕见病上,几例样本就能启动诊断模型。
  • 工业小批量质检:每种新产品的缺陷样本极少(产线刚上线),元学习让缺陷检测模型快速适应新产品,减少停线等待数据积累的时间。
  • 机器人快速学习新任务:机器人面对新物体、新环境的抓取任务,用元学习训练的策略只需几次尝试就能适应,比强化学习从头训练快几个数量级。详见强化学习里程碑。
  • 大模型 in-context learning 的前身:GPT-3 的 few-shot 能力让元学习思路以”in-context 示例”的形式大规模落地——prompt 中给几个示例,模型不更新参数就能完成新任务,这是元学习思想在大模型时代的新形态。
类库语言说明
Learn2LearnPythonPyTorch 生态的元学习开源库,集成 MAML/FOMAML/ProtoNet/ANIL 等
TorchmetaPythonPyTorch 元学习辅助库,提供 episodic 数据加载器与少样本基准
higherPythonFacebook 开源的 PyTorch 高阶微分库,MAML 二阶优化的得力工具
meta-datasetPythonGoogle 的少样本学习基准数据集,覆盖 ImageNet/CIFAR/CLEVR 等多域
Omniglot / Mini-ImageNet数据少样本学习的经典基准数据集,几乎所有元学习论文的标配评测
Meta-Dataset / BSCD-FSL数据跨域少样本学习基准,训练域为 ImageNet,测试域为医学/卫星/艺术图像
术语英文解释
元学习Meta-Learning学习如何学习,在多任务上训练出快速适应新任务的能力
少样本学习Few-Shot Learning (FSL)每个类别只有少量(如 1-5 张)样本就能完成学习的任务设定
K-shot N-wayK-shot N-wayN 个类别、每类 K 个样本的少样本分类基准设定
支持集Support Set少样本任务中给模型看的少量标注样本,用于适应
查询集Query Set少样本任务中用于评估适应后模型表现的样本
MAMLModel-Agnostic Meta-Learning寻找敏感初始参数,一步梯度即可适应新任务的优化类元学习算法
原型网络Prototypical Network每类样本特征取平均作为原型,按距离分类的度量类元学习算法
内循环Inner Loop元学习中任务级适应过程(如 MAML 的几步本地 SGD)
外循环Outer Loop元学习中跨任务的元参数更新过程
情景训练Episodic Training把每个训练 step 组织成支持集 + 查询集的小任务采样方式
二阶优化Second-Order Optimization对梯度本身再求导(涉及 Hessian 矩阵),MAML 的核心机制
域适应Domain Adaptation模型从一个数据分布迁移到另一个不同分布的过程
上下文学习In-Context LearningLLM 通过 prompt 中的示例完成新任务,不更新参数,隐式元学习
  • MAML 开山作:Finn et al., “Model-Agnostic Meta-Learning for Fast Adaptation of Deep Networks”, ICML 2017. 元学习里程碑论文,提出与模型无关的二阶优化框架,引用量过万。
  • Prototypical Network:Snell et al., “Prototypical Networks for Few-shot Learning”, NeurIPS 2017. 原型网络,最简洁优雅的度量类元学习方法,少样本分类的强 baseline。
  • Matching Network:Vinyals et al., “Matching Networks for One Shot Learning”, NeurIPS 2016. 首次提出 episodic training 与 N-way K-shot 设定,奠基少样本学习范式。
  • First-Order MAML / Reptile:Nichol et al., “On First-Order Meta-Learning Algorithms”, arXiv 2018. 论证一阶近似(FOMAML/Reptile)效果接近完整 MAML,大幅降低计算开销。
  • ANIL:Raghu et al., “Rapid Learning or Feature Reuse? Towards Understanding the Effectiveness of MAML”, ICLR 2020. 揭示 MAML 的核心在特征复用而非快速学习,提出 ANIL。
  • 元学习综述:Hospedales et al., “Meta-Learning in Neural Networks: A Survey”, TPAMI 2022. 全面梳理三大流派、应用与开放问题,入门首选综述。