ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch手写MAML:Omniglot小样本分类实战与核心代码解析

PyTorch手写MAML:Omniglot小样本分类实战与核心代码解析 我最初接触MAML的时候犯过几乎所有新手都会犯的错先找论文对着公式推了三页纸以为自己理解了结果一打开代码还是不知道第一轮的网络结构到底该写多深。后来我把顺序彻底反过来先跑通一个最简单的MAML小样本分类器再回头看原论文那些公式突然全活了过来。所以这篇不是一篇论文解读而是直接用PyTorch在Omniglot数据集上手写一个5-way小样本分类器。目标就是让你通过代码理解MAML的核心机制不靠死记硬背原理而是亲手把“内循环适应”和“外循环元更新”跑通。整个过程包含数据管线、任务采样、内循环更新、二阶梯度回传、训练评估完整脚本。这套思路适合所有想入坑元学习和小样本学习的开发者哪怕你之前只写过普通监督学习的分类网络没有任何元学习基础只要耐心跟着代码走一遍也能真正搞清楚MAML在做什么。1. 为什么选MAML和Omniglot作为你入坑小样本学习的第一站1.1 MAML的设计哲学不学结果学“快速学会”的能力MAML全称Model-Agnostic Meta-Learning翻译过来是“模型无关的元学习”。为什么说模型无关因为它的设计不依赖你用的是卷积网络、全连接网络还是Transformer只要模型可微分都能套进去。它的核心主张和普通人想象的不太一样传统小样本学习通常想让模型学会“直接分类”或者学会“一个相似度度量函数”但MAML选择了一个更底层的目标——学一组比较好的初始参数。这组初始参数的特点是当遇到一个新任务时只需要在支撑集上做几步梯度下降模型就能快速适应这个任务并在查询集上取得好效果。你不需要专门设计复杂的网络结构去匹配任务类型也不需要引入额外的度量网络。MAML本身就把“快速适应”这件事变成了训练目标这是它最优雅、也最适合作为入门算法的地方。尤其对新手来说MAML几乎没有花哨的组件核心就是两层循环内循环是针对单个任务临时更新参数外循环是以“临时更新后的表现”为监督信号去更新最初的网络参数。这个结构一旦在代码里跑通整个元学习的图景就打开了。1.2 Omniglot为什么它是比mini-ImageNet更友好的试验田Omniglot数据集可以理解为手写字符界的MNIST但它比MNIST的类别多得多。MNIST只有0到9共10类手写数字Omniglot包含50种字母表的1623类手写字符每个字符只有20个手写样本。1623个类别、每类20个样本这个数据分布天然就是为小样本学习设计的。论文和开源社区一般会把数据切分成两部分训练集background来自30种字母表约963个字符测试集evaluation来自剩余20种字母表约659个字符。注意这两个集合的字母表完全不重叠也就是说测试时看到的都是训练阶段完全没见过的字符类别这才能测出模型真正的跨类别泛化能力而不是简单死记硬背。比起mini-ImageNetOmniglot有明显的入门优势体积小。原图是105x105的灰度图缩放到28x28之后整个训练集也就十几MBCPU都能轻松训练。训练快。我用普通笔记本GPU跑5-way 5-shot任务几千次迭代就能看到准确率逼近99%同配置下跑mini-ImageNet光数据加载就能把人劝退。和MNIST长得像但更难。正因为图像是手写字符不同字母表的形状差异很大模型必须学会“结构特征”而不是“像素模板”这对元学习来说是一个很好的训练场。所以如果你想用最短的时间验证自己对MAML的理解Omniglot就是最好的第一站。别一上来就挑战mini-ImageNet那样你会在数据预处理和算力上消耗大量热情。2. 用PyTorch的执行逻辑拆解MAML的双层结构2.1 内循环在做什么一次“针对当前任务的临时适应”MAML的训练单位是“任务”task而不是一条样本。一个-task由支撑集和查询集组成比如5-way 1-shot任务就是随机选5个字符类别每类挑1张图做支撑集再挑若干张不同样本做查询集。内循环做的事情简单说就是拿着当前模型的初始参数在支撑集上进行几次普通的梯度下降。但这里和普通训练有一个关键区别——内循环更新的是参数的临时副本不是网络本身的参数。每经历一个任务模型会临时变成“适应过这个任务”的新参数然后在查询集上计算损失这个损失才是真正用来更新原始参数的信号。可以这么理解内循环是一次“摸底考试”目标是让模型试着快速掌握当前任务成绩的好坏本身不重要重要的是它能不能凭借初始参数很快适应这个任务。适应得越好说明初始参数越优秀。2.2 外循环在做什么把“适应能力”写回初始参数外循环是MAML的训练主循环。一个batch里有多个任务每个任务经过内循环之后都得到一个查询集损失。把这些损失求和再取平均对模型初始参数求梯度然后用Adam或SGD去更新初始参数。关键点在于这个损失是“经过内循环临时更新之后”的损失也就是说反向传播时信号要穿过内循环那几条梯度下降路径回到最开始的参数上。PyTorch里要做到这一点必须让内循环里每一步更新都保留计算图也就是用到torch.autograd.grad时的create_graphTrue参数。很多人第一次写MAML时会在内循环里直接更新model.parameters()然后对外层loss做backward。这样做出来的效果几乎必然不好因为你在更新已经发生了之后才去算梯度初始参数根本拿不到“适应过程”的梯度信号训练退化成了一种很别扭的监督学习。2.3 二阶梯度这条路为什么代码里绕不开MAML和普通梯度下降最大的差别就在这个“二阶梯度”上。内循环的每一层参数更新都包含初始参数的一阶梯度外循环要对查询集损失求梯度这个梯度里又嵌套了内循环的梯度于是自然就产生了二阶导数。用生活化的话说普通训练相当于跑短跑追求的是终点成绩MAML相当于“训练短跑运动员”但成绩不是看他跑一次有多快而是看他在不熟悉跑道上跑几次之后能进步多少。你要优化的是“进步的速度”而不是直接的跑步成绩。这个“进步的速度”就是二阶梯度传递的信息。在代码实现上如果不设置create_graphTruePyTorch会在内循环更新参数时丢掉临时参数的梯度图外循环拿到的损失对初始参数不可导报错或者效果极差。这一点我在后文的核心代码里会明确标出来。3. 环境准备与数据管线先把Omniglot喂进PyTorch3.1 Anaconda与PyTorch环境安装的完整命令这一节我默认你用的是Anaconda因为conda管理Python环境和CUDA依赖比裸pip方便太多。Windows、Ubuntu、macOS的流程基本一致下面是完整步骤。# 创建独立环境避免污染base环境 conda create -n maml python3.10 -y conda activate maml # CPU版本适合只想跑通代码的读者 pip install torch torchvision # GPU版本以CUDA 11.8为例 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118装完之后验证一下PyTorch能不能正常调用GPUimport torch print(torch.__version__) print(torch.cuda.is_available())如果输出True说明GPU环境OK如果是False不一定是装错了也有可能是CUDA驱动和PyTorch版本不匹配。我建议先确认nvidia-smi能正常显示显卡信息再根据驱动版本选择对应的CUDA安装命令。纯CPU跑Omniglot的小样本任务也完全可行就是训练时间会长一些。3.2 Omniglot数据集的下载、解压与组织Omniglot官方数据在GitHub上有仓库但国内直接访问下载有时很慢。我建议优先找国内GitHub加速镜像或者别人打包好的数据集文件只要能拿到images_background和images_evaluation两个目录就行。拿到数据后目录结构大概是这个样子data/ ├── images_background/ │ ├── Alphabet_of_the_Magi/ │ │ ├── character01/ │ │ │ ├── 01.png │ │ │ └── ... │ │ └── ... │ └── ... └── images_evaluation/ ├── ...我们需要把它读成“类别到图像列表”的映射结构。这里有个细节要注意每一个字符目录里的图片数量各不相同有的正好20张有的可能含子目录或非png文件所以代码里要过滤。import os import random import numpy as np import torch from PIL import Image def load_omniglot(data_dir): label_to_images {} for char_name in sorted(os.listdir(data_dir)): char_dir os.path.join(data_dir, char_name) if not os.path.isdir(char_dir): continue images [] for fname in sorted(os.listdir(char_dir)): if not fname.endswith(.png): continue path os.path.join(char_dir, fname) img Image.open(path).convert(L).resize((28, 28), Image.LANCZOS) img np.array(img, dtypenp.float32) / 255.0 images.append(torch.tensor(img).unsqueeze(0)) if images: label_to_images[char_name] images return label_to_images train_label_to_images load_omniglot(data/images_background) eval_label_to_images load_omniglot(data/images_evaluation)所有图像都归一化到0到1区间格式是[1, 28, 28]后面batch时直接stack即可。个人建议在内存够的时候直接把所有图像都load成Tensor训练时省去反复读磁盘的时间这个预处理策略对Omniglot这种小数据集非常有效。3.3 Task Sampler让每个Task的采样逻辑正确且高效Task Sampler是MAML数据管线里最关键的组件它决定了每次训练拿到的任务长什么样。一个合格的任务采样器至少要满足三个条件从所有类别里随机抽取way个类别每个类别内随机抽shot个支撑样本剩余样本里再抽query个查询样本且查询样本不能和支撑样本重复。数据泄漏的重灾区就在第3点。如果支撑集和查询集里出现了同一张图片模型不需要任何泛化能力直接记住像素就能拿高分训练出来的模型没有意义。class OmniglotTaskSampler: def __init__(self, label_to_images, way5, shot1, query15, num_tasks2000): self.label_to_images label_to_images self.all_labels list(label_to_images.keys()) self.way way self.shot shot self.query query self.num_tasks num_tasks def sample_task(self): task_labels random.sample(self.all_labels, self.way) support_x, support_y [], [] query_x, query_y [], [] for class_idx, label in enumerate(task_labels): imgs self.label_to_images[label] random.shuffle(imgs) support_imgs imgs[:self.shot] query_imgs imgs[self.shot:self.shot self.query] support_x.extend(support_imgs) support_y.extend([class_idx] * self.shot) query_x.extend(query_imgs) query_y.extend([class_idx] * self.query) return ( torch.stack(support_x), torch.tensor(support_y), torch.stack(query_x), torch.tensor(query_y) ) def __iter__(self): for _ in range(self.num_tasks): yield self.sample_task() def __len__(self): return self.num_tasks这里的way5表示5分类问题shot1表示每个类别只给一张支撑图这是小样本学习最经典的5-way 1-shot设置。query数量一般取15因为查询样本太少会让梯度不稳定太多则显存压力大。4. 手写MAML模型定义、内循环更新与元更新的核心代码4.1 网络结构设计四层卷积特征提取器MAML论文里使用的Omniglot骨干网络是一个四层卷积网络每层64个3x3卷积核激活函数用ReLU中间穿插2x2最大池化。原版网络带BatchNorm但我在自己复现时发现BN在MAML里隐藏着不少坑为了先把核心流程跑通这一版我去掉了BN只保留卷积、ReLU和池化效果依然很好。import torch import torch.nn as nn import torch.nn.functional as F from torch.func import functional_call class MAMLConvNet(nn.Module): def __init__(self, in_channels1, hidden_dim64, num_classes5): super().__init__() self.conv1 nn.Conv2d(in_channels, hidden_dim, 3, padding1) self.conv2 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1) self.conv3 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1) self.conv4 nn.Conv2d(hidden_dim, hidden_dim, 3, padding1) self.fc nn.Linear(hidden_dim * 7 * 7, num_classes) def forward(self, x, paramsNone): if params is None: params dict(self.named_parameters()) x F.relu(F.conv2d(x, params[conv1.weight], params[conv1.bias], padding1)) x F.max_pool2d(x, 2) x F.relu(F.conv2d(x, params[conv2.weight], params[conv2.bias], padding1)) x F.max_pool2d(x, 2) x F.relu(F.conv2d(x, params[conv3.weight], params[conv3.bias], padding1)) x F.relu(F.conv2d(x, params[conv4.weight], params[conv4.bias], padding1)) x x.view(x.size(0), -1) x F.linear(x, params[fc.weight], params[fc.bias]) return x这里有个尺寸推导输入28x28经过两次2x2池化后变成7x7所以全连接层的输入维度是64*7*73136。如果输入尺寸变了这里的数字也必须跟着调整新手经常在这一行踩坑。4.2 functional_call临时参数副本的正确注入方式MAML内循环需要临时更新参数但如果直接改model.parameters()里的值就会污染模型本身的状态而且让反向传播的图结构变得混乱。正确做法是维护一个参数副本前向传播时用副本代替原始参数。PyTorch 2.0之后提供了torch.func.functional_call可以很方便地把一袋子参数“注入”到模型里做前向计算而不改变模型本身的参数。from torch.func import functional_call # params是名为conv1.weight等字符串映射到参数的字典 params dict(model.named_parameters()) logits functional_call(model, params, (x,))这样做的本质是把模型当成一个纯函数参数从外部传入。每次内循环更新我们就生成一个新的参数字典然后用这个字典去算查询集损失。外循环的梯度就能沿着这一串临时更新回传到最初的参数上。4.3 inner_loop和meta_train_step核心训练逻辑内循环的实现非常直观。给定支撑集用当前模型参数做预测计算交叉熵损失然后对损失求梯度并更新参数字典。唯一的重点是create_graphTrue。def inner_loop(model, support_x, support_y, inner_lr, inner_steps): params dict(model.named_parameters()) for _ in range(inner_steps): logits functional_call(model, params, (support_x,)) loss F.cross_entropy(logits, support_y) grads torch.autograd.grad(loss, params.values(), create_graphTrue) params {name: p - inner_lr * g for (name, p), g in zip(params.items(), grads)} return params内循环结束后params已经不是原来的初始化参数而是适应过当前任务之后的临时参数。外循环要做的就是用临时参数在查询集上表现来更新原始参数。def meta_train_step(model, meta_optimizer, task_batch, inner_lr, inner_steps): meta_loss 0.0 for support_x, support_y, query_x, query_y in task_batch: adapted_params inner_loop(model, support_x, support_y, inner_lr, inner_steps) logits functional_call(model, adapted_params, (query_x,)) meta_loss meta_loss F.cross_entropy(logits, query_y) meta_loss meta_loss / len(task_batch) meta_optimizer.zero_grad() meta_loss.backward() meta_optimizer.step() return meta_loss.item()这几十行代码就是MAML的全部奥义。外层backward()在PyTorch中会自动处理二阶梯度因为内循环里保留了临时更新的计算图。5. 训练与评估完整脚本、超参数对照和收敛自查5.1 超参数为什么这样设置我自己的训练配置如下这个组合在Omniglot上稳定且不折腾超参数取值说明way5每次任务5类shot1或5支撑集样本数query15查询集样本数inner_lr0.4内循环SGD步长inner_steps5训练时内循环步数eval_inner_steps20评估时内循环步数outer_lr0.001外循环Adam学习率meta_batch_size4每次迭代处理的任务数iterations10000-20000总训练迭代数很多人第一次看到inner_lr0.4都会被吓到觉得学习率怎么会这么大。原因在于内循环只做5步更新它的目标不是精确收敛到某个局部最优而是“朝适应新任务的方向走几步”。步长太小元学习的梯度就感知不到真正的适应方向步长接近普通学习率时MAML想学的“快速适应能力”才会真正显露出来。外循环用Adam是因为外层损失面更崎岖Adam能把每个参数的学习率调平收敛比纯SGD更稳定。在验证MAML时通常不用学习率衰减因为每次任务都在变衰减反而可能让后期模型失去适应能力。5.2 训练脚本与评估脚本把上面的组件拼起来训练代码短得惊人。device torch.device(cuda if torch.cuda.is_available() else cpu) model MAMLConvNet().to(device) model.train() meta_optimizer torch.optim.Adam(model.parameters(), lr0.001) train_sampler OmniglotTaskSampler(train_label_to_images, way5, shot1, query15, num_tasks20000) eval_sampler OmniglotTaskSampler(eval_label_to_images, way5, shot1, query15, num_tasks1000) inner_lr 0.4 inner_steps 5 meta_batch_size 4 for it in range(10000): task_batch [next(iter(train_sampler)) for _ in range(meta_batch_size)] task_batch [(x.to(device), y.to(device), qx.to(device), qy.to(device)) for x, y, qx, qy in task_batch] loss meta_train_step(model, meta_optimizer, task_batch, inner_lr, inner_steps) if it % 100 0: print(fIteration {it}, meta_loss: {loss:.4f}) if it % 1000 0 and it 0: acc meta_evaluate(model, eval_sampler, inner_lr, eval_inner_steps20) print(fEvaluation acc after {it} iterations: {acc:.4f})评估函数和训练的主要区别在于内循环步数可以更大而且create_graphFalse也够用因为评估时不需要更新初始参数。def meta_evaluate(model, eval_sampler, inner_lr, eval_inner_steps20): model.eval() total_acc, total_count 0.0, 0 with torch.no_grad(): # 这里注意torch.no_grad()下autograd.grad的create_graph不能为True for support_x, support_y, query_x, query_y in eval_sampler: support_x support_x.to(device) support_y support_y.to(device) query_x query_x.to(device) query_y query_y.to(device) params dict(model.named_parameters()) for _ in range(eval_inner_steps): logits functional_call(model, params, (support_x,)) loss F.cross_entropy(logits, support_y) grads torch.autograd.grad(loss, params.values(), create_graphFalse) params {name: p - inner_lr * g for (name, p), g in zip(params.items(), grads)} query_logits functional_call(model, params, (query_x,)) preds query_logits.argmax(dim-1) total_acc (preds query_y).float().sum().item() total_count query_y.size(0) model.train() return total_acc / total_count在我自己的笔记本上5-way 1-shot任务训练到大约1万次迭代测试集准确率能到96%以上5-shot任务很快能逼近99%。这个数字低于论文的99.9%主要因为我去掉了BatchNorm和旋转增强但足以证明整套代码逻辑是对的。5.3 跑通之后务必做的三个自检代码跑完之后别急着开心请做三个自检随机初始化模型在测试task上的准确率应该接近1/way。5-way任务大约是20%。如果你的随机模型一上来就60%以上说明Task Sampler或评估逻辑里数据泄漏了。训练loss初期可能有波动这是正常的。因为MAML优化的是二阶信号计算图比普通监督学习复杂初期loss上下抖动不等于训练失败。只要整体趋势向下评估准确率上升就没问题。5-shot准确率应该明显高于1-shot。如果5-shot比1-shot还差几乎可以肯定是支撑集和查询集没分干净或者模型训练时过拟合到了某个固定任务分布上。这三个自检建议当成MAML代码的“出厂测试”每次改完代码都跑一遍能帮你拦住大部分隐藏bug。6. 我在实践中最常踩的坑和调试心得6.1 BN在MAML里是一个隐藏的雷原版MAML Omniglot实现里是有BatchNorm的但BN和MAML协同工作时问题很多。BN层在训练时会持续更新running_mean和running_mean_var这在普通监督学习里没问题但在MAML里内循环每次更新临时参数时BN的统计量也会被更新导致不同task之间相互污染。更麻烦的是如果使用functional_call注入参数BN的buffer需要额外处理。你需要把所有buffer也放进参数词典并且要决定running统计量到底用初始模型维护还是跟随临时参数一起更新。这个决策如果做错训练过程中模型的行为就会很不稳定。我的建议很务实先去掉BN跑通整个流程再回过头来研究怎么把BN正确地加回去。毕竟Omniglot上的效果差异并没有那么夸张去掉BN之后还能省掉一个维度的心智负担。6.2 二阶导带来的显存和时间开销比想象中大create_graphTrue会把整个内循环的计算图保留在内存里5步内循环的计算图深度和复杂度大约是普通backward的2到3倍。我第一次训练时直接把meta_batch_size开到8还在每个task里放30个query样本结果显存瞬间爆掉。经验做法是先把meta_batch_size设成1或2跑通整个流程看显存占用显存不够时优先减少query数量而不是减少way数内循环步数增加到5以上时显存开销上升明显实测中训练时内循环步数到5就够了如果实在显存紧张可以用最近比较流行的一阶近似MAML来替换它在内循环更新时对梯度做detach只保留一阶信息。但我建议第一次学MAML不要用一阶近似因为那会模糊掉你对二阶梯度的理解。先跑标准版本再优化。6.3 数据泄漏的细节隐蔽性很强我在调试早期版本时曾经出现过5-way 1-shot直接冲到70%准确率的情况一开始还挺高兴后来仔细看Task Sampler才发现我把整个类别的所有图片都先stack到一起然后按顺序切分sample和query。由于每类图片本身顺序固定某些类在mini-batch里出现的顺序和位置具有一定规律模型学到了很多“捷径”。把图片顺序每个task都重新shuffle之后准确率才回归正常。这个坑特别容易出现在追求效率的向量化实现里一定要记得random.shuffle必须在每个类别的图像列表上执行而且必须在切分support和query之前。如果你用了多进程DataLoader采样器的随机种子也要在每次worker启动时重置否则不同epoch拿到的任务序列可能完全一致。6.4 调参心得小结最后说说我用下来最有效的几个调参方向。内循环学习率inner_lr是MAML里最敏感的超参。我测试过0.1以下基本学不动0.4到0.5附近最优超过1.0之后训练明显震荡。因为内循环只有几步这个学习率本质上是“适应速度”不是普通意义上的收敛步长所以不用被直觉上的大学习率吓退。outer_lr要小Adam的0.001是个好起点。如果发现外层损失爆炸先降outer_lr而不是降inner_lr。数据增强维度Omniglot有个非常适合的增强方式旋转。手写字符旋转90度、180度、270度后仍然属于同一个字符类别这是Omniglot官方就认可的先验。在加载图片时对每张图做4个角度的旋转相当于把训练数据扩大了4倍任务难度也跟着降低准确率会有明显提升。但要注意训练集和测试集的字符来自不同字母表旋转增强不会造成类别混淆。Omniglot跑通之后再去涉足mini-ImageNet你会发现MAML的代码结构完全不用大改只需要换更重的backbone、更复杂的任务采样策略以及多GPU训练工具。到那个阶段你对“元学习”这三个字的理解就不再是一个需要背诵的名词而是实打实印在代码里的肌肉记忆了。
返回列表