google-deepmind/kfac-jax
Second Order Optimization and Curvature Estimation with K-FAC in JAX.
What it solves
KFAC-JAX は、使いやすい K-FAC (Kronecker-factored Approximate Curvature) オプティマイザと曲率推定器の実装を提供します。二階最適化は通常、一階法よりも計算コストが高くなりますが、本ライブラリはスケーラブルな曲率行列の近似を提供することで、この課題を解決します。
How it works
JAX 上に構築されており、本ライブラリは損失関数の曲率を近似することで、ニューラルネットワークの最適化を行うことができます。ユーザーは、ライブラリが曲率行列を正しく近似できるように、特定の損失関数を登録する必要があります (例: register_softmax_cross_entropy_loss)。オプティマイザは内部的なステージング (JIT と PMAP) を行います。つまり、ユーザーはオプティマイザの step function または損失関数に対して手動で jax.jit や jax.pmap を適用することはできません。
Step function の使用法-
Who it’s for
JAX を使用しており、ニューラルネットワークの学習を改善するために二階最適化技術を実装したい AI 研究者。
Highlights
- JAX-native: 加速器上での高パフォーマンスを実現するために、本ライブラリ専用に構築されています。
- Automatic Staging: オプティマイザの step に対して
jitとpmapを自動的に処理します。 - Adaptive Parameters: 適応的な学習率、適応的なモメンタム、適応的なダンピングをサポートします。
- Curvature Estimation: ニューラルネットワークの学習用に、スケーラブルな曲率近似を提供します。
関連
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト