jax-ml/jax

Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more

何を解決するか

JAX は高性能な数値計算および大規模機械学習のためのシステムを提供します。開発者は NumPy に似たコードを記述でき、自動微分、高速化のためのコンパイル、GPU や TPU などの複数のハードウェアアクセラレータにわたるスケーリングが自動的に行われます。

動作方法

JAX は XLA(Accelerated Linear Algebra)を使用して、Python および NumPy 関数を最適化されたマシンコードにコンパイルします。これは、合成可能な関数変換のシステムとして動作します:

  • jax.grad: 自動微分(逆モードおよび順モードの両方をサポート)を使用して、Python および NumPy 関数の勾配を計算します。
  • jax.jit: XLA を使って関数をエンドツーエンドにコンパイルし、実行速度を向上させます。
  • jax.vmap: 関数を自動的にベクトル化し、配列の軸に沿ってマッピングして、手動でのバッチループを排除します。

対象ユーザー

高性能な数値計算および大規模機械学習に取り組む研究者や開発者向けに設計されており、効率的な勾配計算と数千台のデバイスにわたる計算のスケーリングが必要な方々に最適です。

特徴

  • 合成可能な変換: gradjitvmap を任意の順序で組み合わせて、高度に最適化された関数を作成できます。
  • ハードウェアアクセラレーション: XLA を通じて NVIDIA GPU、Google TPU、その他のアクセラレータをネイティブにサポートします。
  • 柔軟な微分: ループ、分岐、再帰、クロージャーを任意の次数まで微分できます。
  • スケーリングオプション: コンパイラベースの自動並列化、明示的なシャーディング、および手動でのデバイスごとのプログラミングの3つのスケーリングモードを提供します。

関連

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