Run Ray on TPU, Part 2: Ray AI libraries
GoogleはRayのTPU対応について、Ray Serve、Ray Data、Ray Trainで推論・データ処理・分散学習を行う方法を公開しました。TPUスライスの配置をtopology指定で扱い、vLLMによるLLM提供とJAX学習を実行できます。
Ray Serveではaccelerator_config.topologyを指定し、複数ホストにまたがるモデルを同じTPUスライスへ配置してvLLMで提供できます。指定しない場合、ワーカーが別スライスに分散してデプロイが停止するため、実運用では重要な設定です。
Ray Dataのiter_jax_batches()は、デバイス分割済みのJAX配列を渡し、ホスト側コピーによる入力ボトルネックを抑えます。
Ray TrainのJaxTrainerは、TPUトポロジーを指定して分散学習、チェックポイント、障害回復、複数スライスへの拡張を扱えます。JAXは各ワーカー関数内でimportする必要があります。
Ray 2.55以降のTPU対応を前提とする開発者向け解説で、モデル/APIの新規提供や利用料金変更ではありません。
Google Cloud TPUで大規模推論・後学習・バッチ推論を行うチームは、手作業の配置制御を減らしつつ、RayService、公式TPUイメージ、ダッシュボード監視を組み合わせられます。