TorchOpt API参考速查:从优化器到Transform的完整函数地图
TorchOpt API参考速查:从优化器到Transform的完整函数地图
【免费下载链接】torchoptTorchOpt is an efficient library for differentiable optimization built upon PyTorch.项目地址: https://gitcode.com/gh_mirrors/to/torchopt
TorchOpt是一个构建在 PyTorch 之上的高效可微优化(differentiable optimization)库,提供 Optax 风格的功能式 API:用"梯度变换"(Transform)自由组合出优化器,并内置隐式梯度、零阶梯度等可微回传能力。本文整理 TorchOpt 从优化器、Transform 到工具函数、元优化的完整函数地图,帮你 10 分钟定位任意 API。
🗺️ 一分钟看懂:TorchOpt 的双模式 API
TorchOpt 的所有 API 围绕一个核心抽象GradientTransformation(定义在 torchopt/base.py)展开——它是一对纯函数(init_fn, update_fn):init根据参数生成优化器状态,update把梯度变换成最终更新量。
| 风格 | 入口 | 特点 | 源码位置 |
|---|---|---|---|
| 功能式(推荐) | opt.sgd(lr)、opt.adam(lr)等 | 返回GradientTransformation,可任意组合 | torchopt/alias/__init__.py |
| 面向对象 | opt.optim.Adam(model.parameters(), lr=...) | 兼容torch.optim使用习惯 | torchopt/optim/__init__.py |
两种风格共享同一套 Transform 底层,功能式风格还能无缝对接chain、clip_grad_norm等组合工具。
🚀 优化器API速查表:8大预设,两种调用方式
TorchOpt 内置 8 类主流优化器,功能式别名(torchopt/alias/)与面向对象类(torchopt/optim/)一一对应,另有Meta*元优化变体支持二阶可微:
| 优化器 | 功能式 API | OO 类 | 元优化变体 |
|---|---|---|---|
| SGD | opt.sgd(lr, momentum, nesterov) | SGD | MetaSGD |
| AdaGrad | opt.adagrad(lr, eps) | AdaGrad | MetaAdaGrad |
| AdaDelta | opt.adadelta(eps) | AdaDelta | MetaAdaDelta |
| Adam | opt.adam(lr, betas, eps) | Adam | MetaAdam |
| AdamW | opt.adamw(lr, betas, eps, wd) | AdamW | MetaAdamW |
| AdaMax | opt.adamax(lr, betas) | Adamax | MetaAdaMax |
| RAdam | opt.radam(lr, betas, eps) | RAdam | MetaRAdam |
| RMSProp | opt.rmsprop(lr, decay, eps) | RMSProp | MetaRMSProp |
💡 元优化变体统一位于 torchopt/optim/meta/,例如
MetaAdam对超参数也保持可微,是 MAML 等元学习算法的基础。
🧩 Transform梯度变换组件:14个预设函数
torchopt.transform模块(torchopt/transform/init.py)提供 14 个预设变换,是搭建自定义优化器的"乐高积木":
| 函数 | 作用 | 典型用途 |
|---|---|---|
scale(step_size) | 按固定系数缩放更新量 | 学习率 |
scale_by_schedule(scheduler) | 按调度函数缩放 | 学习率衰减 |
scale_by_adam(betas, eps) | Adam 自适应缩放 | 组合 Adam |
scale_by_adamax(betas, eps) | AdaMax 自适应缩放 | 组合 AdaMax |
scale_by_radam(betas, eps) | RAdam 自适应缩放 | 组合 RAdam |
scale_by_rms / scale_by_rss / scale_by_adadelta | 基于均方根等统计量缩放 | 组合 RMS 类优化器 |
scale_by_stddev | 按标准差缩放 | 自适应学习率 |
add_decayed_weights(wd) | L2 权重衰减 | AdamW 组合 |
masked(mask) | 按掩码屏蔽更新 | 稀疏/冻结参数 |
nan_to_num(nan, posinf, neginf) | 替换nan/inf梯度 | 训练稳定性 |
trace(order) | 一阶迹估计(随机迹) | 隐式梯度计算 |
chain(*transforms) | 串联多个变换 | 组装优化器 |
🔗 用 chain 一行组合优化器
chain(torchopt/combine.py)把多个 Transform 串成流水线,再叠加clip_grad_norm(torchopt/clip.py)做梯度裁剪,三行代码即可等价于一个带裁剪的 Adam:
import torchopt as opt adam_with_clip = opt.chain( opt.transform.scale(0.001), opt.transform.scale_by_adam(), opt.clip_grad_norm(1.0) )配套的opt.apply_updates(params, updates)(torchopt/update.py)负责把变换后的更新量写回参数,支持inplace原地更新。
⏱️ 调度器与实用工具函数
- 学习率调度(torchopt/schedule/init.py):
linear_schedule、polynomial_schedule、exponential_decay,可直接传给scale_by_schedule; - 停梯度:
stop_gradient(torchopt/utils.py)阻断对张量的反向传播,是构建可微优化器时的关键技巧; - 状态管理:
extract_state_dict/recover_state_dict用于取出和恢复模块状态,module_clone/module_detach_提供模块克隆与原地 detach; - 梯度 Hook:
register_hook、nan_to_num_hook、zero_nan_hook(torchopt/hook.py),在update前对梯度做自定义拦截; - PyTree 工具:
tree_map、tree_flatten等(torchopt/pytree.py),基于 optree 处理参数树结构。
🧠 元优化与三种可微回传模式
TorchOpt 的最大亮点是可微优化:把优化步骤本身也放进计算图。它提供三种回传模式(源码位于 torchopt/diff/):
| 模式 | API | 原理 | 适用场景 |
|---|---|---|---|
| 显式梯度 | 默认torch.autograd | 直接对整段优化代码求导 | 步数少、结构简单 |
| 隐式梯度 | opt.diff.implicit.custom_root、ImplicitMetaGradientModule | 隐函数梯度定理,跳过展开 | 步数多、内存友好(iMAML) |
| 零阶梯度 | opt.diff.zero_order.zero_order | 有限差分离散扰动估计 | 黑箱、不可导场景 |
配套模块 torchopt/nn/module.py 提供MetaGradientModule、ImplicitMetaGradientModule、ZeroOrderGradientModule以及reparameterize、swap_state等工具,让你把任意nn.Module包装成支持元梯度回传的模块。
以隐式 MAML(iMAML)为例,结合MetaSGD与custom_root训练出的 few-shot 模型在 Omniglot 上的精度曲线如下(完整实现见 examples/iMAML/imaml_omniglot_functional.py):
📊 可视化与进阶模块
- 计算图可视化:
make_dot、resize_graph(torchopt/visual.py),用 Graphviz 渲染 TorchOpt 的计算图,比 torchviz 更能保留元梯度结构:
- 加速算子:
torchopt.accelerated_op提供 CUDA 加速的 Adam 算子,与torch.optim兼容,可用opt.accelerated_op_available()检测; - 线性求解:torchopt/linalg/ 提供共轭梯度(
cg)与非线性共轭梯度(ns),torchopt/linear_solve/ 封装solve_cg、solve_inv、solve_normal_cg,用于隐式梯度的线性系统求解; - 分布式训练:torchopt/distributed/ 提供
parallelize、parallelize_sync等 RPC 并行原语,支持多进程并行训练元优化器(参考 examples/distributed/few-shot/maml_omniglot.py)。
⚡ 速查小结
| 我想…… | 用什么 |
|---|---|
| 快速搭建 Adam | opt.chain(opt.transform.scale(lr), opt.transform.scale_by_adam()) |
| 兼容 torch.optim 写法 | opt.optim.Adam(params, lr=...) |
| 训练可微调的优化器 | opt.MetaAdam/MetaSGD |
| 减少长序列优化的显存 | opt.diff.implicit.custom_root |
| 处理黑箱不可导目标 | opt.diff.zero_order.zero_order |
| 裁剪 / 清洗梯度 | opt.clip_grad_norm(1.0)、opt.nan_to_num |
| 调试计算图 | opt.visual.make_dot |
所有公开 API 均从 torchopt/init.py 统一导出,配套测试用例在 tests/ 目录下,按模块一一对应(如test_optim.py、test_transform.py、test_implicit.py),可作为每个函数的用法示例快速查阅。
【免费下载链接】torchoptTorchOpt is an efficient library for differentiable optimization built upon PyTorch.项目地址: https://gitcode.com/gh_mirrors/to/torchopt
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
