Loading article…
Google JAX ロゴ | |
| 開発者 | グーグル |
|---|---|
| プレビューリリース | v0.4.31 / 2024年7月30日 |
| リポジトリ | github.com/google/jax |
| 書かれた | Python、C++ |
| オペレーティング·システム | Linux、macOS、Windows |
| プラットフォーム | Python、NumPy |
| サイズ | 9.0MB |
| タイプ | 機械学習 |
| ライセンス | アパッチ2.0 |
| Webサイト | jax.readthedocs.io/ja/latest/ より |
Google JAXは、数値関数を変換するための機械学習フレームワークです。[1] [2] [3]これは、autograd(関数の微分による勾配関数の自動取得)の修正版とTensorFlowのXLA(Accelerated Linear Algebra)を組み合わせたものと説明されています。NumPyの構造とワークフローに可能な限り忠実に従うように設計されており、 TensorFlowやPyTorchなどのさまざまな既存のフレームワークと連携します。[4] [5] JAXの主な機能は次のとおりです。[1]
- grad:自動微分
- jit: コンパイル
- vmap:自動ベクトル化
- pmap:単一プログラム、複数データ(SPMD) プログラミング
卒業
以下のコードは、grad関数の自動微分を示しています。
# インポート
jax インポート gradから
jax.numpyを jnpとして インポートする
# ロジスティック関数を定義する
def ロジスティック( x ):
jnp.exp ( x ) / ( jnp.exp ( x ) +1 )を返す
# ロジスティック関数の勾配関数を取得する
grad_logistic = grad (ロジスティック)
# x = 1 におけるロジスティック関数の勾配を評価する
grad_log_out = grad_logistic ( 1.0 )
印刷( grad_log_out )
最後の行は次のように出力される。
0.19661194
ジット
以下のコードは、融合によるJIT関数の最適化を示しています。
# インポート
jax からjitをインポート
jax.numpyを jnpとして インポートする
# キューブ関数を定義する
立方体の定義( x ):
x * x * xを返す
# データを生成する
x = jnp . 1 (( 10000 , 10000 ))
# キューブ関数の JIT バージョンを作成する
jit_cube = jit (キューブ)
# 速度比較のために、cube 関数と jit_cube 関数を同じデータに適用します
立方体( x )
jit_cube ( x )
(行 #17)の計算時間は(行 #16) の計算時間jit_cubeよりも明らかに短くなるはずですcube。行 #7 の値を増やすと、その差はさらに大きくなります。
vマップ
以下のコードは、vmap関数のベクトル化を示しています。
# インポート
jax からvmapを部分的にインポート
jax.numpyを jnpとして インポートする
# 関数を定義する
def grads (自己、 入力):
in_grad_partial = jax.partial ( self._net_grads , self._net_params )は、
grad_vmap = jax.vmap ( in_grad_partial )です。
rich_grads = grad_vmap (入力)
flat_grads = np . asarray ( self . _flatten_batch ( rich_grads ))
flat_grads.ndim == 2かつflat_grads.shape [ 0 ] == inputs.shape [ 0 ]であることをアサートする
flat_gradsを返す
このセクションの右側の GIF は、ベクトル化された加算の概念を示しています。

ピマップ
以下のコードは、行列乗算における pmap関数の並列化を示しています。
# pmap と random を JAX からインポートします。JAX NumPy をインポートします。
jax import pmapからランダム
jax.numpyを jnpとして インポートする
# デバイスごとに 5000 x 6000 の次元のランダム行列を 2 つ生成します
random_keys = random.split ( random.PRNGKey ( 0 ), 2 )ランダムキーを分割します。
行列 = pmap ( lambda key : random . normal ( key , ( 5000 , 6000 )))( random_keys )
# データ転送なしで、各CPU/GPUでローカル行列乗算を並列に実行します
出力 = pmap ( lambda x : jnp . dot ( x , x . T ))(行列)
# データ転送なしで、並列に、各CPU/GPUで両方の行列の平均を個別に取得します
平均 = pmap ( jnp .平均)(出力)
印刷する(手段)
最後の行には値が印刷されます。
[1.1566595 1.1805978]
参照
外部リンク
- ドキュメントː jax.readthedocs.io
- Colab ( Jupyter /iPython) クイックスタート ガイドː colab.research.google.com/github/google/jax/blob/main/docs/notebooks/quickstart.ipynb
- TensorFlowの XLAː www.tensorflow.org/xla (高速線形代数)
- YouTube TensorFlow チャンネル「JAX 入門: 機械学習研究の加速」: www.youtube.com/watch?v=WdTeDXsOSj4
- 原著論文ː mlsys.org/Conferences/doc/2018/146.pdf
参考文献
- ^ ab Bradbury, James; Frostig, Roy; Hawkins, Peter; Johnson, Matthew James; Leary, Chris; MacLaurin, Dougal; Necula, George; Paszke, Adam; Vanderplas, Jake; Wanderman-Milne, Skye; Zhang, Qiao (2022-06-18)、「JAX: Autograd and XLA」、Astrophysics Source Code Library、Google、Bibcode :2021ascl.soft11002B、2022-06-18にオリジナルからアーカイブ、 2022-06-18に取得
- ^ Frostig, Roy; Johnson, Matthew James; Leary, Chris (2018-02-02). 「高レベルトレースによる機械学習プログラムのコンパイル」(PDF)。MLsys : 1–3。2022-06-21時点のオリジナルからのアーカイブ(PDF) 。
{{cite journal}}: CS1 メンテナンス: 日付と年 (リンク) - ^ 「JAX を使用して研究を加速する」www.deepmind.com。 2022 年 6 月 18 日時点のオリジナルよりアーカイブ。 2022 年 6 月 18 日閲覧。
- ^ リンリー、マシュー。「Googleは、支配に向けた最後の大きな取り組みがMetaによって影を潜めた後、静かにAI製品戦略のバックボーンを交換している」。Business Insider。2022年6月21日時点のオリジナルよりアーカイブ。2022年6月21日閲覧。
- ^ 「なぜ Google の JAX はこんなに人気があるのか?」Analytics India Magazine . 2022-04-25. 2022-06-18 時点のオリジナルよりアーカイブ。 2022-06-18に閲覧。
