はじめに:Part 1 の続き

Part 1 では、TPU 上で Ray を動かす際に覚えておくべき唯一の注意点を確認しました。すなわち TPU チップは slice(スライス)と呼ばれる固定グループに配線されており、マルチホストモデルは必ず1つの完全なスライス上に配置されなければならない という点です。スライス間には ICI(Inter-Chip Interconnect)が存在しないため、ワーカー同士が通信できずジョブがハングします。

GKE の Ray Operator アドオンがスライスをプロビジョニングしホストにラベルを付与し、Ray Core の slice_placement_group() がスライス全体を一括予約します。本 Part 2 では、その上に構築する 実務向けライブラリ3種 — Ray Serve、Ray Data、Ray Train(JaxTrainer) — を扱います。

根拠資料: Run Ray on TPU, Part 2: Ray AI Libraries

共通パターンは1つです。topology を宣言し、スライス予約は Core に委ねる。 ライブラリごとに変わるのは「何に対して宣言するか」だけです。

Developer configuring Ray Serve topology field on GKE cluster to gang-schedule multi-host TPU model Dev Environment Setup

Ray Serve:マルチホストモデルを1スライスに gang-schedule する

多くのチームはここから始めます。1ホストに収まらないモデル(例:16チップに tensor-parallel でシャーディングされたモデル)をサービングする際、Serve は たった一行 で問題を解決します。

accelerator_type: TPU-V6E
accelerator_config:
  kind: tpu
  topology: "4x4"   # 16チップスライスの形状

この一行が重要な理由は、topology の記載漏れがマルチホスト TPU 障害の典型パターンだからです。topology を設定すると、Serve の TPU バックエンドは通常の upfront placement group をスキップし、レプリカ起動時に slice placement group を直接生成します。この「委譲」が、tensor-parallel モデルのワーカーを同一 ICI メッシュ上に保持する仕組みです。

記載漏れの場合、Serve は per-chip bundle にフォールバックしますが、マルチホストモデルではこれらの bundle が2つのスライスに分散し得ます。スライス間に ICI はないため、ワーカーは最初の collective で永久に完了できません。クラッシュではなく、DEPLOYING 状態で停止し続け、TPU-hour を消費しながら存在しないバグを探すことになります。 真の原因は YAML 一行の欠落です。

実務では raw RayCluster より RayService を推奨します。公開されている vLLM TPU イメージ上に RayService をデプロイし、Running になるまで待機してから curl でエンドポイントを叩きます。

# RayService の状態確認(Running になるまで待機)
kubectl get rayservice -w

# エンドポイント呼び出しテスト
curl -X POST http://<service-endpoint>/v1/chat/completions \
  -H "Content-Type: application/json" \
  -d '{"model": "llama-3-8b", "messages": [{"role": "user", "content": "hi"}]}'

公式 GKE チュートリアルでは v5e の Llama 3 8B / Mistral 7B、v6e の Llama 3.1 70B、Stable Diffusion までカバーされています。

Ray Data:iter_jax_batches() でデータパイプラインのボトルネックを除去

高速アクセラレータもデータが流れなければ無意味です。TPU は naive loader がボトルネックになるほど高速です。これを解決するのが iter_jax_batches() です。

ds = ray.data.read_parquet("gs://my-bucket/train/")
for batch in ds.iter_jax_batches(batch_size=1024):
    # batch は device-sharded JAX 配列として到着、そのまま学習ステップへ
    loss = train_step(batch)

この API の利点:

  • device sharding を自動処理 — ホスト側の NumPy→JAX コピーでステップが停止しない
  • ragged final batch(batch_size の倍数にならない最終バッチ)を drop/pad/raise で明示的に処理 — 3時間後に shape エラーで落ちる事態を回避

JaxTrainer の入力としても、大規模データセットに対するオフライン一括推論としても利用できます。

Cloud architect diagram showing Ray Data iter_jax_batches pipeline feeding device-sharded JAX arrays into TPU slice System Abstract Visual

JaxTrainer:学習ループを Ray に委譲する

以前は学習が最も混乱する部分でした。topology も考慮し、スライス形状をコードに反映する必要がありました。JaxTrainer がこれを整理します。

from ray.train import ScalingConfig
from ray.train.v2.jax import JaxTrainer

def train_loop_per_worker(config):
    import jax  # ワーカー関数内で import(TPU 要件)
    # ... JAX/Flax 学習ステップ、ホストごとに1回実行 ...

trainer = JaxTrainer(
    train_loop_per_worker=train_loop_per_worker,
    scaling_config=ScalingConfig(
        use_tpu=True,
        topology="4x4",  # スライス形状(チップ数ではない!)
        accelerator_type="TPU-V6E",
    ),
)
trainer.fit()

デバッグ時間を節約するポイント2点:

  1. import jax は train_loop_per_worker の内部に配置 — 各ワーカーが自身の TPU コンテキストで JAX を初期化するため。モジュールスコープに置くと、最初のステップすら到達できず cryptic な device-init エラーと格闘することになります。
  2. topology="4x4" が placement 宣言の全て — 以前は手書きしていた coordination コードブロックが、この一行に置き換わります。

GPU 用の TorchTrainer や JaxTrainer と比較した場合、実際の違いは use_tpu=True と GPU 数ではなく topology を使用する点のみです。残りは Ray が全て実行します。

注意事項と深掘り Tips

  • topology ≠ チップ数:topology="4x4" はスライス形状であり、16個を意味するわけではありません。混同すると上述の DEPLOYING 地獄を再体験します。
  • preemptible capacity の活用:Ray Train が学習ループを所有するため、checkpointing と fault-tolerant restart が付属します。プリエンプティブ TPU で長時間学習を「実際に完走させる」のはこの組み合わせです。
  • マルチスライス拡張:1スライスで不足する場合、topology はマルチスライスに拡張されます。cross-slice coordination は Ray が自動で接続します。
  • 公式イメージの活用:rayproject/ray:*-tpu イメージには jax[tpu]、flax、optax、orbax-checkpoint、プロファイリングツールが同梱されています。TPU 環境を手組みせず、これをベースにしてください。
  • モニタリング:Ray Dashboard の Cluster タブで CPU/GPU の隣に TPU utilization とメモリが表示されます。ray.util.tpu.init_jax_profiler() でワーカー別 JAX プロファイラも接続可能です。

日本開発コミュニティにおける適用文脈

日本では TPU を本番運用する事例はまだ限定的ですが、国内の研究機関・大学の大規模 ML 基盤や、GPU 調達難に直面するスタートアップ において本パターンは有効です。特に Qiita などで散見される「TPU で Ray を動かしたが DEPLOYING から進まない」という相談の多くは topology 記載漏れに起因します。チーム内の YAML テンプレートに topology を必須フィールドとして組み込む運用を推奨します。

次の学習ステップ

  1. get-started サンプルのクローン — Qwen3-4B を v6e スライスで動かす serve/data/train ステップがワーキングコードで提供されています。
  2. --enable-ray-operator でクラスタを1つ立ち上げ、小さいスライスに Ray task を1つ投げてみる。TPU 専門家になる必要はありません、まず動かすことが先決です。
  3. ロードマップの注視:Ray Data / Ray LLM の TPU 統合深化、マルチホスト TPU 上の SkyRL(強化学習・post-training)、動的 super/sub-slice サポートが予定されています。

ML engineer running JaxTrainer distributed training loop on TPU v6e slice with checkpointing dashboard Development Concept Image

まとめ

Part 1 が「なぜスライスが全てなのか」を扱ったのに対し、Part 2 は そのスライス上に実際の AI ライブラリを載せる方法 を整理しました。

  • Ray Serve — accelerator_config.topology 一フィールドでマルチホストモデルを1スライスに gang-schedule
  • Ray Data — iter_jax_batches() で JAX-native バッチをスライスに直結
  • JaxTrainer — ScalingConfig 一つで分散学習ループを実行

GPU で使っていたあの Ray を、TPU でもそのまま使えます。次のアクションは明確です。クラスタを1つ立ち上げ、serve/data/train のいずれかを選んで動かしてみてください。

あわせて読みたい

本コンテンツは、信頼性の高い情報源をもとにAIツールを活用して作成され、編集者によるレビューを経て公開されています。専門家によるアドバイスの代替となるものではありません。