들어가며: 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) — 를 다뤄요.

근거자료: Run Ray on TPU, Part 2: Ray AI Libraries

핵심 패턴은 딱 하나예요. topology를 선언하고, 슬라이스 예약은 Core에 맡긴다. 라이브러리마다 달라지는 건 '무엇에 선언하느냐'뿐입니다.

Developer configuring Ray Serve topology field on GKE cluster to gang-schedule multi-host TPU model Algorithm Concept Visual

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의 입력으로도 쓰고, 그냥 큰 데이터셋에 대한 오프라인 배치 추론에도 씁니다.

Cloud architect diagram showing Ray Data iter_jax_batches pipeline feeding device-sharded JAX arrays into TPU slice Coding Session Visual

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()

디버깅 시간을 아껴주는 포인트 두 가지:

  1. import jax는 train_loop_per_worker 안쪽에 — 각 워커가 자기 TPU 컨텍스트에서 JAX를 초기화하기 때문. 모듈 스코프에 두면 첫 스텝도 못 가고 cryptic device-init 에러와 씨름하게 됩니다.
  2. 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 누락이에요.

다음 단계 학습 방향

  1. get-started 예제 클론 — Qwen3-4B를 v6e 슬라이스에서 돌리는 serve/data/train 스텝이 워킹 코드로 제공됩니다.
  2. --enable-ray-operator로 클러스터 하나 띄우고 작은 슬라이스에 Ray task 하나만 던져보기. TPU 전문가가 될 필요는 없어요, 일단 돌려보는 게 먼저입니다.
  3. 로드맵 주시: Ray Data/Ray LLM의 TPU 통합 심화, 멀티호스트 TPU 위 SkyRL(강화학습·post-training), 동적 super/sub-slice 지원이 예정돼 있어요.

ML engineer running JaxTrainer distributed training loop on TPU v6e slice with checkpointing dashboard Software Concept Art

마무리

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 중 하나를 골라 돌려보세요.

함께 보면 좋은 글

본 콘텐츠는 신뢰할 수 있는 출처를 바탕으로 AI 도구를 활용하여 초안이 작성되었으며, 편집자의 검토를 거쳐 발행되었습니다. 전문가의 조언을 대체하지 않습니다.