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.jitjax.pmap を適用することはできません。 Step function の使用法-

Who it’s for

JAX を使用しており、ニューラルネットワークの学習を改善するために二階最適化技術を実装したい AI 研究者。

Highlights

  • JAX-native: 加速器上での高パフォーマンスを実現するために、本ライブラリ専用に構築されています。
  • Automatic Staging: オプティマイザの step に対して jitpmap を自動的に処理します。
  • Adaptive Parameters: 適応的な学習率、適応的なモメンタム、適応的なダンピングをサポートします。
  • Curvature Estimation: ニューラルネットワークの学習用に、スケーラブルな曲率近似を提供します。

関連

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