The One Rule That Governs Everything
If you've ever launched a multi-host model on TPU and watched the deployment sit in DEPLOYING forever while burning TPU-hours, you already know the pain. The root cause is almost always the same: TPU chips live in fixed groups called slices, wired together through a high-speed interconnect (ICI). A multi-host model has to land on one intact slice, or its workers can't reach each other and the first collective never finishes.
Part 1 covered the foundation — GKE with the Ray Operator add-on provisions slices, and Ray Core's slice_placement_group() reserves a whole slice atomically. This part is about the three libraries you'll actually import: Ray Serve, Ray Data, and Ray Train (via JaxTrainer).
Every library follows the same pattern: declare a topology, let Core handle placement. What changes is only what you declare it on.
For a deeper look at how collectives behave across distributed workers, the PyTorch distributed communication deep dive is a solid companion read.

Ray Serve: One Missing YAML Line Breaks Everything
Serving is where most teams start. A model that needs several GPUs to fit often runs fine on a single TPU host, and TPUs are frequently more available and cheaper for inference. Ray Serve gives you autoscaling, load-balancing, and multi-model composition — and on TPU it serves LLMs through vLLM.
The hard case is a model too big for one host (say, tensor-parallel sharded across 16 chips). That's where topology saves you:
accelerator_type: TPU-V6E
accelerator_config:
kind: tpu
topology: "4x4"
With topology set, Serve's TPU backend skips its usual upfront placement group and defers to the replica, which creates a slice placement group at startup. That deferral is what keeps a tensor-parallel model's workers on one shared ICI mesh.
Leave it off and Serve falls back to per-chip bundles. On a multi-host model those bundles can scatter across two slices — and since there's no ICI between slices, workers never finish their first collective. No crash. No error. Just a deployment stuck in DEPLOYING forever.
Deploying a RayService
# Recommended for production over a raw RayCluster
# After deploying on a published vLLM TPU image:
# kubectl get rayservice -w # wait for Running
# curl http://<endpoint>/v1/completions -d '{...}'
Official GKE tutorials cover Llama 3 8B and Mistral 7B on v5e, Llama 3.1 70B on v6e, and Stable Diffusion.

Ray Data: Stop Letting Your Loader Bottleneck the TPU
A fast accelerator is only as useful as the data you can keep flowing into it. TPUs are fast enough that a naive loader becomes the bottleneck. That's what iter_jax_batches() solves — it hands you batches already as JAX arrays and already device-sharded.
ds = ray.data.read_parquet("gs://my-bucket/train/")
for batch in ds.iter_jax_batches(batch_size=1024):
# batch arrives as device-sharded JAX arrays, ready for the training step
loss = train_step(batch)
The API handles the ragged final batch (the one that isn't a clean multiple of your batch size) with an explicit choice of drop, pad, or raise — instead of a shape error three hours into a run.
Use it as the input side of a JaxTrainer job, or standalone for offline batch inference over a big dataset on a TPU slice.
JaxTrainer: Distributed Training From One Config
Training used to be the confusing part of Ray on TPU. JaxTrainer fixes that — it brings Ray Train's training loop (checkpointing, fault tolerance, multi-slice scale-out) to JAX.
from ray.train import ScalingConfig
from ray.train.v2.jax import JaxTrainer
def train_loop_per_worker(config):
import jax # import jax INSIDE the worker fn (TPU requirement)
# ... your JAX/Flax training step runs here, once per host ...
trainer = JaxTrainer(
train_loop_per_worker=train_loop_per_worker,
scaling_config=ScalingConfig(
use_tpu=True,
topology="4x4", # the slice shape, NOT a chip count
accelerator_type="TPU-V6E",
),
)
trainer.fit()
Two gotchas worth memorizing:
import jaxmust live insidetrain_loop_per_worker, not at module scope. Each worker initializes JAX in its own TPU context; importing at module scope gets you cryptic device-init errors before the first step.topology="4x4"is the entire placement declaration. Set next to a GPUJaxTrainerorTorchTrainer, the only real difference isuse_tpu=Trueand a topology instead of a GPU count.
Because Ray Train owns the loop, you get checkpointing and fault-tolerant restarts — which is what makes long TPU runs on preemptible capacity actually finish.
Tooling: Prebuilt Images and Dashboard
Ray now publishes official rayproject/ray:*-tpu images with the JAX/TPU stack (jax[tpu], flax, optax, orbax-checkpoint) and profiling tooling preinstalled. Base your image on the tagged -tpu one and skip the environment assembly.
The Ray Dashboard now shows TPU utilization and memory next to CPU and GPU on the Cluster tab, with ray.util.tpu.init_jax_profiler() exposing a per-worker JAX profiler the dashboard can attach to.
⚠️ Watch Out For
- Topology is not a chip count.
"4x4"means 4 hosts × 4 chips = 16 chips, not 16 hosts. - Preemptible TPUs will reclaim your slice. Without Ray Train's checkpointing, hours of work evaporate.
- Cross-slice coordination is handled by Ray, but only when you actually need multi-slice — a single-slice topology won't magically scale.
🧭 Next Steps
- Clone the
kubernetes-engine-samplesrepo and run the serve/data/train steps (Qwen3-4B on a v6e slice). - Or just enable
--enable-ray-operatoron a cluster and run one Ray task on a small slice. - Then read up on the React Foundation's new governance model — a good case study in how ecosystem-level decisions ripple through tooling choices.
You don't have to become a TPU expert to use one. Just give it a try. 근거자료: Google Developers Blog — Run Ray on TPU Part 2

The Takeaway
Ray on TPU isn't a different framework — it's the same Ray you already run on GPUs, with one extra field (topology) and one extra flag (use_tpu=True). The libraries do the placement math for you; your job is to declare the slice shape and let Core reserve it atomically.
If you've been avoiding TPU because it felt like a different universe, this is the moment to reconsider. The tooling gap is closing fast, and the cost/availability argument for inference is getting harder to ignore.
Roadmap to watch: deeper Ray Data + Ray LLM TPU integration, SkyRL on multi-host TPU for RL and post-training, and dynamic super/sub-slice support.
Happy building.