Stage 5 · Ray Train 분산 학습 — 데이터 샤딩 · 체크포인트
한눈에 보기
Stage 3에서 데이터를 샤딩하는 법을 배웠다면, 이 단계는 학습 자체를 워커에 분산하는 Ray Train입니다. 핵심 통찰은 — 각 워커가 데이터의 서로 다른 “샤드”를 받고, 자기 샤드로 모델을 학습한다 는 것입니다. 노트북에서 ScalingConfig(num_workers=N)으로 워커 수를 정하면, N개 파이썬 프로세스가 한 머신에서 분산 학습을 합니다.
이번 포스트에서 다루는 핵심 질문은 세 가지입니다.
- TorchTrainer와 ScalingConfig로 분산 학습을 시작하는 방법은?
- 데이터가 워커별로 샤딩되는 원리(
get_dataset_shard/DistributedSampler)와 전역 배치 계산은? - 체크포인트 저장·복원과 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 3의 Dataset을 학습에 직접 넘기면, 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샤드를 반환.- ⚠️
DistributedSampler는IterableDataset(무한 반복)과 안 어울립니다 — 그럴 땐 Ray Data를 쓰세요.
체크포인트와 복원
학습 상태(모델 가중치·옵티마이저)를 저장·복원하려면 RunConfig와 train.report의 checkpoint를 씁니다.
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 Train은
ScalingConfig(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)
- Stage 6 · Ray Serve 서빙 — 학습한 모델을 로컬에서 HTTP 엔드포인트로 배포합니다.
- Stage 3 · Ray Data 파이프라인 — 학습용 데이터 준비와 샤딩 기초.
- Stage 4 · Ray Tune 튜닝 —
Trainer를Tuner로 감싸 하이퍼파라미터를 탐색.