from pyspark.ml.torch.distributor import TorchDistributor
def train_function():
"""각 GPU에서 실행되는 학습 함수"""
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
# 1. 분산 환경 초기화 (TorchDistributor가 자동 설정)
dist.init_process_group("nccl")
rank = dist.get_rank()
local_rank = int(os.environ.get("LOCAL_RANK", 0))
device = torch.device(f"cuda:{local_rank}")
# 2. 모델 정의 및 DDP 래핑
model = MyModel().to(device)
model = DDP(model, device_ids=[local_rank])
# 3. 분산 데이터 로더 설정
sampler = torch.utils.data.distributed.DistributedSampler(dataset)
dataloader = torch.utils.data.DataLoader(
dataset, batch_size=32, sampler=sampler
)
# 4. 학습 루프
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
for epoch in range(10):
sampler.set_epoch(epoch) # 에포크마다 셔플링
for batch in dataloader:
inputs, labels = batch[0].to(device), batch[1].to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = torch.nn.functional.cross_entropy(outputs, labels)
loss.backward()
optimizer.step()
if rank == 0: # 메인 프로세스에서만 로깅
print(f"Epoch {epoch}, Loss: {loss.item():.4f}")
# 5. 모델 저장 (메인 프로세스만)
if rank == 0:
torch.save(model.module.state_dict(), "/tmp/model.pt")
dist.destroy_process_group()
# 4개 GPU에서 분산 학습 실행
distributor = TorchDistributor(
num_processes=4, # 총 프로세스(GPU) 수
local_mode=False, # False: 여러 노드에 분산 / True: 단일 노드
use_gpu=True # GPU 사용
)
result = distributor.run(train_function)