Trainer로 학습 자동화하기

Trainer로 학습 자동화하기

PyTorch 코드를 LightningModule로 정리하고 나면, Trainer가 나머지를 전부 자동화해요. 연구·엔지니어링 보일러플레이트를 크게 줄이면서, 필요한 부분은 여전히 직접 제어할 수 있게 해주는 것이 Trainer의 역할이에요.

출처: PyTorch Lightning Trainer 문서 (공식)

Trainer가 내부에서 하는 일

겉보기엔 "그냥 학습만 돌려주는 도구"처럼 보이지만, 내부에서는 루프의 세부를 알아서 처리해요. 대표적인 예가 이렇습니다.

  • gradient 활성화/비활성화 자동 처리
  • 학습·검증·테스트 데이터로더 실행
  • 콜백을 적절한 시점에 호출
  • 배치와 계산을 올바른 디바이스에 배치

학습 루프만 떼어 본 유사 코드는 이런 모양이에요.

# enable grads
torch.set_grad_enabled(True)
losses = []
for batch in train_dataloader:
    on_train_batch_start()          # 훅 호출
    loss = training_step(batch)     # 학습 스텝
    optimizer.zero_grad()           # 그래디언트 초기화
    loss.backward()                 # 역전파
    optimizer.step()                # 파라미터 갱신
    losses.append(loss)

기본 사용법

Trainer는 아주 간단하게 쓸 수 있어요.

model = MyLightningModule()
trainer = Trainer()
trainer.fit(model, train_dataloader, val_dataloader)

Python 스크립트에서 사용하기

스크립트에서는 main 함수 안에서 Trainer를 호출하는 걸 권장해요. acceleratordevices를 명령줄 인자로 받을 수 있죠.

from argparse import ArgumentParser

def main(hparams):
    model = LightningModule()
    trainer = Trainer(accelerator=hparams.accelerator, devices=hparams.devices)
    trainer.fit(model)

if __name__ == "__main__":
    parser = ArgumentParser()
    parser.add_argument("--accelerator", default=None)
    parser.add_argument("--devices", default=None)
    args = parser.parse_args()
    main(args)

실행 시 이렇게 인자를 넘기면 됩니다.

python main.py --accelerator 'gpu' --devices 2

검증과 테스트

학습 루프 밖에서 별도의 검증 에폭을 수행하고 싶다면 validate를 써요. 초기화 직후나 학습 완료 후에 새 메트릭을 수집할 때 유용해요.

trainer.validate(model=model, dataloaders=val_dataloaders)

학습이 끝나고 논문을 내거나 프로덕션에 올리기 직전엔 테스트셋으로 확인해요.

trainer.test(dataloaders=test_dataloaders)

재현성 확보하기

완전히 재현 가능한 학습을 원한다면 시드를 설정하고 deterministic 플래그를 켜요.

from lightning.pytorch import Trainer, seed_everything

seed_everything(42, workers=True)  # numpy, torch, python.random 시드 설정
model = Model()
trainer = Trainer(deterministic=True)

workers=True로 설정하면 데이터로더 워커와 프로세스마다 torch, numpy, 표준 random에 고유한 시드를 부여해요. 데이터 증강이 워커마다 반복되지 않게 보장해 주죠.

주요 accelerator 설정

accelerator 인자는 하드웨어 유형을 고를 수 있게 해줘요. "cpu", "gpu", "tpu", "hpu", "auto"를 지원하고, 커스텀 accelerator 인스턴스도 넘길 수 있어요.

# CPU accelerator
trainer = Trainer(accelerator="cpu")

# GPU accelerator, GPU 2개
trainer = Trainer(devices=2, accelerator="gpu")

# TPU accelerator, 8개 코어
trainer = Trainer(devices=8, accelerator="tpu")

# DistributedDataParallel 전략
trainer = Trainer(devices=4, accelerator="gpu", strategy="ddp")

"auto"는 현재 머신을 확인하고 적절한 Accelerator를 자동으로 골라 줘요.

# 머신에 GPU가 있으면 GPU accelerator를 자동 사용
trainer = Trainer(devices=2, accelerator="auto")

더 알아보기