MindSpore 元学习(Meta-Learning)实战
一、引言
深度学习的成功,大家已经非常熟悉了——海量标注数据是它的底气。但在很多现实场景里,比如医疗影像诊断、稀有物种识别、工业缺陷检测,弄到大量标注样本,又贵又难。这时候,元学习就站出来了。它也叫“学会学习”,意思就是不止学知识,而是学怎么学得更快,这正好和Few-Shot Learning(小样本学习)一拍即合。

说白了,元学习的思路是:让模型在大量相关任务上摸爬滚打,练出一种通用的“学习能力”。等你给它一个全新的任务,哪怕只给几个样本,它也能快速上手。就像5-way 1-shot分类,要从5个类别里,每个类别只给1个样本,它就得认出来。
这篇文章,我们从元学习的基本理论讲起,重点用MAML和Prototypical Networks这两个经典算法来开刀,然后带着大家用MindSpore 2.0框架,把代码跑起来。
二、元学习基础理论
2.1 元学习与传统学习的区别
传统深度学习是怎么干的?一句话概括:找一个最优参数θ*,让模型在训练集上损失最小。公式看起来就是:。但这种做法有个硬伤:一旦测试数据的分布和训练时不一样,或者测试样本太少,模型就容易“翻车”。
元学习则换了个思路。它面对的是一个任务分布T,每个任务τ都有自己的训练集(support set)和测试集(query set)。它的目标,是找到一组超强的初始参数θ,让你在新任务上,只做几步梯度更新,就能快速适应:。所以,传统学习是“学知识”,元学习是“学如何学知识”——这区别,就是本质上的。
2.2 MAML 算法原理
MAML(Model-Agnostic Meta-Learning)是2017年由Chelsea Finn团队提出的,称得上是元学习领域最闪亮的那颗星。它的核心思想,可以说是非常优雅:找到一组对任务变化特别敏感的初始参数。什么意思?就是你拿它在任意任务上,做一步或几步梯度更新,它就能表现得很棒。
MAML 的数学表述
具体怎么做呢?对于从任务分布中抽出来的一个任务τ,MAML先在内循环里做一次更新:。这里的α是内循环学习率,用的是support set上的损失。
然后,在外面再跑一个外循环,目标是跨任务优化最开始的参数θ:。β是外循环学习率,用的是query set上的损失。
这里的关键点在于,外循环的梯度得穿过内循环的梯度更新过程,也就是计算。这就是关于二阶梯度(second-order gradient)的事儿。展开来看:。二阶梯度算起来成本高,因为它包含Hessian矩阵。很多实际应用里,大家图省事,直接用一阶近似的FOMAML(First-Order MAML)——忽略二阶项,只用来更新θ。不过话说回来,完整的二阶MAML通常还是更胜一筹。
2.3 Prototypical Networks 原理
Prototypical Networks是2017年Snell等人提出的,走的是度量学习的路子,非参数化,想法很直观:为每个类别算出一个“原型(prototype)”——其实就是该类别所有support样本嵌入的平均值。然后,看query样本离哪个原型最近,就归到哪类。
具体来看:给定一个嵌入函数,对于任务τ中的类别k,它的原型就是:。
对于query样本x_q,它属于类别k的概率,通过计算距离后取softmax得到:。这里的距离函数d,通常就用欧氏距离的平方:。
Prototypical Networks的好处是显而易见的:不需要内循环梯度更新,训练效率高,代码也简洁明了。
2.4 元学习的应用场景
- Few-Shot图像分类:比如Omniglot手写字符识别、miniImageNet细粒度分类。
- 强化学习:机器人快速适应新任务,比如在不同地形上运动。
- 自然语言处理:少样本文本分类、跨领域情感分析。
- 医疗影像:只在少量标注样本下进行疾病诊断。
- 个性化推荐:冷启动场景下的快速用户建模。
三、MindSpore 实现元学习
3.1 环境准备
开干之前,先把吃饭的家伙准备好。MindSpore 2.0的配置,可以直接搬过来用:
import mindspore as ms
import mindspore.nn as nn
import mindspore.ops as ops
from mindspore import Tensor
from mindspore.dataset import vision, transforms
import numpy as np
import os
ms.set_context(device_target="GPU", mode=ms.GRAPH_MODE)
ms.set_seed(42)
3.2 数据集构建
要做Few-Shot实验,得先有个能生成N-way K-shot任务的数据集生成器。这里我们用Omniglot风格的任务采样来演示:
class FewShotTaskGenerator:
"""Few-Shot 任务采样器,从任务分布中采样 N-way K-shot 任务"""
def __init__(self, dataset, num_classes, num_samples_per_class):
"""
Args:
dataset: 全量数据集,字典形式 {class_id: [samples]}
num_classes: 每个任务采样的类别数 (N-way)
num_samples_per_class: 每个类别采样的样本数 (K-shot)
"""
self.dataset = dataset
self.num_classes = num_classes
self.num_samples_per_class = num_samples_per_class
self.all_classes = list(dataset.keys())
def generate_task(self):
"""采样一个 N-way K-shot 任务,返回 (support_x, support_y, query_x, query_y)"""
# 随机选择 N 个类别
sampled_classes = np.random.choice(self.all_classes, size=self.num_classes, replace=False)
support_images, support_labels = [], []
query_images, query_labels = [], []
for new_label, cls_id in enumerate(sampled_classes):
samples = self.dataset[cls_id]
indices = np.random.permutation(len(samples))
# 前半作为 support set,后半作为 query set
support_idx = indices[:self.num_samples_per_class]
query_idx = indices[self.num_samples_per_class:]
for idx in support_idx:
support_images.append(samples[idx])
support_labels.append(new_label)
for idx in query_idx:
query_images.append(samples[idx])
query_labels.append(new_label)
return (np.array(support_images, dtype=np.float32),
np.array(support_labels, dtype=np.int32),
np.array(query_images, dtype=np.float32),
np.array(query_labels, dtype=np.int32))
def create_synthetic_dataset(num_classes=50, samples_per_class=20, img_size=28):
"""创建合成 Few-Shot 数据集(用于演示,实际使用请替换为 Omniglot)"""
np.random.seed(42)
dataset = {}
for cls in range(num_classes):
# 每个类别有不同的统计特征,模拟真实数据的类间差异
mean = np.random.randn(1, img_size, img_size) * 0.3
std = np.random.uniform(0.3, 0.8)
samples = np.random.randn(samples_per_class, 1, img_size, img_size) * std + mean
# 归一化到 [0, 1]
samples = (samples - samples.min()) / (samples.max() - samples.min() + 1e-8)
dataset[cls] = samples
return dataset
# 构建数据集
print("构建合成 Few-Shot 数据集...")
full_dataset = create_synthetic_dataset(num_classes=50, samples_per_class=20)
task_gen = FewShotTaskGenerator(full_dataset, num_classes=5, num_samples_per_class=5)
# 验证数据采样
sup_x, sup_y, q_x, q_y = task_gen.generate_task()
print(f"Support set: {sup_x.shape}, labels: {sup_y.shape}")
print(f"Query set: {q_x.shape}, labels: {q_y.shape}")
print(f"类别分布: {np.unique(sup_y, return_counts=True)}")
3.3 MAML 算法的完整实现
现在我们来看怎么用MindSpore把MAML跑起来,而且是包含二阶梯度计算的完整版本。说实话,写MAML最头疼的地方就是那个二阶梯度,但MindSpore的自动微分能力正好能派上用场。
class SimpleCNN(nn.Cell):
"""用于 Few-Shot 分类的轻量 CNN 主干网络"""
def __init__(self, num_classes=5, in_channels=1):
super(SimpleCNN, self).__init__()
self.conv1 = nn.Conv2d(in_channels, 64, kernel_size=3, pad_mode='same', has_bias=True)
self.bn1 = nn.BatchNorm2d(64)
self.conv2 = nn.Conv2d(64, 64, kernel_size=3, pad_mode='same', has_bias=True)
self.bn2 = nn.BatchNorm2d(64)
self.conv3 = nn.Conv2d(64, 64, kernel_size=3, pad_mode='same', has_bias=True)
self.bn3 = nn.BatchNorm2d(64)
self.conv4 = nn.Conv2d(64, 64, kernel_size=3, pad_mode='same', has_bias=True)
self.bn4 = nn.BatchNorm2d(64)
self.relu = nn.ReLU()
self.max_pool = nn.MaxPool2d(kernel_size=2, stride=2)
self.flatten = nn.Flatten()
# 28x28 → 经过4次池化 → 1x1x64 = 64
self.fc = nn.Dense(64, num_classes)
self.dropout = nn.Dropout(keep_prob=0.5)
def construct(self, x):
x = self.max_pool(self.relu(self.bn1(self.conv1(x))))
x = self.max_pool(self.relu(self.bn2(self.conv2(x))))
x = self.max_pool(self.relu(self.bn3(self.conv3(x))))
x = self.max_pool(self.relu(self.bn4(self.conv4(x))))
x = self.flatten(x)
x = self.dropout(x)
return self.fc(x)
class MAML(nn.Cell):
"""MAML 元学习算法的 MindSpore 实现
实现了完整的二阶梯度计算。内循环在每个任务上进行梯度更新,外循环跨任务优化初始参数。
"""
def __init__(self, backbone, inner_lr=0.01, num_inner_steps=1):
super(MAML, self).__init__()
self.backbone = backbone
self.inner_lr = inner_lr
self.num_inner_steps = num_inner_steps
# 外循环优化器由外部管理
def inner_loop_update(self, x, y, weights):
"""执行内循环的前向传播和损失计算(用于获取适配后的参数)"""
logits = self.backbone.construct(x)
loss = nn.CrossEntropyLoss()(logits, y)
grads = ops.grad(self.backbone.construct)(x)
# 通过自动微分获取梯度
return loss, logits
def construct(self, sup_x, sup_y, q_x, q_y):
"""
MAML 的前向传播:
1. 在 support set 上执行内循环梯度更新
2. 在 query set 上计算元损失(外循环)
3. 返回元损失用于外循环优化
Args:
sup_x: support set 图像
sup_y: support set 标签
q_x: query set 图像
q_y: query set 标签
Returns:
meta_loss: 外循环损失
accuracy: query set 准确率
"""
# ---- 内循环 ----
# 在 support set 上计算损失并获取梯度
logits_sup = self.backbone(sup_x)
loss_sup = nn.CrossEntropyLoss()(logits_sup, sup_y)
# 获取模型参数的梯度
grads = ops.grad(self._inner_loss)(sup_x, sup_y)
# 手动应用梯度更新(内循环)
params = self.backbone.trainable_params()
adapted_params = []
for p, g in zip(params, grads):
adapted_params.append(p - self.inner_lr * g)
# ---- 外循环 ----
# 用适配后的参数在 query set 上计算元损失
# 注意:这里需要保留计算图以支持二阶梯度
# 通过 ops.value_and_grad 实现高阶微分
meta_loss = self._compute_meta_loss(adapted_params, q_x, q_y)
# 计算 query set 准确率
q_logits = self._forward_with_params(adapted_params, q_x)
accuracy = self._compute_accuracy(q_logits, q_y)
return meta_loss, accuracy
def _inner_loss(self, x, y):
"""内循环损失函数,用于计算梯度"""
logits = self.backbone(x)
return nn.CrossEntropyLoss()(logits, y)
def _compute_meta_loss(self, adapted_params, q_x, q_y):
"""用适配后的参数计算 query set 上的损失"""
logits = self._forward_with_params(adapted_params, q_x)
return nn.CrossEntropyLoss()(logits, q_y)
def _forward_with_params(self, adapted_params, x):
"""用给定参数进行前向传播"""
# MindSpore 中通过 Functional API 实现
from mindspore import mutable
from mindspore.ops import composite
# 使用 mindspore.ops 的函数式 API 绑定参数
net = self.backbone
# 保存原始参数
original_params = [p.clone() for p in net.trainable_params()]
# 替换为适配后的参数
for orig, adapted in zip(net.trainable_params(), adapted_params):
orig.assign_value(adapted)
logits = net(x)
# 恢复原始参数
for orig, sa ved in zip(net.trainable_params(), original_params):
orig.assign_value(sa ved)
return logits
@staticmethod
def _compute_accuracy(logits, labels):
"""计算分类准确率"""
preds = ops.argmax(logits, axis=1)
correct = ops.equal(preds, labels).astype(ms.float32)
return ops.mean(correct)
class MAMLTrainer:
"""MAML 训练器,管理训练循环和任务采样"""
def __init__(self, num_classes=5, num_shots=5, num_query=15, inner_lr=0.01, outer_lr=0.001, num_inner_steps=1):
self.num_classes = num_classes
self.num_shots = num_shots
self.num_query = num_query
self.inner_lr = inner_lr
self.outer_lr = outer_lr
self.num_inner_steps = num_inner_steps
# 初始化模型
self.backbone = SimpleCNN(num_classes=num_classes)
self.maml = MAML(self.backbone, inner_lr=inner_lr, num_inner_steps=num_inner_steps)
# 外循环优化器
self.optimizer = nn.Adam(params=self.backbone.trainable_params(), learning_rate=outer_lr)
# 定义训练步骤
self.train_step_fn = self._build_train_step()
def _build_train_step(self):
"""构建训练步骤函数,使用 value_and_grad 实现二阶梯度"""
grad_fn = ms.ops.value_and_grad(self.maml.construct,
grad_position=None,
# 对所有参数求梯度
weights=self.backbone.trainable_params())
def train_step(sup_x, sup_y, q_x, q_y):
(meta_loss, accuracy), grads = grad_fn(sup_x, sup_y, q_x, q_y)
self.optimizer(grads)
return meta_loss, accuracy
return train_step
def train_epoch(self, task_generator, num_tasks=100):
"""训练一个 epoch,在 num_tasks 个任务上进行元更新"""
total_loss = 0.0
total_acc = 0.0
for _ in range(num_tasks):
sup_x, sup_y, q_x, q_y = task_generator.generate_task()
# 转换为 MindSpore Tensor
sup_x_t = Tensor(sup_x, ms.float32)
sup_y_t = Tensor(sup_y, ms.int32)
q_x_t = Tensor(q_x, ms.float32)
q_y_t = Tensor(q_y, ms.int32)
loss, acc = self.train_step_fn(sup_x_t, sup_y_t, q_x_t, q_y_t)
total_loss += loss.asnumpy()
total_acc += acc.asnumpy()
return total_loss / num_tasks, total_acc / num_tasks
def evaluate(self, task_generator, num_tasks=200):
"""在多个任务上评估模型"""
self.backbone.set_train(False)
total_acc = 0.0
for _ in range(num_tasks):
sup_x, sup_y, q_x, q_y = task_generator.generate_task()
sup_x_t = Tensor(sup_x, ms.float32)
sup_y_t = Tensor(sup_y, ms.int32)
q_x_t = Tensor(q_x, ms.float32)
q_y_t = Tensor(q_y, ms.int32)
# 执行内循环适配
logits_sup = self.backbone(sup_x_t)
loss_sup = nn.CrossEntropyLoss()(logits_sup, sup_y_t)
grads = ms.ops.grad(self.maml._inner_loss)(sup_x_t, sup_y_t)
params = self.backbone.trainable_params()
adapted_params = [p - self.inner_lr * g for p, g in zip(params, grads)]
# 在 query set 上评估
q_logits = self.maml._forward_with_params(adapted_params, q_x_t)
acc = self.maml._compute_accuracy(q_logits, q_y_t)
total_acc += acc.asnumpy()
self.backbone.set_train(True)
return total_acc / num_tasks
# 运行 MAML 训练
print("\n" + "=" * 60)
print("开始 MAML 训练 (5-way 5-shot)")
print("=" * 60)
# 使用更多 query 样本的 task generator
eval_task_gen = FewShotTaskGenerator(full_dataset, num_classes=5, num_samples_per_class=5)
trainer = MAMLTrainer(num_classes=5, num_shots=5, num_query=15,
inner_lr=0.01, outer_lr=0.001, num_inner_steps=1)
num_epochs = 30
for epoch in range(1, num_epochs + 1):
loss, acc = trainer.train_epoch(eval_task_gen, num_tasks=50)
if epoch % 5 == 0 or epoch == 1:
eval_acc = trainer.evaluate(eval_task_gen, num_tasks=100)
print(f"Epoch {epoch:3d} | Loss: {loss:.4f} | "
f"Train Acc: {acc:.4f} | Eval Acc: {eval_acc:.4f}")
print("\nMAML 训练完成!")
弄清楚了MAML的“内外循环”机制,再回头看代码,思路就会顺畅很多。至少我是这么觉得的——每一步梯度的传递、参数的拷贝和恢复,都是为了让训练过程既保留计算图,又不影响主网络的原始参数。
3.4 Prototypical Networks 实现
相比MAML,Prototypical Networks的实现就轻快多了。它不需要内循环,也没有二阶梯度的烦恼。一句话:算原型,算距离,分类。
class PrototypicalNetwork(nn.Cell):
"""Prototypical Networks 的 MindSpore 实现
通过计算类别原型和查询样本嵌入之间的距离进行分类。
不需要内循环梯度更新,训练效率高。
"""
def __init__(self, backbone, num_classes=5):
super(PrototypicalNetwork, self).__init__()
self.backbone = backbone # 只负责特征提取
self.num_classes = num_classes
self.squared_euclidean = self._squared_euclidean_distance
def construct(self, sup_x, sup_y, q_x, q_y):
"""
前向传播:
1. 提取 support 和 query 的嵌入
2. 计算每个类别的原型
3. 基于 embedding 到原型的距离进行分类
Returns:
loss: 交叉熵损失
accuracy: 分类准确率
"""
# 提取特征嵌入
sup_embeddings = self.backbone(sup_x) # [N*K, D]
q_embeddings = self.backbone(q_x) # [Q, D]
# 计算每个类别的原型(support 嵌入的均值)
prototypes = self._compute_prototypes(sup_embeddings, sup_y)
# 计算距离并执行分类
logits = self._compute_logits(q_embeddings, prototypes)
# 计算损失和准确率
loss = nn.CrossEntropyLoss()(logits, q_y)
accuracy = self._compute_accuracy(logits, q_y)
return loss, accuracy
def _compute_prototypes(self, embeddings, labels):
"""计算每个类别的原型向量
Args:
embeddings: [N*K, D] 特征嵌入
labels: [N*K] 类别标签
Returns:
prototypes: [num_classes, D] 每个类别的原型
"""
num_classes = self.num_classes
emb_dim = embeddings.shape[1]
prototypes = ops.Zeros()((num_classes, emb_dim), ms.float32)
counts = ops.Zeros()((num_classes,), ms.float32)
# 累加每个类别的嵌入
for cls in range(num_classes):
mask = ops.equal(labels, cls).astype(ms.float32) # [N*K]
mask = ops.expand_dims(mask, 1) # [N*K, 1]
cls_embeddings = embeddings * mask # 遮蔽非该类样本
count = ops.reduce_sum(mask) + 1e-8 # 防除零
prototype = ops.reduce_sum(cls_embeddings, axis=0) / count
prototypes = ops.TensorScatterUpdate(prototypes,
Tensor([[cls]], ms.int32),
ops.expand_dims(prototype, 0))
return prototypes
def _compute_logits(self, query_embeddings, prototypes):
"""基于到原型的负欧氏距离计算 logits
logits[i, k] = -||z_i - c_k||_2^2
Args:
query_embeddings: [Q, D]
prototypes: [K, D]
Returns:
logits: [Q, K]
"""
# 计算距离矩阵 [Q, K]
# 使用展开技巧高效计算
q = ops.expand_dims(query_embeddings, 1) # [Q, 1, D]
p = ops.expand_dims(prototypes, 0) # [1, K, D]
distances = ops.reduce_sum((q - p) ** 2, axis=2) # [Q, K]
# 负距离作为 logits(距离越小,logits 越大,概率越高)
return -distances
@staticmethod
def _compute_accuracy(logits, labels):
preds = ops.argmax(logits, axis=1)
correct = ops.equal(preds, labels).astype(ms.float32)
return ops.mean(correct)
@staticmethod
def _squared_euclidean_distance(a, b):
return ops.reduce_sum((a - b) ** 2)
class FeatureEncoder(nn.Cell):
"""Prototypical Networks 的特征编码器(4 层卷积 + 嵌入)"""
def __init__(self, in_channels=1, hidden_dim=64, embedding_dim=64):
super(FeatureEncoder, self).__init__()
self.encoder = nn.SequentialCell([
nn.Conv2d(in_channels, hidden_dim, 3, pad_mode='same', has_bias=True),
nn.BatchNorm2d(hidden_dim),
nn.ReLU(),
nn.MaxPool2d(2, 2),
nn.Conv2d(hidden_dim, hidden_dim, 3, pad_mode='same', has_bias=True),
nn.BatchNorm2d(hidden_dim),
nn.ReLU(),
nn.MaxPool2d(2, 2),
nn.Conv2d(hidden_dim, hidden_dim, 3, pad_mode='same', has_bias=True),
nn.BatchNorm2d(hidden_dim),
nn.ReLU(),
nn.MaxPool2d(2, 2),
nn.Conv2d(hidden_dim, embedding_dim, 3, pad_mode='same', has_bias=True),
nn.BatchNorm2d(embedding_dim),
nn.ReLU(),
nn.MaxPool2d(2, 2),
])
self.flatten = nn.Flatten()
def construct(self, x):
x = self.encoder(x)
return self.flatten(x)
class ProtoNetTrainer:
"""Prototypical Networks 训练器"""
def __init__(self, num_classes=5, learning_rate=0.001):
self.num_classes = num_classes
# 初始化编码器和 ProtoNet
encoder = FeatureEncoder(in_channels=1, hidden_dim=64, embedding_dim=64)
self.protonet = PrototypicalNetwork(encoder, num_classes=num_classes)
# 优化器
self.optimizer = nn.Adam(params=self.protonet.trainable_params(),
learning_rate=learning_rate)
# 训练步骤
self.train_step_fn = self._build_train_step()
def _build_train_step(self):
grad_fn = ms.ops.value_and_grad(self.protonet.construct,
grad_position=None,
weights=self.protonet.trainable_params())
def train_step(sup_x, sup_y, q_x, q_y):
(loss, acc), grads = grad_fn(sup_x, sup_y, q_x, q_y)
self.optimizer(grads)
return loss, acc
return train_step
def train_epoch(self, task_generator, num_tasks=100):
total_loss = 0.0
total_acc = 0.0
for _ in range(num_tasks):
sup_x, sup_y, q_x, q_y = task_generator.generate_task()
sup_x_t = Tensor(sup_x, ms.float32)
sup_y_t = Tensor(sup_y, ms.int32)
q_x_t = Tensor(q_x, ms.float32)
q_y_t = Tensor(q_y, ms.int32)
loss, acc = self.train_step_fn(sup_x_t, sup_y_t, q_x_t, q_y_t)
total_loss += loss.asnumpy()
total_acc += acc.asnumpy()
return total_loss / num_tasks, total_acc / num_tasks
def evaluate(self, task_generator, num_tasks=200):
self.protonet.set_train(False)
total_acc = 0.0
for _ in range(num_tasks):
sup_x, sup_y, q_x, q_y = task_generator.generate_task()
sup_x_t = Tensor(sup_x, ms.float32)
sup_y_t = Tensor(sup_y, ms.int32)
q_x_t = Tensor(q_x, ms.float32)
q_y_t = Tensor(q_y, ms.int32)
_, acc = self.protonet(sup_x_t, sup_y_t, q_x_t, q_y_t)
total_acc += acc.asnumpy()
self.protonet.set_train(True)
return total_acc / num_tasks
# 运行 Prototypical Networks 训练
print("\n" + "=" * 60)
print("开始 Prototypical Networks 训练 (5-way 5-shot)")
print("=" * 60)
proto_trainer = ProtoNetTrainer(num_classes=5, learning_rate=0.001)
num_epochs = 30
for epoch in range(1, num_epochs + 1):
loss, acc = proto_trainer.train_epoch(eval_task_gen, num_tasks=50)
if epoch % 5 == 0 or epoch == 1:
eval_acc = proto_trainer.evaluate(eval_task_gen, num_tasks=100)
print(f"Epoch {epoch:3d} | Loss: {loss:.4f} | "
f"Train Acc: {acc:.4f} | Eval Acc: {eval_acc:.4f}")
print("\nPrototypical Networks 训练完成!")
3.5 训练与评估流程
为了方便对比这两种算法,我们干脆封装一个统一的评估流程,这样跑一次实验,结果就一目了然:
def run_full_experiment():
"""完整的实验流程:训练 + 评估 + 对比"""
np.random.seed(42)
ms.set_seed(42)
# 构建数据集
dataset = create_synthetic_dataset(num_classes=50, samples_per_class=20)
# 1-shot 和 5-shot 实验配置
configs = [
{"num_shots": 1, "name": "1-shot"},
{"num_shots": 5, "name": "5-shot"},
]
results = {}
for config in configs:
print(f"\n{'=' * 60}")
print(f"实验配置: 5-way {config['name']}")
print(f"{'=' * 60}")
task_gen = FewShotTaskGenerator(dataset, num_classes=5,
num_samples_per_class=config['num_shots'] + 10)
# ---- MAML ----
print("\n--- MAML ---")
maml_trainer = MAMLTrainer(num_classes=5,
num_shots=config['num_shots'],
num_query=10,
inner_lr=0.01,
outer_lr=0.001)
for epoch in range(1, 21):
loss, acc = maml_trainer.train_epoch(task_gen, num_tasks=50)
if epoch % 5 == 0:
eval_acc = maml_trainer.evaluate(task_gen, num_tasks=100)
print(f"Epoch {epoch:3d} | Eval Acc: {eval_acc:.4f}")
maml_final = maml_trainer.evaluate(task_gen, num_tasks=200)
results[f"MAML-{config['name']}"] = maml_final
print(f"最终准确率: {maml_final:.4f}")
# ---- Prototypical Networks ----
print("\n--- Prototypical Networks ---")
proto_trainer = ProtoNetTrainer(num_classes=5, learning_rate=0.001)
for epoch in range(1, 21):
loss, acc = proto_trainer.train_epoch(task_gen, num_tasks=50)
if epoch % 5 == 0:
eval_acc = proto_trainer.evaluate(task_gen, num_tasks=100)
print(f"Epoch {epoch:3d} | Eval Acc: {eval_acc:.4f}")
proto_final = proto_trainer.evaluate(task_gen, num_tasks=200)
results[f"ProtoNet-{config['name']}"] = proto_final
print(f"最终准确率: {proto_final:.4f}")
# 打印汇总结果
print("\n" + "=" * 60)
print("实验结果汇总")
print("=" * 60)
print(f"{'方法':<25} {'准确率':>10}")
print("-" * 37)
for name, acc in results.items():
print(f"{name:<25} {acc:>10.4f}")
return results
if __name__ == "__main__":
results = run_full_experiment()
四、实验与结果分析
实验做完了,结论其实很清晰。
4.1 实验设置
| 配置 | 描述 |
|---|---|
| 任务分布 | 50 个类别,每类 20 个样本 |
| 评估方式 | 5-way 1-shot / 5-way 5-shot |
| 内循环步数 | 1(MAML) |
| 训练任务数/epoch | 50 |
| 评估任务数 | 200 |
| 训练轮数 | 20 epochs |
4.2 结果分析
跑完整个实验,你能看到几件有意思的事。
首先,元学习显著优于随机初始化。 传统方法在5-way 1-shot的设置下,就靠那5个样本,基本是瞎猜(准确率接近20%)。但经过元训练的MAML和Prototypical Networks,在新任务上能快速“进入状态”。这不就是“学会学习”的魔法吗?
其次,MAML和Prototypical Networks各有千秋。
- MAML 强在通过内循环梯度更新进行任务适配,理论上更灵活,但代价是二阶梯度计算的开销。如果用一阶近似的FOMAML能快很多,不过性能会稍微打点折扣。
- Prototypical Networks 走的是度量学习的路,没有内循环更新,训练和推理都格外高效。在类别区分度高的数据上,它的表现相当亮眼。
- 总的来说,如果训练资源有限,Prototypical Networks是个不错的入门选择;而在面对更复杂的任务时,MAML的潜力更值得挖掘。
再次,K-shot的数量很关键。 从1-shot增加到5-shot,两种方法的准确率都拔高了一大截。道理简单:更多的支持样本,让原型估计(ProtoNet)更靠谱,也让梯度方向(MAML)更稳定。
4.3 消融实验建议
想再往前深挖一步?这几组消融实验值得一试:
- 二阶 vs. 一阶MAML:对比完整MAML和FOMAML的性能差距与计算时间。
- 内循环步数:试试1-step、3-step、5-step,看看多步更新收益如何。
- 不同嵌入维度:对比32/64/128维嵌入对ProtoNet性能的影响。
- 不同的距离度量:ProtoNet中用余弦距离替换欧氏距离,效果会有何不同。
五、总结与展望
这篇文章从理论讲到实战,把元学习在MindSpore里的实现拆了个清楚。MAML和Prototypical Networks的代码,大家可以直接上手跑。
核心收获
- 元学习的本质,是优化初始参数或学习策略,让模型在新任务上快速适配。这比传统的迁移学习,更像一个系统性的few-shot解决方案。
- MAML通过内外循环的二阶梯度优化,找到敏感的初始参数,适合需要显式任务适配的场景。
- Prototypical Networks则用度量学习的方式计算类别原型,简洁高效,是不需要内循环的经典代表。
- MindSpore的自动微分能力(
ops.value_and_grad、ops.grad)天生就支持二阶梯度计算,做MAML这类算法,非常顺手。
未来方向
元学习现在还是个很活跃的领域,以下几个方向值得关注:
- 元强化学习(Meta-RL):把元学习和强化学习结合起来,让机器人快速学会新技能。
- 元学习的可扩展性:解决大规模任务下的计算瓶颈,比如用隐式梯度方法。
- 与预训练大模型结合:用元学习的思路,增强大语言模型的few-shot能力。
- 垂直领域落地:在医疗诊断、自动驾驶、工业质检这些标注成本高昂的场景里,元学习大有可为。
元学习的核心哲学——“让机器像人一样举一反三”,正在推动人工智能走向更通用、更高效的方向。而MindSpore作为国产框架,凭借它在自动微分和高阶梯度计算上的硬实力,为这一领域的研究提供了扎实的工程基础。
