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.

Developer configuring Ray Serve topology field for multi-host TPU deployment on GKE cluster Algorithm Concept Visual

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.

Cloud dashboard showing Ray Data iter_jax_batches pipeline streaming device-sharded JAX arrays to TPU slice Programming Illustration

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 jax must live inside train_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 GPU JaxTrainer or TorchTrainer, the only real difference is use_tpu=True and 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

  1. Clone the kubernetes-engine-samples repo and run the serve/data/train steps (Qwen3-4B on a v6e slice).
  2. Or just enable --enable-ray-operator on a cluster and run one Ray task on a small slice.
  3. 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

JaxTrainer ScalingConfig code snippet displayed on laptop next to TPU utilization metrics in Ray Dashboard Technical Structure Concept

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.

This content was drafted using AI tools based on reliable sources, and has been reviewed by our editorial team before publication. It is not intended to replace professional advice.