A Regra Que Governa Tudo

Olha só isso: se você já subiu um modelo multi-host em TPU e ficou olhando o deployment travado em DEPLOYING queimando TPU-horas, você sabe a dor. A causa quase sempre é a mesma — os chips de TPU vivem em grupos fixos chamados slices, conectados por um link de alta velocidade (ICI). Um modelo multi-host precisa cair numa slice inteira, senão os workers não se enxergam e o primeiro collective nunca termina.

A Parte 1 cobriu a base — GKE com o add-on Ray Operator provisiona slices, e o slice_placement_group() do Ray Core reserva a slice inteira de uma vez. Esta parte é sobre as três libs que você vai importar de verdade: Ray Serve, Ray Data e Ray Train (via JaxTrainer).

Todas seguem o mesmo padrão: declare uma topologia, deixe o Core cuidar do placement. O que muda é só em cima do quê você declara.

Se você quer entender melhor como collectives se comportam entre workers distribuídos, dá uma olhada nesse guia de comunicação distribuída no PyTorch. Vale muito a pena!

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

Ray Serve: Uma Linha de YAML Que Quebra Tudo

Serving é onde a maioria dos times começa. Um modelo que precisa de várias GPUs pra caber muitas vezes roda tranquilo num único host TPU, e TPUs costumam ser mais disponíveis e baratas pra inferência. O Ray Serve te dá autoscaling, load-balancing e composição multi-modelo — e no TPU ele serve LLMs via vLLM.

O caso difícil é quando o modelo não cabe num host só (tipo, tensor-parallel espalhado em 16 chips). É aí que o topology salva:

accelerator_type: TPU-V6E
accelerator_config:
  kind: tpu
  topology: "4x4"

Com topology setado, o backend TPU do Serve pula o placement group inicial e delega pra réplica, que cria um slice placement group no startup. Essa delegação é o que mantém os workers de um modelo tensor-parallel numa única malha ICI compartilhada.

Sem isso, o Serve cai no fallback de bundles por chip. Num modelo multi-host, esses bundles podem se espalhar por duas slices — e como não tem ICI entre slices, os workers nunca terminam o primeiro collective. Sem crash. Sem erro. Só um deployment preso em DEPLOYING pra sempre.

Deployando um RayService

# Recomendado pra produção em vez de um RayCluster cru
# Depois de deployar numa imagem vLLM TPU publicada:
# kubectl get rayservice -w   # espera chegar em Running
# curl http://<endpoint>/v1/completions -d '{...}'

Os tutoriais oficiais do GKE cobrem Llama 3 8B e Mistral 7B no v5e, Llama 3.1 70B no v6e, e Stable Diffusion.

Cloud dashboard showing Ray Data iter_jax_batches pipeline streaming device-sharded JAX arrays to TPU slice Dev Environment Setup

Ray Data: Pare de Deixar o Loader Virar Gargalo

Um acelerador rápido só é útil se você consegue manter dados fluindo nele. TPUs são rápidas o suficiente pra que um loader ingênuo vire o gargalo. É isso que o iter_jax_batches() resolve — ele te entrega lotes já como arrays JAX e já device-sharded.

ds = ray.data.read_parquet("gs://meu-bucket/train/")
for batch in ds.iter_jax_batches(batch_size=1024):
    # o batch chega como arrays JAX device-sharded, pronto pro step de treino
    loss = train_step(batch)

A API cuida do batch final incompleto (aquele que não é múltiplo limpo do seu batch size) com uma escolha explícita entre drop, pad ou raise — em vez de um shape error três horas depois do início.

Use como lado de input de um job JaxTrainer, ou sozinho pra batch inference offline em cima de um dataset grande numa slice TPU.

JaxTrainer: Treino Distribuído com Um Config Só

Treino costumava ser a parte confusa do Ray no TPU. O JaxTrainer resolve isso — ele traz o training loop do Ray Train (checkpointing, fault tolerance, scale-out multi-slice) pro JAX.

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

def train_loop_per_worker(config):
    import jax  # importa jax DENTRO da worker fn (requisito do TPU)
    # ... seu step de treino JAX/Flax roda aqui, uma vez por host ...

trainer = JaxTrainer(
    train_loop_per_worker=train_loop_per_worker,
    scaling_config=ScalingConfig(
        use_tpu=True,
        topology="4x4",           # o formato da slice, NÃO uma contagem de chips
        accelerator_type="TPU-V6E",
    ),
)
trainer.fit()

Duas pegadinhas que valem memorizar:

  • import jax precisa ficar dentro do train_loop_per_worker, não no topo do arquivo. Cada worker inicializa o JAX no seu próprio contexto TPU; importar no module scope te dá erros crípticos de device-init antes do primeiro step.
  • topology="4x4" é toda a declaração de placement. Ao lado de um JaxTrainer ou TorchTrainer de GPU, a única diferença real é use_tpu=True e uma topologia no lugar da contagem de GPUs.

Como o Ray Train é dono do loop, você ganha checkpointing e restarts tolerantes a falha — é isso que faz runs longos em capacidade preemptible realmente terminarem.

Tooling: Imagens Prontas e Dashboard

O Ray agora publica imagens oficiais rayproject/ray:*-tpu com a stack JAX/TPU (jax[tpu], flax, optax, orbax-checkpoint) e ferramentas de profiling já instaladas. Baseie sua imagem na tag -tpu e esqueça a montagem manual do ambiente.

O Ray Dashboard agora mostra utilização e memória de TPU ao lado de CPU e GPU na aba Cluster, com ray.util.tpu.init_jax_profiler() expondo um profiler JAX por worker que o dashboard consegue anexar.

⚠️ Fica de Olho Nisso

  • Topology não é contagem de chips. "4x4" significa 4 hosts × 4 chips = 16 chips, não 16 hosts.
  • TPU preemptible vai reclamar sua slice. Sem o checkpointing do Ray Train, horas de trabalho evaporam.
  • Coordenação cross-slice é do Ray, mas só quando você realmente precisa de multi-slice — uma topologia single-slice não escala sozinha.

🧭 Próximos Passos

  1. Clone o repo kubernetes-engine-samples e rode os steps de serve/data/train (Qwen3-4B numa slice v6e).
  2. Ou só habilite --enable-ray-operator num cluster e rode uma task Ray numa slice pequena.
  3. Depois dá uma lida no novo modelo de governança da React Foundation — um belo case de como decisões de ecossistema repercutem nas escolhas de ferramentas.

Você não precisa virar especialista em TPU pra usar uma. Só dá uma chance. 근거자료: Google Developers Blog — Run Ray on TPU Part 2

JaxTrainer ScalingConfig code snippet displayed on laptop next to TPU utilization metrics in Ray Dashboard Developer Related Image

A Moral da História

Ray no TPU não é um framework diferente — é o mesmo Ray que você já roda em GPU, com um campo extra (topology) e uma flag extra (use_tpu=True). As libs fazem a matemática do placement pra você; seu trabalho é declarar o formato da slice e deixar o Core reservar atomicamente.

Se você estava evitando TPU porque parecia outro universo, esse é o momento de reconsiderar. O gap de tooling está fechando rápido, e o argumento de custo/disponibilidade pra inferência está cada vez mais difícil de ignorar.

Roadmap pra ficar de olho: integração mais profunda de Ray Data + Ray LLM no TPU, SkyRL em TPU multi-host pra RL e post-training, e suporte dinâmico a super/sub-slice.

Bora construir! 💪

Este conteúdo foi elaborado com o auxílio de ferramentas de IA, com base em fontes confiáveis, e revisado pela nossa equipe editorial antes da publicação. Não substitui o aconselhamento de um profissional especializado.