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 中實現可組合的函數式最佳化。

相關

  • 專案
  • 專案
  • 專案
  • 專案
  • 專案