Colossal-AI Booster API — 플러그인·훈련 루프·체크포인트

Colossal-AI Booster API

새 설계에서 colossalai.booster는 이전 colossalai.initialize의 역할을 대체해 훈련 컴포넌트(모델·옵티마이저·데이터로더)에 기능을 매끄럽게 주입해요. 이 API로 모델을 병렬 기능과 통합하고, 훈련 루프에 들어가기 전에 colossalai.booster를 호출하는 게 표준 절차예요.

출처: https://colossalai.org/docs/basics/booster_api/

플러그인

플러그인은 병렬 설정을 관리하는 중요한 컴포넌트예요. 현재 지원하는 플러그인은 다음과 같아요.

  • HybridParallelPlugin — 하이브리드 병렬 훈련 가속 솔루션. 텐서 병렬·파이프라인 병렬·데이터 병렬(DDP, ZeRO 포함)의 조합 인터페이스
  • GeminiPlugin — Gemin 가속 솔루션. 청크 기반 메모리 관리의 ZeRO
  • TorchDDPPlugin — PyTorch의 DDP 가속 솔루션. 모듈 수준 데이터 병렬로 여러 머신에 걸쳐 실행
  • LowLevelZeroPlugin — ZeRO 옵티마이저의 1/2 단계. Stage 1은 옵티마이저 상태를, Stage 2는 옵티마이저 상태+그래디언트를 데이터 병렬 워커/GPU에 샤딩
  • TorchFSDPPlugin — PyTorch의 FSDP 가속 솔루션. zero-dp로 모델 훈련

일부 플러그인은 lazy initialization을 지원해 큰 모델 초기화 시 메모리를 아낄 수 있어요.

Booster 클래스

colossalai.booster.Booster(device=None, mixed_precision=None, plugin=None)로 훈련을 통일된 인터페이스로 다뤄요. mixed_precision은 문자열로 'fp16', 'fp16_apex', 'bf16', 'fp8' 중 하나를 줄 수 있거나 MixedPrecision 객체일 수 있어요.

# pseudo code
colossalai.launch(...)
plugin = GeminiPlugin(...)
booster = Booster(precision='fp16', plugin=plugin)

model = GPT2()
optimizer = HybridAdam(model.parameters())
dataloader = plugin.prepare_dataloader(train_dataset, batch_size=8)
lr_scheduler = LinearWarmupScheduler()
criterion = GPTLMLoss()

model, optimizer, criterion, dataloader, lr_scheduler = booster.boost(model, optimizer, criterion, dataloader, lr_scheduler)

for epoch in range(max_epochs):
    for input_ids, attention_mask in dataloader:
        outputs = model(input_ids.cuda(), attention_mask.cuda())
        loss = criterion(outputs.logits, input_ids)
        booster.backward(loss, optimizer)
        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()

핵심 메서드

  • booster.boost(model, optimizer, criterion, dataloader, lr_scheduler) — 모델·옵티마이저·크라이테리언·스케줄러·데이터로더에 기능을 주입
  • booster.backward(loss, optimizer) — 훈련 스텝의 역전파 실행
  • booster.save_model(model, checkpoint, shard=..., size_per_shard=..., use_safetensors=...) / save_optimizer(...) — 체크포인트 저장 (샤딩·safetensors 지원)
  • booster.load_model(model, checkpoint, ...) / load_optimizer(...) / load_lr_scheduler(...) — 체크포인트 로드
  • booster.execute_pipeline(...) — 파이프라인 병렬 시 forward/backward 실행 (파이프라인 병렬 학습엔 보통 방법 대신 이 함수 사용)

LoRA를 쓸 때는 booster.lora(model, pretrained_dir=..., lora_config=..., quantize=...)를 쓸 수 있고, 구현은 Hugging Face peft를 사용해 인자가 같아요.

더 알아보기