JAX
Overview
- JAX / Flax 入門 Google Cloud Japan
- JAX とは
- Google Brain の数値計算ライブラリ
- Numpy と同じようなインターフェース
- 違いは、CPU 以外での計算をサポートしてる
- 独自機能
- JIT コンパイルで早い
- 自動微分:勾配ベクトル計算しやすい
- 基本的に JAX は、上位ライブラリと組み合わせての利用
- Flax, Optax など
- jax 版の numpy module がある
- Flax は Keras みたいなもの (nn.Module)
- 違いは、条件分岐を記述可能
- Keras だと変なことしなきゃいけないっぽい -> Flax 特有知識があんまり必要ないっぽい設計らしい
- Jax+Flax は pytorch っぽい感じで blackbox ではない
- 誤差関数も学習ステップも Jax は自分で実装しなきゃいけない
- ここで出てくるのは optax (cross entropy とか optimizer が実装されている)
- Jax+Flax は model の param は python の dict として外部に用意されている
- JAX とは
- 今こそはじめるJAX/Flax入門 Part 1
- 2019年に登場
- 特長
- 自動微分のサポート: 深層学習の基礎となる順方向/逆方向の勾配ベースの演算をネイティブにサポートしています。grad, hessian, jacfwd, jacrevなどの関数変換を用いて簡単に実現できます。
- ベクトル化:深層学習では、バッチ全体のLossを計算したり、複数GPUによる分散学習など、単一の処理や関数を多くのデータやデバイスに適用することがよくあります。JAXは任意の関数をvmapを介して並列化したり、単一のデバイスでは大きすぎる処理をpmapにより分散させたりすることが可能です。
- JITコンパイル:JAXはXLA(Accelerated Linear Algebra)に基づくJIT(Just-In-Time)コンパイルによる高速化が可能です。XLAは線形代数のためのドメイン固有のコンパイラで、計算グラフを最適化し、効率的な実行を目指すもので、特にGoogle TPU(Tensor Processing Unit)にやNVIDIA GPU向けに最適化されています。
- grad の使い方
import jax.numpy as jnp
from jax import grad
def tanh(x):
y = jnp.exp(-2.0 * x)
return (1.0 - y) / (1.0 + y)
grad_tanh = grad(tanh)
print(grad_tanh(1.0)) # 0.4199743
print(grad(grad(grad(tanh)))(1.0)) # 0.62162673
-
vmap()やpmap() - PDF 機械学習で楽しむ JAX/NumPyro v0.3
- JAX は google が開発中の自動微分と XLA(accelerated Linear Algebra)
- 自動微分の実装法としては JAX では Jacovian Vector Product (JVP) を用いたものが使用されている
- いろんな plot のコードが紹介されている
- ZOZO NEXT の Sai Htaung Kham さんの記事
- Practical な ユースケース
- JAX と Flax を使用した Gemma のファインチューニング
- Slide JAX: Accelerated Machine Learning Research via Composable Function Transformations in Python (Matt Johnson, Google Brain)
- jax/flaxの思想:オブジェクト指向との違い
- いろんなソースコードがある
- Googleが開発しているJAX用の高性能ニューラルネットワークライブラリFlaxを用いてツイート分類を行うレシピ
- オーソドックスな学習アルゴリズムの実装を紹介
- 拡散生成モデルで学ぶJax/Flaxによる深層学習プログラミング
- 詳しそうな解説
- JAXとFlaxを使って、ナウい機械学習をしたい
Misc
- Google からのライブラリとしては、TensorFlow, Keras もある