AI・機械学習
上級

XLA(エックスエルエー)

Google が開発した機械学習専用コンパイラで、TensorFlow や JAX の計算グラフを TPU・GPU・CPU 向けに最適化コンパイルする Accelerated Linear Algebra の略称。

0 回閲覧
0 いいね

XLA とは

XLA(Accelerated Linear Algebra)は Google が開発したドメイン固有コンパイラであり、線形代数演算を中心とする機械学習ワークロードを各種ハードウェア向けに最適化コンパイルする。TensorFlow のデフォルトコンパイラバックエンドとして組み込まれ、JAX では唯一の実行エンジンとして機能する。

Google の TPU 上で LLM を学習・推論する場合、XLA は不可欠な存在である。Gemini・PaLM・T5 といった大規模モデルの学習はすべて XLA を通じて TPU クラスタ上で実行されている。

XLA のコンパイルパイプライン

HLO(High Level Optimizer)

XLA の中間表現である HLO IR は、行列乗算・畳み込み・リダクションなどの高水準演算をノードとする計算グラフである。フレームワーク(TensorFlow, JAX, PyTorch/XLA)からのコードはまず HLO に変換される。

HLO 最適化パス

HLO レベルで適用される主要な最適化は以下の通り。

  1. 演算融合(Fusion): Element-wise 演算やリダクションを上流の演算と融合
  2. レイアウト割り当て(Layout Assignment): テンソルのメモリレイアウトをターゲットデバイスに最適化
  3. バッファ割り当て(Buffer Assignment): 中間テンソルのメモリ割り当てを最適化し、メモリ再利用を最大化
  4. Algebraic Simplification: 数学的等価変換による演算簡約(例: X * 1 → X)
  5. While Loop 最適化: ループ不変コードの外出し、ループアンローリング

StableHLO

XLA の HLO を標準化した StableHLO は、フレームワーク間のポータビリティを向上させる取り組みである。MLIR(Multi-Level Intermediate Representation)上に構築され、異なるコンパイラバックエンド間での互換性を提供する。

コンパイルステージ処理内容最適化例
HLO 生成フレームワークからの変換型推論・形状推論
HLO 最適化グラフレベル最適化融合・定数畳み込み・CSE
レイアウト割り当てメモリ配置決定TPU MXU 対応タイリング
バッファ割り当てメモリ管理Liveness 解析・再利用
コード生成デバイス固有コード出力PTX / TPU 命令 / LLVM IR

TPU における XLA の役割

MXU(Matrix Multiply Unit)最適化

TPU v4/v5 の MXU は 128x128 の行列乗算を 1 サイクルで実行する。XLA はテンソル演算を 128x128 タイルに自動分割し、MXU の利用率を最大化する。パディング挿入やテンソル転置の自動挿入もコンパイラが担当する。

メガコア(Megacore)対応

TPU v4 以降はチップ内に 2 つのコアを持つ Megacore 構成となっている。XLA は計算グラフを自動的に 2 分割し、コア間の通信を最小化する配置を決定する。

ICI(Inter-Chip Interconnect)通信最適化

TPU Pod 内のチップ間通信は ICI 経由で行われる。XLA の SPMD パーティショナが AllReduce・AllGather などの集団通信を最適なトポロジに配置し、通信レイテンシを削減する。

JAX における XLA の活用

JAX は XLA をネイティブの実行エンジンとして使用し、以下の機能を提供する。

jax.jit

Python 関数を XLA コンパイル対象としてマークする。初回呼び出し時にトレース → HLO 変換 → コンパイルが行われ、以降はコンパイル済みカーネルが直接実行される。

jax.pmap / jax.pjit

データ並列・テンソル並列の両方を XLA の SPMD パーティショナ経由で実現する。PartitionSpec でテンソルの分割軸を宣言すると、XLA が自動的に通信プリミティブを挿入する。

jax.grad

自動微分も XLA の HLO レベルで処理される。フォワードパスの HLO からバックワードパスの HLO を自動生成し、両パスを通じた最適化(勾配チェックポインティングの自動挿入など)が適用される。

PyTorch/XLA

PyTorch モデルを XLA バックエンド上で実行するためのブリッジライブラリ。torch_xla.core.xla_model 経由でテンソルを XLA デバイスに配置すると、PyTorch の Eager 実行が XLA の Lazy Tensor 方式に切り替わる。

利点

  • Google Cloud TPU 上で PyTorch モデルを実行可能
  • FSDP(Fully Sharded Data Parallelism)の XLA 対応
  • GSPMD(General SPMD)による自動並列化

制約

  • Eager Mode との挙動差異(Lazy Tensor の同期タイミング)
  • Dynamic Shapes のサポートが限定的
  • torch.compile との統合は OpenXLA プロジェクトで進行中

XLA と他のコンパイラの関係

XLA は Google エコシステムの中核コンパイラとして位置付けられ、StableHLO を通じて外部コンパイラとの相互運用が進んでいる。IREE(Intermediate Representation Execution Environment)は StableHLO を入力として受け取り、モバイル・エッジデバイス向けのコード生成を行う。

よくある質問(FAQ)

Q1: XLA は NVIDIA GPU でも使えますか?

はい、XLA は NVIDIA GPU(CUDA バックエンド)にも対応しています。ただし、NVIDIA GPU 上では TensorRT や Triton(torch.compile 経由)の方が成熟度が高く、一般的には TPU 上での利用が XLA の主戦場です。JAX を NVIDIA GPU で使う場合は自動的に XLA の CUDA バックエンドが使用されます。

Q2: XLA のコンパイル時間が長い場合の対策は?

XLA のコンパイルキャッシュ(XLA_FLAGS=--xla_gpu_persistent_compilation_cache_dir=/path)を有効にすると、同じ計算グラフの再コンパイルを回避できます。また、JAX では jax.jit のトレース結果をキャッシュする AOT(Ahead-of-Time)コンパイルモードも利用可能です。

Q3: 自作 PC で XLA を試すことはできますか?

はい、JAX をインストールすれば NVIDIA GPU 上で XLA を利用できます(pip install jax[cuda12])。ただし、XLA の真価は TPU クラスタでの大規模分散学習にあるため、小規模実験には torch.compile + Triton の方が手軽です。Google Colab の無料 TPU ランタイムで XLA + JAX を体験するのも良い選択肢です。