metaopt/torchopt
TorchOpt is an efficient library for differentiable optimization built upon PyTorch.
解决的问题
TorchOpt 提供了 PyTorch 中可微分优化的框架,使用户能够通过优化过程计算梯度。这对于双层优化问题(如元学习)至关重要,其中外层参数必须根据内层优化循环的结果进行优化。
工作原理
TorchOpt 实现了三种主要的微分模式,以应对不同的优化场景:
- 显式梯度 (EG): 通过展开的优化路径进行反向传播,将每个梯度步骤视为可微函数。适用于内层循环步数较少的情况。
- 隐式梯度 (IG): 使用隐函数定理,在内层循环的驻点处找到解析导数,避免了展开整个优化路径的需要。
- 零阶微分 (ZD): 当内层循环不可微或 Hessian 计算成本过高时,使用有限差分或进化策略(ES)等零阶方法来估计梯度。
它提供了两种 API:一种是类似 JAX/Optax 的函数式 API,另一种是类似标准 PyTorch torch.optim 的面向对象 API,以适应不同的编码偏好。
适用人群
专为从事元学习、超参数优化及其他双层优化任务的研究人员和开发者设计,提供一种高效、灵活的方式来对 PyTorch 优化器进行微分。
特性亮点
- 三种微分模式: 支持显式、隐式和零阶梯度。
- 灵活的 API: 提供类似 JAX 的函数式接口和类似 PyTorch 的面向对象接口。
- 性能优化: 包含 C++/CUDA 加速算子和基于 RPC 的分布式训练框架。
- 函数式集成: 与
functorch对接,支持在 PyTorch 中实现可组合的函数式优化。
相关
- 项目
- 项目
- 项目
- 项目
- 项目