PyTorch FSDP를 사용한 Llama 2 70B 파인튜닝
Hugging Face는 Hugging Face Transformers, Accelerate, TRL 라이브러리를 활용하여 PyTorch Fully Sharded Data Parallelism (FSDP)를 사용하여 Llama 2 70B 모델을 파인튜닝하는 방법론을 상세히 설명했습니다. 이 접근 방식은 옵티마이저 상태(optimizer states), 그래디언트(gradients), 파라미터(parameters)를 여러 장치에 샤딩(sharding)함으로써 멀티 노드, 멀티 GPU 설정에서 대규모 모델의 학습을 가능하게 합니다.
모델 로딩 중 CPU RAM 병목 현상 극복
Llama 2 70B 모델을 로드하는 데는 일반적으로 상당한 CPU RAM이 필요합니다. 만약 노드의 모든 프로세스가 모델을 로드한다면, 약 2TB의 CPU RAM이 필요할 수 있습니다 (70B parameters * 4 bytes * 8 GPUs). 메모리 부족(OOM) 오류를 방지하기 위해 Hugging Face는 transformers 및 accelerate에 구현된 특정 초기화 전략을 사용합니다:
- Meta Device Initialization: 모델이 모든 랭크(rank)에서
meta디바이스를 사용하여 생성됩니다. 즉, 가중치 없이 초기화됩니다. - Rank 0 Loading: state dict는 rank 0에서만 로드됩니다.
- Empty Parameter Allocation: 다른 모든 랭크는
torch.empty()를 사용하여meta디바이스에 빈 파라미터를 생성합니다. - State Broadcasting:
sync_module_states=True로 설정하면, FSDP는 학습이 시작되기 전에 rank 0에서 다른 모든 랭크로 가중치를 브로드캐스트합니다.
이 방법은 노드당 하나의 프로세스만 사전 학습된 모델을 CPU RAM에 로드하도록 보장하여, 설정 단계에서의 메모리 사용량을 획기적으로 줄여줍니다.
Sharded State Dicts를 이용한 효율적인 체크포인팅
rank 0에서 CPU 오프로딩과 함께 FULL_STATE_DICT를 사용하여 전체 중간 체크포인트를 저장하면 종종 NCCL Timeout 오류와 상당한 지연이 발생합니다. 이를 해결하기 위해 Hugging Face는 FSDP 설정에서 SHARDED_STATE_DICT를 사용할 것을 권장합니다.
- Intermediate Checkpoints:
SHARDED_STATE_DICT는 GPU당 샤드를 별도로 저장하여 학습의 빠른 저장 및 재개를 가능하게 합니다. - Final Model Export: 배포를 위한 표준 모델 state dict를 얻으려면,
trainer.save_model()을 호출하기 전 학습 마지막 단계에서만 state dict 유형을FULL_STATE_DICT로 전환합니다.
VRAM 및 학습 속도 최적화
계산 비용을 줄이고 학습 속도를 높이기 위해, 구현에는 두 가지 주요 기술인 Gradient Checkpointing과 Flash Attention이 사용됩니다.
Flash Attention
표준 어텐션 메커니즘은 요소별 연산(masking, softmax, dropout) 중 중복된 High Bandwidth Memory (HBM) 읽기/쓰기로 인해 메모리 제한을 받는 경우가 많습니다. Flash Attention은 다음과 같이 이를 최적화합니다:
- Kernel Fusion: 중간 단계들을 SRAM에 유지하고 최종 결과만 HBM에 한 번 씁니다.
- Tiling: NxN softmax/scores 계산을 SRAM 제한 내에 들어오도록 블록 단위로 분할하며, 온라인 softmax 알고리즘을 활용합니다.
- Recomputation: 역전파(backward pass) 시 순전파(forward pass)의 전체 NxN 행렬을 저장하는 대신 필요한 값만 다시 계산하여 메모리 소비를 크게 줄입니다.
Gradient Checkpointing
70B 파라미터 모델의 파인튜닝 중에 더 큰 배치 크기나 더 긴 시퀀스 길이를 허용하기 위해 VRAM 사용량을 추가로 줄이는 Gradient checkpointing이 활성화됩니다.
구현 및 하드웨어 사양
하드웨어 구성
파인튜닝은 다음 하드웨어를 사용하여 수행되었습니다:
- Nodes: 2개 노드 (최소 1개 필요).
- GPUs: 노드당 8개의 A100 (80GB) GPU.
- Interconnects: NVLink (intra-node) 및 Elastic Fabric Adapter (inter-node).
- System RAM: 노드당 1TB.
- CPU: 노드당 96 코어.
학습 실행
학습 프로세스는 FULL_SHARD 전략과 TRANSFORMER_BASED_WRAP auto wrap 정책을 사용하는 accelerate launch 명령어를 사용했습니다. 혼합 정밀도(Mixed precision) 학습은 bf16을 사용하여 활성화되었습니다. 8개의 A100 80GB GPU를 사용하는 단일 노드 설정의 경우, 메모리 관리를 위해 bitsandbytes의 paged_adamw_32bit 옵티마이저를 권장합니다.
파인튜닝은 meta-llama/Llama-2-70b-chat-hf 모델과 smangrul/code-chat-assistant-v1 데이터셋을 사용하여 약 13.5시간 만에 완료되었습니다.