JAX: Python と NumPy プログラムのための合成可能な変換
Python+NumPy プログラムのコンポーザブル変換: 微分、ベクトル化、JIT から GPU/TPU など。
ひと目でわかる
- これは何?
- jax-ml/jax の README は、grad、jit、vmap の三つの変換、三つのスケーリングモード、そして各プラットフォームの CPU・GPU・TPU サポートを説明している。
- 誰に向いている?
- README は JAX を、Python と NumPy プログラムの上で grad、jit、vmap の三つの変換を合成し、三つのデバイスモードのいずれかでスケールさせるシステムとして位置づけている。README の例と表の先にあるものはすべて、docs.jax.dev のリンクされたドキュメントに記載されている。
- 商用利用できる?
- できます。Apache-2.0 は寛容なライセンスで、著作権表示とライセンス表示を残せば、使用・改変・販売が可能です。
- 今もメンテナンスされている?
- されています。直近 1 日以内に新しいコミットがあります。
- 何の言語で書かれている?
- 主に Python です(GitHub の言語統計による)。
回答はプロジェクトの GitHub データ(最終同期:2026年9月15日)と当サイトの分析に基づくもので、法的助言ではありません。
オープンソース詳細解説
jax-ml-jax-deep-analysis: 数値関数を変換するためのライブラリ
JAX は、アクセラレータ向け配列計算とプログラム変換のための Python ライブラリで、高性能な数値計算と大規模機械学習を目的としている。README の中心的な主張は、JAX がネイティブの Python 関数と NumPy 関数を自動微分でき、それがループ、分岐、再帰、クロージャを通しても機能するという点だ。逆モード微分は jax.grad として提供され、順モード微分もサポートされ、両者は任意の次数で合成できる。その下では XLA が NumPy プログラムを TPU、GPU、その他のハードウェアアクセラレータ上でコンパイルし実行する。README はまた、コンパイルと自動微分が任意に合成できるとも述べている。
JAX は Python と NumPy の数値関数に変換を適用するライブラリで、jax.grad、jax.jit、jax.vmap を中核としている。grad は逆モードと順モードを組み合わせられ、jit は XLA によるコンパイルを行い、vmap は配列軸への写像を処理する。README の tanh、二乗誤差、per-example gradient の例が三者の接点を示す。
jax-ml-jax-deep-analysis: 中核となる三つの変換: grad、jit、vmap
README は三つの変換をシステムの中核として紹介している。jax.grad は逆モード勾配を計算し、例では tanh 関数の一階と三階の導関数、および if/else 分岐を持つ関数の微分を示している。jax.jit は XLA で関数をエンドツーエンドにコンパイルし、デコレータとしても高階関数としても使え、関数内で使える Python 制御フローの種類に制約を課す。jax.vmap は関数を配列の軸に沿ってマッピングし、ループを関数のプリミティブ操作に押し下げる。これにより行列ベクトル積が行列行列積になり、バッチ次元をコードで持ち歩く必要がなくなる。README は詳細を Autodiff Cookbook とリファレンスドキュメントに委ねている。
jax.jit は任意の Python コードをそのまま実行する機能ではなく、制御フローに制約がある。まず純粋な関数を CPU で呼び、jax.grad の数値、jit 後の初回と二回目の時間、vmap の出力 shape を記録すると、変換ごとの役割を分けて確認できる。README の印字例は f32[512@data,512] である。
jax-ml-jax-deep-analysis: 一つの損失関数に grad、jit、vmap を合成する
README の主要な例は、パラメータのペアをループで回す予測関数と二乗誤差損失を定義し、grad_loss = jax.jit(jax.grad(loss)) と perex_grads = jax.jit(jax.vmap(grad_loss, in_axes=(None, 0, 0))) を構築して、高速なサンプルごとの勾配を得る。同じパターンはスケーリングの節でも現れ、jit と grad を適用した損失が、明示的にシャーディングされたパラメータとデータに対して評価される。例の出力は f32[512@data,512] のような型を印字し、パラメータテンソルが data 軸のシャーディング注釈を持つことを示している。
スケーリングには compiler-based automatic parallelization、explicit sharding and automatic partitioning、manual per-device programming の三つがある。前者はグローバルな視点で分割を任せ、後者二つは型や collectives を通じて分割を見えるようにする。jax.typeof と jax.device_put の結果を確認すれば、データ軸の指定を追跡できる。
jax-ml-jax-deep-analysis: 自動から手動までの三つのスケーリングモード
数千台のデバイスにスケールするために、README は三つのアプローチを記録している。コンパイラベースの自動並列化は、単一のグローバルマシンでプログラミングしているかのように書き、データのシャーディングと計算の分割をコンパイラが選択する。ユーザーがいくつかの制約を与えることはできる。明示的シャーディングと自動分割はグローバルな視点を保つが、データのシャーディングは JAX の型の中で可視になり、jax.typeof で検査できる。手動のデバイス単位プログラミングはデバイス単位の視点に切り替え、明示的なコレクティブを許す。README の表は、三つのモードをグローバルかデバイス単位かの視点、明示的シャーディングの有無、明示的コレクティブの有無で要約している。
CPU は Linux x86_64、Linux aarch64、Mac aarch64、Windows x86_64、Windows WSL2 x86_64 に対応すると README は表で示す。NVIDIA GPU、AMD GPU、TPU、Apple GPU、Intel GPU には OS ごとの制限や実験扱いがある。pip install -U jax、jax[cuda13]、jax[tpu] などの入口を実機に合わせ、公式サポート表と結果を照合する必要がある。
jax-ml-jax-deep-analysis: プラットフォームサポートと文書化されたインストールコマンド
README のプラットフォーム表は、CPU が Linux x86_64、Linux aarch64、Mac aarch64、Windows x86_64、Windows WSL2 x86_64 でサポートされることを示している。NVIDIA GPU は Linux の二つのバリアントでサポートされ、Windows WSL2 では実験的、その他では利用できない。Google TPU は Linux x86_64 のみ。AMD GPU は Linux x86_64 でサポートされ、Windows WSL2 では実験的。Apple GPU は Mac aarch64 で実験的、Intel GPU は Linux x86_64 で実験的。README に記載されたインストールコマンドは、CPU が pip install -U jax、NVIDIA GPU が pip install -U "jax[cuda13]"、Google TPU が pip install -U "jax[tpu]"、Linux の AMD GPU が pip install -U "jax[rocm7-local]" である。Intel GPU のインストールは Intel 自身の手順に従う。ソースからのビルド、Docker、他の CUDA バージョン、コミュニティの conda ビルドなどの代替手段はドキュメントに委ねられている。
jax-ml-jax-deep-analysis: プロジェクトの位置づけ、引用、Apache-2.0 ライセンス
README は、JAX が研究プロジェクトであり Google の公式製品ではないと明言し、鋭い縁があると警告して、gotchas ノートブックとイシュートラッカーを指している。引用エントリは著者をアルファベット順に並べ、jax/version.py の 0.3.13 をバージョン番号とし、オープンソースリリースの年として 2018 を挙げている。自動微分と XLA コンパイルのみを備えた初期バージョンは SysML 2018 の論文で説明された。README は、より包括的で最新の論文を準備中としている。リポジトリは Apache-2.0 のもとで公開され、永続的、全世界的、非独占的、無償、ロイヤリティフリー、取消不能の著作権ライセンスと、特許訴訟を提起した場合に終了する特許ライセンスを付与する。ライセンスの抜粋は保証、サポート、セキュリティ体制については触れていない。リポジトリのメタデータは、この記事の準備時点で 36,103 スター、3,716 フォーク、2,554 のオープンイシューを示している。
編集部の結論
README は JAX を、Python と NumPy プログラムの上で grad、jit、vmap の三つの変換を合成し、三つのデバイスモードのいずれかでスケールさせるシステムとして位置づけている。README の例と表の先にあるものはすべて、docs.jax.dev のリンクされたドキュメントに記載されている。
コミュニティノート