들어가며: Part 1 다음 이야기
Part 1에서 우리는 TPU 위에서 Ray를 돌릴 때 딱 하나만 기억하면 된다는 걸 확인했죠. 바로 TPU 칩은 slice(슬라이스)라는 고정 그룹으로 묶여 있고, 멀티호스트 모델은 반드시 하나의 온전한 슬라이스 위에 올라가야 한다는 것입니다. 슬라이스 간에는 ICI(Inter-Chip Interconnect)가 없어서, 워커들이 서로 못 만나면 그냥 job이 멈춰버려요.
GKE의 Ray Operator 애드온이 슬라이스를 프로비저닝하고 호스트에 라벨을 붙여주고, Ray Core의 slice_placement_group()이 슬라이스 전체를 한 번에 예약해줍니다. 이번 Part 2는 그 위에 올라가는 실전 라이브러리 3총사 — Ray Serve, Ray Data, Ray Train(JaxTrainer) — 를 다뤄요.
핵심 패턴은 딱 하나예요. topology를 선언하고, 슬라이스 예약은 Core에 맡긴다. 라이브러리마다 달라지는 건 '무엇에 선언하느냐'뿐입니다.

Ray Serve: 멀티호스트 모델을 한 슬라이스에 gang-schedule 하기
대부분의 팀이 여기서 시작해요. 한 호스트에 안 들어가는 모델(예: 16칩에 tensor-parallel로 샤딩된 모델)을 서빙할 때, Serve는 딱 한 줄로 문제를 해결합니다.
accelerator_type: TPU-V6E
accelerator_config:
kind: tpu
topology: "4x4" # 16칩 슬라이스 형태
이 한 줄이 왜 중요하냐면, topology를 빼먹는 게 멀티호스트 TPU 실패의 전형적인 패턴이기 때문이에요. topology가 설정되면 Serve의 TPU 백엔드는 평소 쓰던 upfront placement group을 건너뛰고, replica가 시작할 때 slice placement group을 직접 만듭니다. 이 "양보"가 바로 tensor-parallel 모델의 워커들을 같은 ICI 메시 위에 붙잡아 두는 장치예요.
빼먹으면? Serve가 per-chip bundle로 폴백하는데, 멀티호스트 모델에서는 이 bundle들이 두 슬라이스에 흩어질 수 있어요. 슬라이스 사이엔 ICI가 없으니 워커들은 첫 collective에서 영원히 못 끝냅니다. 크래시가 나는 게 아니라, DEPLOYING 상태로 영원히 멈춰서 TPU-hour를 태우며 없는 버그를 찾게 되죠. 진짜 원인은 YAML 한 줄 누락입니다.
실무에서는 raw RayCluster보다 RayService를 권장해요. 퍼블리시된 vLLM TPU 이미지 위에 RayService를 올리고, Running 될 때까지 기다린 다음 curl로 엔드포인트를 찔러보면 됩니다.
# RayService 상태 확인 (Running 될 때까지 대기)
kubectl get rayservice -w
# 엔드포인트 호출 테스트
curl -X POST http://<service-endpoint>/v1/chat/completions \
-H "Content-Type: application/json" \
-d '{"model": "llama-3-8b", "messages": [{"role": "user", "content": "hi"}]}'
공식 GKE 튜토리얼에는 v5e의 Llama 3 8B / Mistral 7B, v6e의 Llama 3.1 70B, Stable Diffusion까지 커버돼 있습니다.
Ray Data: iter_jax_batches()로 데이터 파이프라인 병목 제거
빠른 가속기도 데이터가 안 흘러들어가면 무용지물이에요. TPU는 naive loader가 병목이 될 만큼 빠릅니다. 이걸 해결하는 게 iter_jax_batches()예요.
ds = ray.data.read_parquet("gs://my-bucket/train/")
for batch in ds.iter_jax_batches(batch_size=1024):
# batch는 device-sharded JAX 배열로 도착, 바로 학습 스텝에 투입 가능
loss = train_step(batch)
이 API가 좋은 이유:
- device sharding을 자동 처리 — 호스트 쪽 NumPy→JAX 복사로 스텝이 멈추는 일이 없음
- ragged final batch(batch_size의 배수로 안 떨어지는 마지막 배치)를
drop/pad/raise로 명시적 처리 — 3시간 뒤에 shape 에러로 죽는 일 방지
JaxTrainer의 입력으로도 쓰고, 그냥 큰 데이터셋에 대한 오프라인 배치 추론에도 씁니다.

JaxTrainer: 학습 루프를 Ray에 위임하기
예전엔 학습이 제일 헷갈리는 부분이었어요. topology도 챙겨야 하고, 슬라이스 형태를 코드에 반영해야 했거든요. JaxTrainer가 이걸 정리합니다.
from ray.train import ScalingConfig
from ray.train.v2.jax import JaxTrainer
def train_loop_per_worker(config):
import jax # 워커 함수 안에서 import (TPU 요구사항)
# ... JAX/Flax 학습 스텝, 호스트당 한 번 실행 ...
trainer = JaxTrainer(
train_loop_per_worker=train_loop_per_worker,
scaling_config=ScalingConfig(
use_tpu=True,
topology="4x4", # 슬라이스 형태 (칩 개수 아님!)
accelerator_type="TPU-V6E",
),
)
trainer.fit()
디버깅 시간을 아껴주는 포인트 두 가지:
import jax는train_loop_per_worker안쪽에 — 각 워커가 자기 TPU 컨텍스트에서 JAX를 초기화하기 때문. 모듈 스코프에 두면 첫 스텝도 못 가고 cryptic device-init 에러와 씨름하게 됩니다.topology="4x4"가 placement 선언의 전부 — 예전엔 손으로 짜던 coordination 코드 블록이 이 한 줄로 대체됩니다.
GPU용 TorchTrainer나 JaxTrainer와 비교하면, 실제 차이는 use_tpu=True와 GPU 개수 대신 topology를 쓴다는 것뿐이에요. 나머지는 Ray가 다 돌립니다.
주의사항과 심화 팁
- topology ≠ 칩 개수:
topology="4x4"는 슬라이스 형태지, 16개를 뜻하는 게 아닙니다. 헷갈리면 위에서 설명한DEPLOYING지옥을 다시 만나요. - preemptible capacity 활용: Ray Train이 학습 루프를 소유하므로 checkpointing과 fault-tolerant restart가 공짜로 따라옵니다. 선점형 TPU에서 긴 학습을 실제로 "끝내게" 만들어주는 게 이 조합이에요.
- 멀티 슬라이스 확장: 슬라이스 하나로 부족하면 topology가 멀티 슬라이스로 확장됩니다. cross-slice coordination은 Ray가 알아서 연결해요.
- 공식 이미지 활용:
rayproject/ray:*-tpu이미지에jax[tpu], flax, optax, orbax-checkpoint, 프로파일링 툴이 이미 들어있습니다. 직접 TPU 환경 조립하지 말고 이걸 베이스로 쓰세요. - 모니터링: Ray Dashboard의 Cluster 탭에서 CPU/GPU 옆에 TPU utilization과 메모리가 표시됩니다.
ray.util.tpu.init_jax_profiler()로 워커별 JAX 프로파일러도 붙일 수 있어요.
한국 개발 생태계에서의 적용 맥락
국내에서 TPU를 프로덕션에 쓰는 팀은 아직 많지 않지만, 국내 SI/금융권의 온프레미스 ML 플랫폼에서 이 패턴이 특히 유효합니다. GPU 확보가 어려운 상황에서 추론 워크로드를 TPU로 돌리고, 멀티호스트 모델은 반드시 topology를 명시하는 관례를 팀 내 YAML 템플릿으로 못 박아두는 걸 권합니다. "왜 DEPLOYING에서 안 넘어가죠?"라는 문의의 90%는 topology 누락이에요.
다음 단계 학습 방향
- get-started 예제 클론 — Qwen3-4B를 v6e 슬라이스에서 돌리는 serve/data/train 스텝이 워킹 코드로 제공됩니다.
--enable-ray-operator로 클러스터 하나 띄우고 작은 슬라이스에 Ray task 하나만 던져보기. TPU 전문가가 될 필요는 없어요, 일단 돌려보는 게 먼저입니다.- 로드맵 주시: Ray Data/Ray LLM의 TPU 통합 심화, 멀티호스트 TPU 위 SkyRL(강화학습·post-training), 동적 super/sub-slice 지원이 예정돼 있어요.

마무리
Part 1이 "왜 슬라이스가 전부인가"를 다뤘다면, Part 2는 그 슬라이스 위에 실제 AI 라이브러리를 얹는 방법을 정리했습니다.
- Ray Serve —
accelerator_config.topology한 필드로 멀티호스트 모델을 한 슬라이스에 gang-schedule - Ray Data —
iter_jax_batches()로 JAX-native 배치를 슬라이스에 직결 - JaxTrainer —
ScalingConfig하나로 분산 학습 루프 실행
GPU에서 쓰던 그 Ray, 이제 TPU에서도 그대로 씁니다. 다음 액션은 명확해요. 클러스터 하나 띄우고, serve/data/train 중 하나를 골라 돌려보세요.