Accelerate ND-Parallel: 효율적인 다중 GPU 훈련 가이드

Hugging Face는 Accelerate와 Axolotl에 ND-Parallelism을 도입하여, 데이터 병렬화(DP), 완전 샤드 데이터 병렬화(FSDP), 텐서 병렬화(TP), 컨텍스트 병렬화(CP)와 같은 여러 병렬화 전략을 단일 훈련 스크립트에 결합하는 간소화된 방식을 제공합니다. 이 통합을 통해 개발자는 수십에서 수백억 개 파라미터를 가진 모델을 다중 노드 GPU 클러스터에서 훈련할 때 메모리 사용량과 통신 오버헤드 사이의 균형을 최적화할 수 있습니다.

핵심 병렬화 전략

데이터 병렬화 (DP)

Data Parallelism은 전체 모델, 그래디언트 및 옵티마이저 상태를 모든 디바이스에 복제합니다. 각 디바이스는 서로 다른 서브 배치를 처리하고, 파라미터 업데이트 전에 그래디언트를 모든 디바이스 간에 동기화합니다. 이는 처리량을 증가시키지만 전체 모델이 단일 GPU에 들어가야 합니다.

완전 샤드 데이터 병렬화 (FSDP)

FSDP는 모델 가중치, 그래디언트 및 옵티마이저 상태를 GPU에 샤드하여 디바이스당 메모리 사용량을 줄입니다. 순전파 또는 역전파를 수행하기 위해 FSDP는 특정 레이어(보통 트랜스포머 디코더 블록)에 필요한 가중치를 모은 뒤, 연산이 끝나면 다시 샤드합니다. 이는 통신 오버헤드를 증가시키는 대신 피크 메모리 사용량을 크게 낮춥니다.

텐서 병렬화 (TP)

Tensor Parallelism은 큰 선형 레이어(예: 피드포워드 레이어 또는 어텐션 프로젝션)를 디바이스에 나눕니다. FSDP의 동적 샤딩과 달리 TP는 정적인 메모리 파티션을 생성합니다. TP는 빈번한 활성화 동기화가 필요하기 때문에 고대역폭 링크(NVLink 등)를 사용하는 단일 노드 내에서 가장 효과적이며, PCIe 연결 GPU에서는 권장되지 않습니다.

컨텍스트 병렬화 (CP)

Context Parallelism은 입력 시퀀스를 GPU에 샤드하여 어텐션의 2차 스케일링으로 인해 GPU 메모리를 초과할 수 있는 매우 긴 시퀀스 길이를 처리합니다. RingAttention을 사용하면 각 GPU가 query, key, value 행렬의 일부를 보유하고, key-value 샤드를 GPU 링을 따라 순환시켜 각 query가 전체 시퀀스에 대한 어텐션 점수를 계산하도록 하면서 연산 및 메모리 부하를 분산합니다.

ND-Parallelism: 다중 노드 확장을 위한 전략 조합

다중 노드 훈련은 종종 노드 간 지연 및 메모리 제약으로 병목 현상이 발생합니다. ND-Parallelism은 클러스터를 다차원 토폴로지로 간주하여 통신을 최적화합니다.

하이브리드 샤드 데이터 병렬화 (HSDP)

HSDP는 2D 병렬화 접근법으로, 노드 내부에서는 FSDP를 수행하고(빠른 노드 내 링크 활용) 노드 간에는 DP를 수행합니다. 이는 느린 노드 간 통신을 단일 그래디언트 동기화 단계로 최소화하여 순수 FSDP에 비해 메모리 사용량이 증가하는 대신 처리량을 높입니다.

FSDP + 텐서 병렬화

FSDP와 TP를 결합하면 모델을 FSDP를 통해 노드 간에 샤드하고, TP를 사용해 노드 내부에서 레이어를 분할합니다. 이는 FSDP 지연을 감소시키고, 단일 디바이스에 맞지 않을 정도로 큰 레이어의 훈련을 가능하게 하며, 전역 배치 크기를 줄일 수 있게 합니다.

FSDP + 컨텍스트 병렬화

이 2D 전략은 매우 긴 시퀀스로 훈련할 때 사용됩니다. CP는 이미 FSDP와 통합되어 있지만, CP 위에 FSDP를 추가하면 모델 가중치와 옵티마이저 상태에 필요한 메모리 예산을 더욱 줄일 수 있습니다.

하이브리드 샤드 데이터 병렬화 + 텐서 병렬화

이 3D 계층 구조는 DP를 사용해 노드 그룹 간에 모델을 복제하고, 해당 그룹 내에서는 FSDP로 모델을 샤드하며, 각 노드 내에서는 TP로 레이어를 분할합니다. 이 구성은 특정 하드웨어와 확장 제약에 맞게 최대 유연성을 제공합니다.

구현 및 사용 시 참고 사항

Accelerate와 Axolotl에서의 설정

사용자는 Accelerate의 ParallelismConfig 클래스나 Axolotl의 특정 설정 필드를 통해 이러한 전략을 구성할 수 있습니다:

  • dp_shard_size: FSDP의 정도
  • dp_replicate_size: DP의 정도
  • tp_size / tensor_parallel_size: TP의 정도
  • cp_size / context_parallel_size: CP의 정도

메모리 및 안정성 최적화

  • CPU RAM Efficient Loading: 단일 디바이스에 들어가지 않을 정도로 큰 모델의 경우, cpu_ram_efficient_loadingSHARDED_STATE_DICTFullyShardedDataParallelPlugin에서 활성화하는 것이 중요합니다.
  • Effective Batch Size: 유효 배치 크기는 micro_batch_size * gradient_accumulation_steps * dp_world_size 로 계산되며, 여기서 dp_world_size = (dp_shard_size * dp_replicate_size) / tp_size 입니다.
  • Learning Rate Scaling: 유효 배치 크기가 증가함에 따라 학습률을 선형 또는 제곱근 스케일링으로 조정하여 안정성을 유지해야 합니다.
  • Gradient Checkpointing: 역전파 시 중간 활성화를 재계산함으로써 연산량을 메모리 절감과 교환합니다. 이는 활성화 메모리를 60-80% 줄이는 대신 훈련 시간을 약 20-30% 증가시킵니다.

Sources