Stage 5 · Ray Train 분산 학습 — 데이터 샤딩 · 체크포인트

Ray Train — 데이터를 워커 수만큼 샤드로 나눠 각 워커가 자기 샤드로 학습, 체크포인트 공유 데이터셋 get_dataset_shard 샤딩 Worker 0 · 샤드 0 · DDP Worker 1 · 샤드 1 · DDP Worker 2 · 샤드 2 · DDP 체크포인트 (공유 저장소) RunConfig · CheckpointConfig 전역 배치 크기 = 워커별 배치 × world_size
Ray Train — 하나의 데이터셋이 **num_workers만큼의 서로 다른 샤드**로 나뉘고, 각 워커는 `get_dataset_shard()`/`DistributedSampler`가 준 자기 샤드만으로 DDP 학습을 진행합니다. **전역 배치 크기 = 워커별 배치 × world_size**이며, 체크포인트는 공유 저장소(로컬 경로)에 기록됩니다.

한눈에 보기

Stage 3에서 데이터를 샤딩하는 법을 배웠다면, 이 단계는 학습 자체를 워커에 분산하는 Ray Train입니다. 핵심 통찰은 — 각 워커가 데이터의 서로 다른 “샤드”를 받고, 자기 샤드로 모델을 학습한다 는 것입니다. 노트북에서 ScalingConfig(num_workers=N)으로 워커 수를 정하면, N개 파이썬 프로세스가 한 머신에서 분산 학습을 합니다.

이번 포스트에서 다루는 핵심 질문은 세 가지입니다.

  1. TorchTrainer와 ScalingConfig로 분산 학습을 시작하는 방법은?
  2. 데이터가 워커별로 샤딩되는 원리(get_dataset_shard/DistributedSampler)와 전역 배치 계산은?
  3. 체크포인트 저장·복원과 Train V2 주의점은?

TorchTrainer 기본

Ray Train은 프레임워크별 Trainer를 제공합니다(PyTorch, Lightning, HuggingFace, TensorFlow, XGBoost 등). 파이토치 예시로 시작합니다.

pip install -U "ray[train]" torch

학습은 각 워커가 실행하는 함수로 작성합니다. ScalingConfig(num_workers=N)이 워커 수를 정합니다.

import torch
import torch.nn as nn
import ray
from ray import train
from ray.train import ScalingConfig
from ray.train.torch import TorchTrainer

def train_func(config):
    model = nn.Linear(4, 2)
    # DDP 래핑: 모델을 GPU/장치로 옮기고 분산 학습 지원
    model = train.torch.prepare_model(model)

    loader = torch.utils.data.DataLoader(
        simple_dataset(), batch_size=32, shuffle=True,
    )
    # 데이터 로더에 DistributedSampler 추가 (각 워커가 다른 샤드)
    loader = train.torch.prepare_data_loader(loader)

    loss_fn = nn.MSELoss()
    opt = torch.optim.SGD(model.parameters(), lr=0.01)

    for epoch in range(10):
        for x, y in loader:
            opt.zero_grad()
            out = model(x)
            loss = loss_fn(out, y)
            loss.backward()
            opt.step()
        train.report({"loss": loss.item()})   # 메트릭 보고

trainer = TorchTrainer(
    train_func,
    scaling_config=ScalingConfig(num_workers=4, use_gpu=False),  # 노트북 CPU 4 워커
)
result = trainer.fit()
print(result.metrics)
  • ScalingConfig(num_workers=N) : 분산 학습 워커 프로세스 수. 로컬에선 N ≤ 논리 코어 수여야 합니다.
  • prepare_model : 모델을 장치로 옮기고 DDP(분산 데이터 병렬)로 래핑.
  • prepare_data_loader : 데이터 로더에 DistributedSampler를 붙여 워커 간 데이터를 샤딩.
  • train.report(metrics) : 메트릭 보고 + 선택적으로 체크포인트 첨부.

데이터 샤딩 — 각 워커는 자기 샤드만

DistributedSampler — 워커별 disjoint 샤드

prepare_data_loader가 사용하는 DistributedSampler는 데이터를 워커 수만큼 서로 안 겹치는 조각으로 나눕니다. 이게 바로 “학습 샤딩” 입니다.

  • 각 워커는 전체 데이터가 아니라 자기 샤드만 반복합니다(겹침 없음).
  • 자원의 맥락은 ray.train.get_context()로 얻습니다.
from ray import train

ctx = train.get_context()
print("세계 크기(world_size):", ctx.get_world_size())
print("내 rank:", ctx.get_world_rank())

전역 배치 크기

batch_size워커별 값입니다. 전체(전역) 배치는:

global_batch_size = worker_batch_size × world_size

예: batch_size=32, num_workers=4 → 전역 배치 128. 하이퍼파라미터를 워커 수와 무관하게 맞추려면 이 공식으로 worker_batch_size를 계산해 쓰는 것이 좋습니다.

Ray Data 연동 — get_dataset_shard

Stage 3Dataset을 학습에 직접 넘기면, Ray가 알아서 워커별 샤드로 나눠줍니다.

from ray.data import from_items

train_ds = from_items([{"x": [i, i], "y": [i]} for i in range(1000)])

def train_func(config):
    shard = train.get_dataset_shard()        # 이 워커가 받은 데이터 샤드
    for batch in shard.iter_torch_batches(batch_size=32):
        x = batch["x"]; y = batch["y"]
        # ... 학습 ...
        break

trainer = TorchTrainer(
    train_func,
    scaling_config=ScalingConfig(num_workers=4),
    datasets={"train": train_ds},
)
  • train.get_dataset_shard() : 이 워커 전용 DataIterator 샤드를 반환.
  • ⚠️ DistributedSamplerIterableDataset(무한 반복)과 안 어울립니다 — 그럴 땐 Ray Data를 쓰세요.

체크포인트와 복원

학습 상태(모델 가중치·옵티마이저)를 저장·복원하려면 RunConfigtrain.reportcheckpoint를 씁니다.

from ray.train import RunConfig

def train_func(config):
    model = train.torch.prepare_model(nn.Linear(4, 2))
    ckpt_dir = "/tmp/my_ckpt"                # 로컬 단일 노드면 로컬 경로 OK
    for step in range(20):
        # ... 학습 ...
        train.report(
            {"loss": value},
            checkpoint=train.Checkpoint.from_directory(ckpt_dir),
        )

trainer = TorchTrainer(
    train_func,
    scaling_config=ScalingConfig(num_workers=4),
    run_config=RunConfig(
        storage_path="/tmp/ray_train_runs",  # 단일 노드: 로컬 경로
        name="my_experiment",
    ),
)
result = trainer.fit()
print(result.checkpoint)    # 최종/최고 체크포인트 경로
print(result.error)         # 오류 시
  • 체크포인트로 복원하려면: train.get_checkpoint() 로 워커 내부에서 다시 읽습니다.
  • Train V2 주의: 최신 Ray(2.43+·3.x)는 Train V2 API가 기본화되고 있습니다(구 V1과 RAY_TRAIN_V2_ENABLED=1로 전환). 설치한 Ray 버전의 API 시그니처를 확인하세요. 스토리지 경로는 다중노드에선 공유 스토리지가 필요하지만, 단일 노드에선 로컬 경로로 충분합니다.

로컬 리소스 주의: num_workers를 코어 수보다 크게 잡으면 워커들이 서로 CPU를 다퉈 오히려 느려집니다. OMP_NUM_THREADS를 낮춰 워커마다 스레드 풀이 폭주하는 것도 막아야 합니다(Stage 1 참고).

Summary

  • Ray TrainScalingConfig(num_workers=N)으로 모델을 분산 학습합니다. 노트북에선 N ≤ 코어 수.
  • 데이터 샤딩: prepare_data_loader/DistributedSampler 또는 Ray Data의 get_dataset_shard()가 각 워커에 겹치지 않는 데이터 샤드를 줍니다.
  • 전역 배치 = 워커별 배치 × world_size.
  • 체크포인트: train.report(..., checkpoint=...) + RunConfig(storage_path=...)로 저장·복원하며, 단일 노드면 로컬 경로로 충분. Train V2 기본화에 주의.

다음 학습 (Next Learning)