PyTorch实现MAML:Omniglot 5-way 1-shot小样本分类实战

发布时间:2026/9/20 14:36:52
PyTorch实现MAML:Omniglot 5-way 1-shot小样本分类实战
说句实话MAML 这套东西我前前后后看了三版论文公式推了又推总觉得懂了一到自己写训练循环就卡壳。后来痛定思痛决定不跟公式死磕直接用 PyTorch 在 Omniglot 上把 5-way 1-shot 的小样本分类器撸出来。代码跑通那一刻再回头去看 meta-update 那堆符号脑子里自动就有计算图了。这篇文章就是我当时逐步实现的完整记录从 Omniglot 数据怎么加载、任务怎么采样到 MAML 的内循环和外循环怎么用代码表达最后附上可以直接跑的完整代码。适合那些已经对元学习有基本概念、但还没真正手写过 MAML 的人也适合想拿 Omniglot 做 baseline 的科研新手。1. 整体设计拆解为什么选 Omniglot 当实验场1.1 小样本学习到底难在哪普通分类任务里模型见过某个类别的几百上千张图片测试时要求正确识别这个类别。小样本学习完全不是这个玩法训练阶段给你 5 个新类别每类只有 1 张或 5 张图称为 support set然后用这些极少样本去判断接下来一批 query 图各属于哪一类。人类能轻松完成这种事但传统监督学习在这种设定下基本报废因为模型根本没有足够的统计量去拟合一个新的分类边界。于是元学习Meta-Learning的思路不再是学会识别某个类别而是学会快速识别新类别。这句话看起来很像废话但本质上改变了优化目标训练过程的目标函数不是某个任务上的 loss而是模型在经历少数几步梯度更新后的任务表现。MAML 就是这种思路的代表作之一它的核心诉求非常直白找到一个优秀的初始参数让模型遇到任何新任务时只需要几步梯度下降就能快速适应。1.2 为什么 Omniglot 是验证 MAML 的最佳沙盒Omniglot 被形象地称为手写字母界的 MNIST。整个数据集包含 50 个字母体系alphabet比如拉丁字母、韩文、梵文等每个体系下有若干字符类别总共约 1623 个字符类。每个字符类只有 20 张手写样本图像是 105x105 的灰度图。这个规模在深度学习里小得感人但正因为小跑 MAML 这种带二阶梯度计算的算法时才不会把 GPU 显存直接塞爆。更关键的是Omniglot 有标准的训练/评估划分backgroundTrue 对应训练集30 个字母体系964 类backgroundFalse 对应评估集20 个字母体系659 类。这种划分保证了元学习器在训练时完全接触不到测试时的字符类别测试时面对的是全新类别。所以拿它做 MAML 的实验场地最合适不过。1.3 本文代码结构的组织方式为了避免代码和啰嗦的理论纠缠在一起我按下面几个层次组织数据层封装 Omniglot 数据提供 N-way K-shot 任务采样模型层定义一个小巧但有效的 4 层卷积网络元学习层实现 MAML 的内循环更新与外循环元更新评估层在新类别任务上验证模型快速适应能力这种分层的好处是当你以后想换成 Mini-ImageNet 或者 CIFAR-FS只需要替换数据层和模型层元学习层几乎不用动。2. 白话版 MAML 原理先忘掉二阶导这回事2.1 元学习本质上是在学一个好的初始点MAML 不学分类器本身的参数它学的是初始参数这个起点。想象你手里有一个调好的哑铃重量任何新手拿到这个重量稍微练几天就能达到不错的水平。MAML 的目标就是找到这个最优初始哑铃重量使得任何新任务在它的起点上只需要少量梯度更新就能快速收敛。论文里那张图虽然抽象但理解成寻找一个对所有任务都敏感的参数平原确实贴切。如果初始点选得好朝任何方向的梯度下降都能快速降低任务损失如果初始点选得差有些任务怎么走都走不到低损失区域。2.2 内循环在单个任务上快速适应MAML 训练时每轮采样一批任务每个任务就是一次完整的 N-way K-shot 分类问题。对于其中一个任务模型从当前初始参数 θ 出发用这个任务的 support set 计算损失做几步梯度下降。这里梯度下降的步数和学习率是人为设置的超参数一般步数很少常见配置是 5 步。经过这几步更新后得到的参数 θ就是模型对该任务快速适应后的结果。注意内循环是逐任务独立进行的。每个任务有自己的 support set所以每个任务都会得到一份独立的 θ。多个任务之间互不干扰最后在外循环统一汇总梯度。2.3 外循环让初始点在所有任务上都好适应得到每个任务的 θ 后我们把 θ 放到这个任务的 query set 上计算损失。如果初始参数 θ 选得好那么经过内循环快速适应得到的 θ 在 query set 上应该表现不错。所以外循环的目标就是最小化所有任务在 θ 上的 query loss 之和。关键点来了外循环求导时要优化的是初始参数 θ而不是 θ。这意味着计算图必须从 θ 一路连到每个 θ 再到 query loss对 θ 求梯度时梯度路径会穿过内循环的梯度更新过程于是自然引入了二阶导数。这就是 MAML 公式里那个二次梯度项的来源。2.4 二阶导到底可不可怕很多教程喜欢把二阶导渲染得很玄学实际在 PyTorch 里二阶导不需要你手动去算 Hessian 矩阵。你只需要在内循环计算梯度时开启 create_graphTruePyTorch 的自动求导引擎会把这个梯度算子本身也当作计算图的一部分记录下来。外循环再对初始参数求梯度时梯度自动就包含了二阶信息框架替你把脏活累活干完了。如果显存紧张还可以用一阶近似 FOMAML。简单说就是在外循环计算梯度时假装内循环的更新过程不参与求导只保留 θ 到 query loss 这一段梯度。这样计算量几乎减半在某些任务上效果损失不大。代码实现上只需要把内循环的 create_graph 关掉即可。3. 数据准备Omniglot 与 N-way K-shot 任务采样3.1 Omniglot 数据集的结构细节用 torchvision 加载 Omniglot 时得到的是一个列表每个样本是一对PIL 图字符标签。这里有个容易踩坑的地方torchvision 的 Omniglot 返回的字符标签是一个表示字符类别的整数覆盖全部 964 或 659 个类而不是某次任务内部的临时标签。所以我们可以在整个数据集合上做类别采样。Omniglot 原始图像是 105x105但 MAML 相关实验通常会把图像缩放到 28x28一方面与 MNIST 保持一致另一方面大幅减少网络计算量。训练阶段我额外加了随机 90 度倍数的旋转增强这是因为原文使用了旋转增强来扩充类别数每类从 20 张变成 80 张更利于元学习器学到旋转不变的特征表达。3.2 数据加载与预处理的几个选择我封装了一个 FewShotOmniglot 类内部直接复用 torchvision 自带的 Omniglot 下载与解析逻辑然后用 _flat_character_images 把图片和标签拆成两个 list。transform 部分训练集用旋转增强评估集只用缩放和归一化避免评估时引入随机性。class FewShotOmniglot(Dataset): def __init__(self, root./data, backgroundTrue, transformNone): self.dataset Omniglot(rootroot, backgroundbackground, downloadTrue) self.transform transform self.images [img for img, _ in self.dataset._flat_character_images] self.labels [label for _, label in self.dataset._flat_character_images] def __len__(self): return len(self.images) def __getitem__(self, idx): img self.images[idx] label self.labels[idx] if self.transform is not None: img self.transform(img) return img, label这里有一点要提醒如果你在国内网络环境下载失败别反复重试去官网把 omniglot-py 压缩包手动下载下来解压后放到代码指定的 data/omniglot-py 目录下再把 download 参数设为 False 即可。这条路最稳。3.3 任务采样构造 N-way K-shot 的关键代码一个任务的具体含义是随机选 N 个类别每个类别取 K 张图作为 support set再取 Q 张图作为 query set。support set 用来做内循环快速适应query set 用来计算外循环的元目标。我在 FewShotOmniglot 里写了一个 sample_task 方法每次调用生成一组 support_x、support_y、query_x、query_ydef sample_task(self, n_way5, k_shot1, q_query15, devicecpu): labels np.array(self.labels) all_classes np.unique(labels) chosen np.random.choice(all_classes, sizen_way, replaceFalse) support_x, support_y [], [] query_x, query_y [], [] for new_label, cls in enumerate(chosen): idx np.where(labels cls)[0] np.random.shuffle(idx) for i in range(k_shot): img, _ self[idx[i]] support_x.append(img) support_y.append(new_label) for i in range(k_shot, k_shot q_query): img, _ self[idx[i]] query_x.append(img) query_y.append(new_label) return (torch.stack(support_x).to(device), torch.tensor(support_y, dtypetorch.long).to(device), torch.stack(query_x).to(device), torch.tensor(query_y, dtypetorch.long).to(device))注意这里我把原始字符类别重新映射成 0 到 N-1 的临时标签否则 CrossEntropyLoss 会认为你有几百个输出维度。这个映射完全没问题因为每个任务内部的类别集是随机采样的临时标签和原始标签之间没有固定对应关系模型每轮任务都在处理 5 分类的新问题。3.4 任务采样的两个小细节第一个细节是每个类别的 support 和 query 样本必须来自同类别的不同样本不能重复。我的实现里先对类别内索引做整体 shuffle然后按顺序切分这样就保证 support 和 query 没有交集。第二个细节是 query 数量。原论文一般每类取 15 张 query这样每个任务有 5×1575 个 query 样本足够稳定地估计任务适应后的表现。如果 query 太少外循环梯度噪声会变大训练容易震荡。4. 模型定义与 MAML 训练循环逐行拆解4.1 模型选择小巧的 4 层卷积网络Omniglot 图像很小不需要很深很大的网络。我参考 MAML 原文使用了一个 4 层卷积网络每层是 64 个 3×3 卷积核卷积后接归一化和 ReLU再做 2×2 max pooling。最后的分类层输入维度是特征图的通道数输出维度是 N-way 的类别数。这里有个我踩过的坑第一版我用了 BatchNorm结果训练曲线震荡得很厉害准确率长期在 20% 上下。原因在于 BatchNorm 依赖 batch 内的统计量而 MAML 每个任务的数据量非常少support 只有 5 张图batch 统计量极其不稳定而且内循环临时参数更新后 running mean 的处理也很麻烦。换成 GroupNorm 后问题迎刃而解它与 batch 大小解耦在小样本场景下明显更稳。class CNN4(nn.Module): def __init__(self, num_classes5): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 64, 3, padding1), nn.GroupNorm(8, 64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3, padding1), nn.GroupNorm(8, 64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3, padding1), nn.GroupNorm(8, 64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3, padding1), nn.GroupNorm(8, 64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Linear(64, num_classes) def forward(self, x): x self.features(x) x F.adaptive_avg_pool2d(x, 1).flatten(1) return self.classifier(x)加粗提醒functional_call需要传入模型全部参数所以你必须保证模型里的 BatchNorm 在推理时不会额外维护 buffer否则参数字典不完整会报错。我用 GroupNorm 之后就没有这个烦恼。4.2 核心机制functional_call 与参数替换MAML 内循环的每一步都要基于更新后的临时参数重新前向计算。传统做法是把模型复制一份在副本上更新参数但这样会让计算图断开无法传播到初始参数。另一个曲线做法是手动保存原始参数内循环时修改 module 的 weight 和 bias外循环再恢复操作起来很麻烦且容易出错。我推荐直接用torch.func.functional_call。这个函数允许你传入一个参数字典和输入按字典里的参数计算前向结果而模型对象本身的参数完全不动。这样我们就可以定义一个由初始参数字典逐步演化出临时参数字典的纯函数式循环。from collections import OrderedDict from torch.func import functional_call # 获取模型初始参数快照 params OrderedDict(model.named_parameters()) # 用 params 执行前向计算 logits functional_call(model, params, (support_x,))4.3 内循环更新函数内循环本质上是针对单个任务做几步普通梯度下降只不过参数载体是 OrderedDict且必须开启 create_graphTrue这样才能保留二阶梯度路径。def inner_update(model, params, support_x, support_y, inner_lr, steps, create_graphTrue): for _ in range(steps): logits functional_call(model, params, (support_x,)) loss F.cross_entropy(logits, support_y) grads torch.autograd.grad(loss, params.values(), create_graphcreate_graph) params OrderedDict( (name, param - inner_lr * grad) for (name, param), grad in zip(params.items(), grads) ) return params这里的param - inner_lr * grad会生成一个新的 tensor它记录了从原始参数出发的完整计算路径。到第五步结束params 里的每个 tensor 都带着一串长长的梯度更新历史。这就是二阶导数的来源。4.4 外循环元更新逻辑每个任务独立做内循环更新后我们都得到一份任务特定的 fast_params。接下来用这份 fast_params 在 query set 上算 loss把多个任务的 loss 平均再对最开始的模型参数求梯度。def meta_train_step(model, train_set, meta_batch_size, n_way, k_shot, q_query, inner_lr, inner_steps, meta_lr, device): task_losses [] task_accs [] for _ in range(meta_batch_size): support_x, support_y, query_x, query_y train_set.sample_task( n_way, k_shot, q_query, device ) params OrderedDict(model.named_parameters()) fast_params inner_update(model, params, support_x, support_y, inner_lr, inner_steps, create_graphTrue) logits_q functional_call(model, fast_params, (query_x,)) loss_q F.cross_entropy(logits_q, query_y) task_losses.append(loss_q) pred_q logits_q.argmax(dim1) task_accs.append((pred_q query_y).float().mean().item()) meta_loss torch.stack(task_losses).mean() meta_acc np.mean(task_accs) meta_grads torch.autograd.grad(meta_loss, model.parameters()) with torch.no_grad(): for p, g in zip(model.parameters(), meta_grads): p.sub_(meta_lr * g) return meta_loss.item(), meta_acc有个细节值得强调这里更新参数用的是手动梯度下降p.sub_(meta_lr * g)不是optimizer.step()。原因是我们并不想引入动量或权重衰减这些额外因素而且 meta loss 的梯度是手动从计算图里取出来的用 optimizer 反而容易搞混哪些参数该更新。如果你想用 Adam 来做外循环优化可以把 meta_grads 喂给 optimizer 的step之外的逻辑但为了复现 MAML 原文SGD 手动更新最直接。4.5 评估函数评估时我们要模拟真实使用场景模型拿到一个新任务先在 support set 上快速更新几步然后在 query set 上看准确率。此时不需要再构建二阶计算图所以 create_graphFalse省显存也省时间。def evaluate(model, eval_set, n_way5, k_shot1, q_query15, inner_lr0.4, inner_steps5, num_tasks200, devicecpu): model.eval() accs [] with torch.enable_grad(): for _ in range(num_tasks): support_x, support_y, query_x, query_y eval_set.sample_task( n_way, k_shot, q_query, device ) params OrderedDict(model.named_parameters()) fast_params inner_update(model, params, support_x, support_y, inner_lr, inner_steps, create_graphFalse) logits_q functional_call(model, fast_params, (query_x,)) pred_q logits_q.argmax(dim1) accs.append((pred_q query_y).float().mean().item()) return np.mean(accs)注意这里torch.enable_grad()是有意保留的因为 inner_update 里即使 create_graphFalse仍然需要调用torch.autograd.grad来计算临时参数的更新梯度纯torch.no_grad()下会直接报错。4.6 完整可运行代码整理把前面的代码块拼起来加上超参数和主循环就能直接跑起来了。import os import random import numpy as np from collections import OrderedDict import torch import torch.nn as nn import torch.nn.functional as F import torchvision.transforms as transforms from torch.utils.data import Dataset from torchvision.datasets import Omniglot from torch.func import functional_call # ---------- 超参数 ---------- META_LR 1e-3 INNER_LR 0.4 INNER_STEPS 5 META_BATCH_SIZE 16 N_WAY 5 K_SHOT 1 Q_QUERY 15 EVAL_INTERVAL 50 META_ITERS 3000 SEED 42 # ---------- 数据封装 ---------- class FewShotOmniglot(Dataset): def __init__(self, root./data, backgroundTrue, transformNone): self.dataset Omniglot(rootroot, backgroundbackground, downloadTrue) self.transform transform self.images [img for img, _ in self.dataset._flat_character_images] self.labels [label for _, label in self.dataset._flat_character_images] def __len__(self): return len(self.images) def __getitem__(self, idx): img self.images[idx] label self.labels[idx] if self.transform is not None: img self.transform(img) return img, label def sample_task(self, n_way5, k_shot1, q_query15, devicecpu): labels np.array(self.labels) all_classes np.unique(labels) chosen np.random.choice(all_classes, sizen_way, replaceFalse) support_x, support_y [], [] query_x, query_y [], [] for new_label, cls in enumerate(chosen): idx np.where(labels cls)[0] np.random.shuffle(idx) for i in range(k_shot): img, _ self[idx[i]] support_x.append(img) support_y.append(new_label) for i in range(k_shot, k_shot q_query): img, _ self[idx[i]] query_x.append(img) query_y.append(new_label) return (torch.stack(support_x).to(device), torch.tensor(support_y, dtypetorch.long).to(device), torch.stack(query_x).to(device), torch.tensor(query_y, dtypetorch.long).to(device)) # ---------- 模型 ---------- class CNN4(nn.Module): def __init__(self, num_classes5): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 64, 3, padding1), nn.GroupNorm(8, 64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3, padding1), nn.GroupNorm(8, 64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3, padding1), nn.GroupNorm(8, 64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3, padding1), nn.GroupNorm(8, 64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Linear(64, num_classes) def forward(self, x): x self.features(x) x F.adaptive_avg_pool2d(x, 1).flatten(1) return self.classifier(x) # ---------- MAML 核心 ---------- def inner_update(model, params, support_x, support_y, inner_lr, steps, create_graphTrue): for _ in range(steps): logits functional_call(model, params, (support_x,)) loss F.cross_entropy(logits, support_y) grads torch.autograd.grad(loss, params.values(), create_graphcreate_graph) params OrderedDict( (name, param - inner_lr * grad) for (name, param), grad in zip(params.items(), grads) ) return params def meta_train_step(model, train_set, meta_batch_size, n_way, k_shot, q_query, inner_lr, inner_steps, meta_lr, device): task_losses [] task_accs [] for _ in range(meta_batch_size): support_x, support_y, query_x, query_y train_set.sample_task( n_way, k_shot, q_query, device ) params OrderedDict(model.named_parameters()) fast_params inner_update(model, params, support_x, support_y, inner_lr, inner_steps, create_graphTrue) logits_q functional_call(model, fast_params, (query_x,)) loss_q F.cross_entropy(logits_q, query_y) task_losses.append(loss_q) pred_q logits_q.argmax(dim1) task_accs.append((pred_q query_y).float().mean().item()) meta_loss torch.stack(task_losses).mean() meta_acc np.mean(task_accs) meta_grads torch.autograd.grad(meta_loss, model.parameters()) with torch.no_grad(): for p, g in zip(model.parameters(), meta_grads): p.sub_(meta_lr * g) return meta_loss.item(), meta_acc # ---------- 评估 ---------- def evaluate(model, eval_set, n_way5, k_shot1, q_query15, inner_lr0.4, inner_steps5, num_tasks200, devicecpu): model.eval() accs [] with torch.enable_grad(): for _ in range(num_tasks): support_x, support_y, query_x, query_y eval_set.sample_task( n_way, k_shot, q_query, device ) params OrderedDict(model.named_parameters()) fast_params inner_update(model, params, support_x, support_y, inner_lr, inner_steps, create_graphFalse) logits_q functional_call(model, fast_params, (query_x,)) pred_q logits_q.argmax(dim1) accs.append((pred_q query_y).float().mean().item()) return np.mean(accs) # ---------- 主循环 ---------- if __name__ __main__: torch.manual_seed(SEED) np.random.seed(SEED) random.seed(SEED) device torch.device(cuda if torch.cuda.is_available() else cpu) train_transform transforms.Compose([ transforms.RandomRotation([0, 90, 180, 270]), transforms.Resize((28, 28)), transforms.ToTensor(), ]) test_transform transforms.Compose([ transforms.Resize((28, 28)), transforms.ToTensor(), ]) train_set FewShotOmniglot(root./data, backgroundTrue, transformtrain_transform) eval_set FewShotOmniglot(root./data, backgroundFalse, transformtest_transform) model CNN4(num_classesN_WAY).to(device) for it in range(META_ITERS): loss, acc meta_train_step( model, train_set, META_BATCH_SIZE, N_WAY, K_SHOT, Q_QUERY, INNER_LR, INNER_STEPS, META_LR, device ) if (it 1) % EVAL_INTERVAL 0: eval_acc evaluate(model, eval_set, N_WAY, K_SHOT, Q_QUERY, INNER_LR, INNER_STEPS, num_tasks100, devicedevice) print(fIter {it1:5d} | train_loss {loss:.4f} | train_acc {acc:.3f} | eval_acc {eval_acc:.3f})这段代码我实测在 CPU 上跑 3000 步大约需要几十分钟GPU 上会快很多。如果你的机器配置一般可以把 META_BATCH_SIZE 调成 8或者把 META_ITERS 降到 1000先确认整个流程能通。5. 训练策略与关键参数调节5.1 内循环学习率为什么要设成 0.4普通监督学习的学习率一般是 0.001 到 0.01但 MAML 内循环的学习率却大得离谱常用 0.4。这不是手滑而是设计使然内循环的目标不是精细拟合而是快速适应。模型需要在 5 步之内从初始参数快速移动到该任务的最优参数附近步子必须迈大。我试过把内循环学习率改成 0.01结果无论怎么调外循环学习率模型都无法在 query 上取得满意效果因为 5 步梯度下降对于 0.01 的步长来说几乎等于原地踏步。你可以在代码里试试不同值体验一下模型学不会快速适应是什么样的。5.2 内循环步数与计算成本内循环步数 INNER_STEPS 直接决定计算图的深度。步数越多二阶梯度路径越长显存和耗时都线性增加。5 步是 Omniglot 场景下的经典选择再往上收益很小还容易造成过拟合。如果你显存紧张可以先从 2 步跑通再逐步增加。5.3 外循环学习率与任务 batch 大小的配合外循环学习率 META_LR 用的是 0.001。这个值不能照搬普通 SGD 的经验因为 meta loss 的梯度是跨多个任务平均后的结果噪声相对可控但梯度本身是二阶梯度量级和普通梯度不一样。我建议在 0.0005 到 0.003 之间搜索。META_BATCH_SIZE 我用了 16。原论文用的是 32理论上 batch 越大meta 梯度越稳定但显存开销也越大。在 Omniglot 这种简单任务上16 已经能获得不错效果。有个直观判断标准如果训练 loss 曲线波动非常剧烈可以考虑调大 meta batch size如果训练孤零零地稳定但准确率上不去更应该怀疑其他参数。5.4 训练曲线的判断经验很多第一次跑 MAML 的人会被训练 loss 的形态吓到因为它不像普通分类任务那样单调下降。MAML 的训练 loss 是每个 meta step 里多个任务的 query loss 平均在前期波动很大这是正常的。我习惯同时盯三个指标训练 loss、训练 acc、评估 acc。训练 acc 反映当前 meta batch 内任务快速适应后的平均能力评估 acc 反映是否真的学到了跨任务泛化的初始参数。如果训练 acc 在涨而评估 acc 不动大概率是过拟合到了训练时的类别分布如果两边都不动先检查代码里参数更新有没有生效。5.5 数据增强对最终准确率的贡献Omniglot 原实验里的 4 倍旋转增强非常有效。每个字符旋转 0、90、180、270 度后等于把类别数扩大了 4 倍。这让元学习器见过的任务种类更多样训练出的初始参数通用性更强。我在代码里用 RandomRotation 实现了同样的效果实测评估准确率能提高 5 到 10 个百分点。加了旋转增强后训练和评估时的 transform 要区分开。如果评估时也做随机旋转会因为随机性导致准确率波动不好复现。评估只用 Resize 和 ToTensor。6. 常见问题与排查技巧实录6.1 训练 loss 不下降怎么办优先检查三点内循环更新得到 fast_params 后有没有真正用到 query 上计算 loss。很多人写错成用初始模型直接算 query loss那样 MAML 就退化成了普通多任务学习。外循环更新的是不是模型原始参数。确认 meta_grads 是torch.autograd.grad(meta_loss, model.parameters())而不是对 fast_params 求梯度。学习率是否合适。内循环学习率太小会导致快速适应失效外循环学习率太大容易震荡。如果 code 没问题但 loss 还是不动可以打印某个任务的 support loss 和 query loss 对比support loss 应该在内循环过程中逐步下降query loss 在训练初期可能会偏高这都正常。6.2 二阶梯度显存爆炸META_BATCH_SIZE 太大、INNER_STEPS 太多都会让计算图非常庞大。如果你的显存在 8G 以下建议按 16 batch、5 steps 先跑如果还是爆把 batch 降到 8。另一个思路是改用 FOMAML把 inner_update 的 create_graph 设成 False二阶导数直接消失显存压力大幅下降代价是最终准确率可能会有小幅损失。6.3 评估时 functional_call 报参数不匹配错误信息通常会提示 expected 某参数但没找到。绝大多数情况是模型里有 BatchNorm 的 running buffer而你的参数字典只包括了 named_parameters。解决办法有两种把 BatchNorm 换成 GroupNorm我在代码中就是这个方案或者把模型里的 buffer 也一起传入 functional_call。第二种方案还需要维护 buffer 在任务间的独立性比较麻烦不建议新手尝试。6.4 测试准确率远低于训练准确率先检查是不是把 backgroundTrue 的训练集拿去测试了。Omniglot 元学习场景必须保证测试类别在训练阶段完全没见过。详细排查方向整理成速查表现象可能原因处理建议训练 acc 高评估 acc 低训练/评估类别分布不一致确认 background 参数设置正确训练与评估 acc 都很低内循环 lr 过小尝试 0.1 到 0.5 的范围loss 早期下降后期震荡外循环 lr 偏大降低 meta_lr 到 0.0005评估 acc 波动大评估任务数太少扩大到 500 个任务取平均加旋转增强后效果反而变差transform 同时用在了测试集区分训练和测试 transform6.5 一个隐藏很深的坑手动更新参数后 optimizer 失效如果你把模型交给 optimizer 管理又在 meta_train_step 里手动用p.sub_更新了参数之后调用 optimizer.step() 会导致梯度信息混乱。我的建议是干脆不用 optimizerMAML 外循环的手动更新就是最原汁原味的做法。后面你如果要做外循环 Adam只需要把优化器维护的 param group 和 meta_grads 对齐而不是盲目的 optimizer.step()。7. 写在最后的一点感想跑通这个项目之后我对 MAML 的整体感觉确实完全变了。原来看公式时觉得晦涩难懂的 inner loop、outer loop在实际代码里就是两个简单的 for 循环加一个 OrderedDict 参数列表。第一步内循环算出临时参数第二步外循环把这个临时参数在 query 上的 loss 沿计算图反传回初始参数。真就这么简单。如果让我给刚接触元学习的人一个建议我会说不要先去弄懂 Hessian 矩阵和二阶近似的所有数学细节先把 5-way 1-shot 在 Omniglot 上跑起来看到准确率一点点从 30% 涨到 90% 以上然后再回去看原论文。你会发现公式里每一个符号都对应代码中的某一行。等什么时候你想进一步优化二阶计算或者想把自己的研究任务套上元学习框架再回过来研究 FOMAML、最近邻分类器、proto-net 这些变体就水到渠成了。我后来在这个基础上做过一次 Quick 实验把模型换成 ResNet12再把内循环替换成两步对比学习目标效果提升虽然有限但这个 MAML 骨架代码几乎没怎么动。这也是我希望你跑通后能继续用起来的原因——小样本学习这条路起点越简单后续换新组件越方便。