metaopt/torchopt

TorchOpt is an efficient library for differentiable optimization built upon PyTorch.

何を解決するか

TorchOptはPyTorchにおける微分可能な最適化のフレームワークを提供し、最適化プロセスを介して勾配を計算できるようにします。これは、外側のパラメータが内側の最適化ループの結果に基づいて最適化される必要がある、メタラーニングなどの二段階最適化問題において不可欠です。

動作方法

TorchOptは、異なる最適化シナリオに対応する3つの主要な微分モードを実装しています:

  • 明示的勾配 (EG): 展開された最適化パスを逆伝播し、各勾配ステップを微分可能な関数として扱います。内側ループのステップ数が少ない場合に最適です。
  • 陰的勾配 (IG): 陰関数定理を用いて、内側ループの定常点における解析的導関数を求めるため、最適化パス全体を展開する必要がありません。
  • ゼロ次微分 (ZD): 内側ループが微分不可能またはヘッセ行列の計算がコストが高すぎる場合、有限差分や進化的戦略(ES)などのゼロ次方法で勾配を推定します。

関数型API(JAX/Optaxに似たもの)とオブジェクト指向API(標準的なPyTorch torch.optimに似たもの)の両方を提供し、異なるコーディングスタイルに対応しています。

対象ユーザー

メタラーニング、ハイパーパラメータ最適化、その他の二段階最適化タスクに取り組んでいる研究者や開発者向けに設計されており、PyTorchオプティマイザを微分可能に効率的かつ柔軟に扱えるようにします。

特徴

  • 3つの微分モード: 明示的、陰的、ゼロ次勾配をサポート。
  • 柔軟なAPI: JAX風の関数型とPyTorch風のオブジェクト指向インターフェースの両方を提供。
  • パフォーマンス最適化: C++/CUDA加速演算子とRPCベースの分散学習フレームワークを含む。
  • 関数型統合: functorchと統合され、PyTorchにおける合成可能な関数型最適化を可能にします。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト