Colossal-AI Booster API — 플러그인·훈련 루프·체크포인트
Colossal-AI Booster API
새 설계에서 colossalai.booster는 이전 colossalai.initialize의 역할을 대체해 훈련 컴포넌트(모델·옵티마이저·데이터로더)에 기능을 매끄럽게 주입해요. 이 API로 모델을 병렬 기능과 통합하고, 훈련 루프에 들어가기 전에 colossalai.booster를 호출하는 게 표준 절차예요.
플러그인
플러그인은 병렬 설정을 관리하는 중요한 컴포넌트예요. 현재 지원하는 플러그인은 다음과 같아요.
- 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를 사용해 인자가 같아요.