Google опубликовала практическую инструкцию по работе Ray Serve, Ray Data и Ray Train на TPU. Для команды, уже собравшей инфраструктуру вокруг Ray, это важнее очередного бенчмарка ускорителей: TPU теперь можно проверить на своей нагрузке, не заменяя весь orchestration-слой. Но обещание «просто поменять GPU на TPU» работает лишь там, где стек уже совместим с JAX или vLLM.
Что именно заработало
Главная особенность TPU в том, что чипы объединены в фиксированные группы, slices. Хосты внутри slice связаны через ICI, а между разными slices такой связи нет. Поэтому workers одной multi-host модели должны попасть в целую группу. Иначе, пишет Google Developers Blog, задача может не упасть с ошибкой, а навсегда остаться в состоянии DEPLOYING, продолжая расходовать TPU-hours.
Нижний уровень этой механики закрывают GKE с Ray Operator и примитив Ray Core slice_placement_group(). Разработчик задаёт topology, например 4x4 для 16 чипов, после чего Ray резервирует целый slice. Поверх этого работают три библиотеки.
Ray Serve сохраняет autoscaling, балансировку нагрузки и композицию нескольких моделей. LLM на TPU он обслуживает через vLLM. Для модели, которая не помещается на одном хосте, нужно добавить поле topology: тогда replica сама создаёт slice placement group и удерживает workers на общей ICI-сети. Google рекомендует разворачивать production-нагрузку через RayService, а не raw RayCluster. В официальных GKE-инструкциях есть примеры для Llama 3 8B и Mistral 7B на v5e, Llama 3.1 70B на v6e, а также Stable Diffusion.
Ray Data получил iter_jax_batches(). Метод отдаёт batches уже как JAX arrays и сразу распределяет их по устройствам, без промежуточного копирования NumPy-to-JAX на host. Для последнего неполного batch можно явно выбрать drop, pad или raise. Это применимо и как вход для обучения через JaxTrainer, и для offline batch inference.
Ray Train теперь предлагает JaxTrainer с checkpointing, fault-tolerant restarts и масштабированием на несколько slices. Пользователь передаёт training function и форму slice, а Ray запускает по worker на host и собирает их в mesh. По описанию Google, рядом с GPU JaxTrainer или TorchTrainer ключевые изменения сводятся к use_tpu=True и указанию topology вместо числа GPU.
Есть и менее заметные, но практичные детали. Ray публикует официальные образы rayproject/ray:*-tpu с JAX/TPU-стеком, включая flax, optax и orbax-checkpoint. Ray Dashboard показывает загрузку и память TPU рядом с CPU и GPU, а JAX profiler можно подключать к отдельным workers.
Где заканчивается бесшовная миграция
Ray действительно убирает необходимость строить отдельный orchestration-стек для TPU. Та же система управляет размещением, serving, данными, checkpointing и восстановлением задач. Для инфраструктурной команды это означает, что эксперимент можно ограничить новым кластером в GKE, TPU-образом и изменениями в конфигурации пайплайна, а не начинать ещё одну платформенную стройку.
Но совместимость здесь не универсальная. Serve опирается на vLLM, Data предлагает специальный путь для JAX, а Train переносит на TPU именно JAX-задачи через JaxTrainer. Более того, Google отдельно предупреждает: jax нужно импортировать внутри train_loop_per_worker, потому что каждый worker инициализирует его в собственном TPU-контексте. Импорт на уровне модуля приводит к ошибкам инициализации устройств.
Поэтому TPU стала практической альтернативой GPU не для любой команды на Ray, а прежде всего для тех, у кого модель уже обслуживается через vLLM или обучение можно выразить через JAX. Если основная ценность пайплайна зашита в другом framework и GPU-специфичном коде, Ray сохранит оркестрацию, но не отменит переписывание вычислительной части. Ускоритель поменялся, физика не подписывала договор об обратной совместимости.
Что проверять сейчас
Разумный следующий шаг для команды на Ray состоит не в выборе TPU «на будущее», а в коротком сравнительном контуре на одной реальной нагрузке. Маркеров три: запускается ли multi-host модель с одним полем topology, держит ли iter_jax_batches() ускоритель загруженным и переживает ли JaxTrainer перезапуск из checkpoint без ручной координации.
Если эти три проверки проходят, а стоимость и производительность на вашей задаче лучше текущего GPU-контура, TPU можно обсуждать как рабочую альтернативу. Если миграция упирается в замену framework или переписывание training loop, новость остаётся полезной, но уже не про дешёвый перенос. Она про то, что Ray сохранился, а всё остальное ещё предстоит посчитать.
