higher
PyTorch元学习核心库:穿越优化器迭代,计算梯度之上的梯度
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
PyTorch元学习核心库:穿越优化器迭代,计算梯度之上的梯度
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。

项目:facebookresearch/higher · PyTorch · Apache-2.0 · 1.6k★
higher 是 Facebook Research 出品的 PyTorch 元学习核心工具包——它能让开发者"穿越"优化器的迭代过程,计算梯度之上的梯度,从而支撑 MAML、Reptile 等元学习算法的端到端反向传播。 这个库的核心价值在于打破 PyTorch 原有梯度流的限制:通常情况下,参数更新发生在优化器内部,梯度链会在第一次更新后就断开;higher 通过 monkey-patch 机制将 torch.nn.Module 转为"无状态函数",把参数以外部张量形式传入 forward,让梯度穿过任意步数的优化循环。
要理解 higher,先要理解元学习(Meta-Learning)的核心诉求。传统深度学习是"学习一个模型去完成一个任务",而元学习是"学习一个学习方法来生成模型"——用通俗的话说,就是训练一个"会学习"的算法,让它学会快速适应新任务。
这个过程通常涉及双层优化(bi-level optimization):
问题来了:PyTorch 的标准优化器(torch.optim.SGD / Adam)是黑盒状态机,参数更新发生在 step() 内部,梯度链在第一次参数更新后就断了。你无法直接计算"模型参数经过 5 步更新后的梯度"。这就是 higher 要解决的核心问题。
higher 的技术核心分为两个部分。
PyTorch 的 torch.nn.Module 是有状态的——参数存储在 self._parameters 里,forward 依赖这些内部状态。higher 通过 higher.patch.monkeypatch() 将任意 Module 转换为无状态函数版本,把参数以外部列表(List[Tensor])形式传入:
import higher
# 原模型(stateful)
model = MyResNet()
# Patch 后(stateless)
fmodel, diff_opt = higher.make_functional(model)
# 正常 forward:把参数作为第一个参数传入
output = fmodel(params, x) # params 是外部张量列表,不依赖内部状态
_patch.py 的实现原理:
_fast_params[time] 时间索引管理,每次优化器 step 后追加新参数快照_expand_params() 将扁平参数列表映射回 Module 的各层结构_patch.py 配合 higher/optim.py 里的 DifferentiableOptimizer 基类,将标准 PyTorch 优化器(SGD、Adam)改造为梯度可传递的版本:
# 获取可微分优化器
diff_opt = higher.get_diff_optim(
torch.optim.Adam(model.parameters()),
fmodel
)
# 内部循环:多次微分优化器 step
for step in range(k):
loss = task_loss(fmodel(params, x_task), y_task)
grads = torch.autograd.grad(loss, params)
diff_opt.step(grads)
optim.py 的核心逻辑:
_stateregister_optim 注册自定义优化器higher.innerloop_ctx 是最常用的入口 API,将以上两部分封装为一个上下文管理器:
with higher.innerloop_ctx(model, opt) as (fmodel, diff_opt):
# 在这个上下文里,fmodel 是无状态的,diff_opt 可微分
for _ in range(k):
loss = task_loss(fmodel(x), y)
grads = torch.autograd.grad(loss, fmodel.parameters())
diff_opt.step(grads)
# fmodel.parameters() 现在是 k 步后的参数
# 可以计算 P[k] 对原始参数 P[0] 的梯度
higher 最早也是最直接的应用场景是实现 MAML。原始 MAML 论文需要手写梯度推导,而用 higher 可以直接写出端到端的计算图:
# examples/maml-omniglot.py 核心逻辑
with innerloop_ctx(model, opt) as (fmodel, diff_opt):
# 内层:快速适应每个任务
q_loss = 0
for task_x, task_y in batch_tasks:
# 几步内层更新
fast_loss = criterion(fmodel(task_x), task_y)
diff_opt.step(torch.autograd.grad(fast_loss, params))
# 外层:用所有任务验证损失的反梯度调整初始参数
meta_loss = sum(criterion(fmodel(x), y) for x, y in batch_tasks)
meta_opt.zero_grad()
meta_loss.backward()
meta_opt.step()
当内层优化步数较多时,直接展开计算图会显存爆炸。higher 支持近似方案和隐式梯度方法,适用于上百步内层更新的场景。
用 higher 可以将对模型架构参数的梯度纳入优化循环,实现可微分的 Neural Architecture Search。
higher/
├── __init__.py # API 导出:innerloop_ctx, monkeypatch, get_diff_optim
├── patch.py # 核心:无状态化 Module 实现(21KB)
├── optim.py # 核心:可微分优化器封装(46KB)
└── utils.py # 辅助工具函数
patch.py 核心类:
_MonkeyPatchBase:抽象基类,定义了 _fast_params、_parma_mapping 等核心数据结构FunctionalModule:无状态函数式模块,参数作为 forward 第一参数MonkeyPatchable + patched_type:monkeypatch 装饰器实现optim.py 核心类:
DifferentiableOptimizer:抽象基类,step() 用 torch.no_grad() + autograd 操作替代原地更新DiffSGD、DiffAdam:标准优化器的可微分版本register_optim():运行时注册自定义可微分优化器register_optim 轻松添加自定义可微分优化器higher 是元学习领域的底层基础设施库,而非面向终端用户的应用库。它的核心贡献在于:把原本需要手动推导梯度公式的元学习算法,变成了可以端到端自动微分的 PyTorch 代码。
该项目代表了 2019-2021 年元学习热潮中"可微分编程"思路的典型实现。虽然近年来元学习热度有所回落,但 higher 提出的"优化器穿越梯度"范式——即通过无状态化让梯度穿过参数更新——至今仍是隐式学习(Implicit Learning)和梯度元学习研究的重要工具链。
对于想入门元学习的开发者,建议从 examples/maml-omniglot.py 入手,理解 innerloop_ctx 的使用方式,再逐步深入 patch.py 和 optim.py 的源码。