Articlehuggingface.co·2025년 8월 12일·2

Accelerate ND-Parallel: A guide to Efficient Multi-GPU Training

Quick Summary

Accelerate와 Axolotl은 데이터·완전 샤딩 데이터·텐서·컨텍스트 병렬화를 조합해 대규모 모델의 메모리 사용량, 계산량, 장치 간 통신 비용을 균형 있게 설계하도록 지원한다.

Accelerate ND-Parallel: A guide to Efficient Multi-GPU Training 관련 대표 이미지

🖼️ 인포그래픽

Accelerate ND-Parallel: A guide to Efficient Multi-GPU Training 내용을 설명하는 본문 이미지

🖼️ 4컷 인포그래픽

Accelerate ND-Parallel: A guide to Efficient Multi-GPU Training의 핵심 내용을 4단계로 요약한 인포그래픽
Accelerate ND-Parallel: A guide to Efficient Multi-GPU Training 핵심 내용을 4단계로 압축한 4컷 인포그래픽

💡 한 줄 요약

Accelerate와 Axolotl은 데이터·완전 샤딩 데이터·텐서·컨텍스트 병렬화를 조합해 대규모 모델의 메모리 사용량, 계산량, 장치 간 통신 비용을 균형 있게 설계하도록 지원한다.

📌 핵심 요약

  • Accelerate의 ParallelismConfig와 Axolotl 설정 필드를 사용하면 데이터 병렬화, 완전 샤딩 데이터 병렬화, 텐서 병렬화, 컨텍스트 병렬화의 차수를 선언적으로 지정하고 함께 적용할 수 있다.
  • 데이터 병렬화는 모델 전체를 장치마다 복제해 처리량을 높이지만 모델이 단일 장치에 들어가야 하며, 다른 병렬화 전략과 조합할 때 최상위 복제 계층으로 작동한다.
  • 완전 샤딩 데이터 병렬화는 가중치·그래디언트·옵티마이저 상태를 여러 장치에 나눠 메모리를 절약하는 대신, 순전파와 역전파 과정에서 매개변수를 수집하고 다시 분산하는 통신 비용을 부담한다.
  • 텐서 병렬화는 큰 선형 계층의 행렬 연산과 매개변수를 장치별로 정적으로 분할해 메모리와 계산량을 줄이지만, 빈번한 활성값 동기화 때문에 빠른 노드 내부 연결에서 사용하는 것이 적합하다.
  • 매우 긴 문맥에서는 어텐션 메모리가 문맥 길이의 제곱에 비례해 증가하므로 컨텍스트 병렬화가 필요하며, 실제 구성은 모델 크기와 시퀀스 길이뿐 아니라 노드 구조와 장치 간 통신 속도를 함께 고려해야 한다.

🧩 주요 포인트

  1. Accelerate의 ParallelismConfig와 Axolotl 설정 필드를 사용하면 데이터 병렬화, 완전 샤딩 데이터 병렬화, 텐서 병렬화, 컨텍스트 병렬화의 차수를 선언적으로 지정하고 함께 적용할 수 있다.
  2. 데이터 병렬화는 모델 전체를 장치마다 복제해 처리량을 높이지만 모델이 단일 장치에 들어가야 하며, 다른 병렬화 전략과 조합할 때 최상위 복제 계층으로 작동한다.
  3. 완전 샤딩 데이터 병렬화는 가중치·그래디언트·옵티마이저 상태를 여러 장치에 나눠 메모리를 절약하는 대신, 순전파와 역전파 과정에서 매개변수를 수집하고 다시 분산하는 통신 비용을 부담한다.
  4. 텐서 병렬화는 큰 선형 계층의 행렬 연산과 매개변수를 장치별로 정적으로 분할해 메모리와 계산량을 줄이지만, 빈번한 활성값 동기화 때문에 빠른 노드 내부 연결에서 사용하는 것이 적합하다.
  5. 매우 긴 문맥에서는 어텐션 메모리가 문맥 길이의 제곱에 비례해 증가하므로 컨텍스트 병렬화가 필요하며, 실제 구성은 모델 크기와 시퀀스 길이뿐 아니라 노드 구조와 장치 간 통신 속도를 함께 고려해야 한다.

🧠 상세 정리

1. 다차원 병렬화의 통합 목적

대규모 모델을 여러 그래픽 처리 장치에서 학습하려면 병렬화 방식마다 다른 메모리 절감 효과와 통신 특성을 이해해야 한다. 글은 Accelerate와 Axolotl이 데이터 병렬화, 완전 샤딩 데이터 병렬화, 컨텍스트 병렬화, 텐서 병렬화를 임의로 조합할 수 있도록 통합한 기능을 소개한다. 목표는 특정 전략 하나를 강제하는 것이 아니라 모델 규모, 입력 길이, 장치 수, 노드 간 연결 조건에 맞춰 각 병렬화 차수를 구성하는 것이다. 특히 수십억에서 수천억 개 매개변수를 가진 모델로 확장할수록 단순히 장치를 늘리는 것만으로는 충분하지 않으며, 장치 사이의 통신 오버헤드를 최소화하도록 전략 간 상호작용을 설계해야 한다.

2. Accelerate와 Axolotl의 구성 방식

Accelerate에서는 ParallelismConfig에 dp_shard_size, dp_replicate_size, cp_size, tp_size를 지정해 각 병렬화 차수를 설정하며, 값이 1이면 해당 전략을 비활성화한다. 예시 구성은 네 값을 모두 2로 설정하므로 최소 16개의 장치가 필요하고, 완전 샤딩 데이터 병렬화를 위해 버전 2의 FullyShardedDataParallelPlugin과 변환기 블록 단위 자동 래핑 정책을 함께 사용한다. Accelerator가 생성한 장치 메시를 모델을 불러올 때 전달한 뒤 accelerator.prepare로 준비하면 설정된 병렬화 구성이 모델에 적용된다. Axolotl에서도 dp_shard_size, dp_replicate_size, context_parallel_size, tensor_parallel_size 필드를 기존 설정에 추가할 수 있으며, 제공된 예시 설정과 명령으로 최소 월드 크기 16의 학습을 실행할 수 있다.

3. 데이터 병렬화의 구조와 조합 규칙

데이터 병렬화는 모델, 그래디언트, 옵티마이저 상태 전체를 각 장치에 복제하고 데이터 배치를 장치별 하위 배치로 균등하게 나눈다. 각 장치는 서로 다른 데이터를 처리하지만 매개변수를 갱신하기 전에 그래디언트를 동기화하므로 단일 장치 학습보다 처리량을 크게 높일 수 있다. 다만 모델 전체와 학습 상태가 각 장치의 메모리에 들어가야 하므로, 단일 장치 용량을 넘는 모델에는 이것만으로 대응할 수 없다. dp_replicate_size는 모델 복제본 수를 정하며 데이터 병렬화는 최상위 계층으로 작동한다. 예를 들어 데이터 병렬화 차수와 텐서 병렬화 차수를 각각 2로 설정하면 모델 복제본이 두 개 만들어지고, 각 복제본 내부가 다시 두 개의 텐서 병렬 샤드로 나뉜다.

4. 완전 샤딩 데이터 병렬화의 메모리 절감 원리

완전 샤딩 데이터 병렬화는 단일 장치에 모델이 들어가지 않는 문제를 해결하기 위해 가중치, 그래디언트, 옵티마이저 상태를 여러 장치에 균등하게 분산한다. 각 장치는 전체 데이터 배치의 일부를 받지만, 순전파와 역전파를 수행할 때 필요한 매개변수를 수집해야 하며 처리 후에는 이를 다시 샤딩할 수 있다. 모델 전체를 계속 복제하는 대신 일반적으로 한 번에 하나의 변환기 디코더 블록에 해당하는 가중치만 모으기 때문에 최고 메모리 사용량을 낮출 수 있다. 그러나 계층마다 더 세밀하게 수집하고 재분산할수록 메모리는 더 절약되는 반면 통신 횟수와 비용은 증가한다. 따라서 샤딩 효과와 매개변수 수집 비용 사이의 균형을 결정하는 래핑 단위가 핵심 조정 요소가 된다.

5. 노드 구조와 완전 샤딩의 확장 한계

글은 여러 그래픽 처리 장치를 수용하는 한 대의 머신을 노드라고 부르고, 전체 프로세스에 참여하는 장치 수를 월드 크기로 정의한다. 한 노드 내부에서는 엔브이링크와 같은 빠른 연결을 사용할 수 있지만, 여러 노드 사이에서는 인피니밴드와 같은 상대적으로 느린 통신 경로를 거친다. 완전 샤딩 데이터 병렬화를 여러 노드 전체에 적용하면 모든 장치를 하나의 큰 집합처럼 취급해 샤딩하며, 올리듀스와 리듀스 스캐터 연산이 노드 내부와 노드 사이를 모두 통과한다. 예를 들어 그래픽 처리 장치 8개를 가진 노드 4개라면 32개 장치 전체에 걸쳐 샤딩할 수 있지만, 노드 간 통신 비용도 함께 증가한다. 이 때문에 글은 일반적으로 완전 샤딩 범위를 한 노드보다 크게 만들지 않으려 하며, 더 큰 규모에서는 하이브리드 샤딩 등 다른 전략과의 조합이 필요하다고 설명한다.

6. 텐서 병렬화의 계산 분할 방식

텐서 병렬화는 모델의 큰 선형 계층을 여러 장치에 영구적으로 나누고, 모든 장치가 동일한 데이터 배치를 받아 행렬 곱셈의 일부씩 계산하도록 한다. 연속된 두 선형 계층을 구성할 때 첫 번째 계층은 열 방향으로, 다음 계층은 행 방향으로 분할하면 샤딩된 출력을 결합하기 위한 올리듀스 연산을 한 번으로 줄일 수 있다. 변환기 모델의 피드포워드 계층뿐 아니라 어텐션의 쿼리, 키, 값, 출력 투영에도 거의 추가 통신 비용 없이 적용할 수 있다고 설명한다. 완전 샤딩처럼 실행 중 매개변수를 동적으로 모으는 방식이 아니라 정적인 메모리 파티션을 유지하므로, 텐서 병렬화 그룹 크기에 따라 일정한 메모리 절감 효과를 얻는다. 특히 완전 샤딩 과정에서 디코더 계층 하나조차 메모리에 올릴 수 없는 초대형 모델에 중요하다.

7. 텐서 병렬화의 통신 조건과 적용 범위

텐서 병렬화에서는 각 장치가 출력의 일부만 계산하므로 다음 연산으로 넘어가기 전에 다른 장치의 결과와 활성값을 빈번하게 동기화해야 한다. 이 특성 때문에 노드 수를 늘리는 것만으로 선형적으로 확장하기 어렵고, 빠른 노드 내부 연결을 갖춘 단일 노드 범위에서 가장 효과적이다. 여러 노드로 학습을 확장하려면 텐서 병렬화 그룹은 노드 내부에 유지하면서 데이터 병렬화나 완전 샤딩 데이터 병렬화 같은 다른 전략을 노드 사이에 배치해야 한다. 통신량이 큰 만큼 피시아이 익스프레스만으로 연결된 그래픽 처리 장치에는 권장되지 않는다. Accelerate에서는 tp_size로, Axolotl에서는 tensor_parallel_size로 텐서 병렬화 차수를 지정한다.

8. 긴 문맥 학습이 제기하는 컨텍스트 병렬화 문제

추론 능력을 강화한 대규모 언어 모델이 복잡한 작업을 해결하는 데 더 많은 토큰을 사용하면서, 미세 조정 과정에서도 매우 긴 시퀀스를 처리해야 하는 요구가 커졌다. 글은 필요한 문맥 길이가 경우에 따라 백만 토큰에 이를 수 있다고 설명하지만, 변환기의 어텐션 연산은 문맥 길이의 제곱에 비례해 증가하므로 단일 그래픽 처리 장치로 감당하기 어렵다고 지적한다. 제시된 예에서는 어텐션 헤드 32개를 사용하는 70억 매개변수 규모 모델에 12만 8천 토큰을 적용할 경우, 단일 어텐션 행렬의 활성값 메모리가 헤드 전체에서 약 1테라바이트에 이른다. 이 수치는 긴 문맥 학습에서 모델 매개변수만 분할해서는 충분하지 않고 시퀀스와 어텐션 계산 자체를 여러 장치에 분산해야 하는 이유를 보여준다. 제공된 원문 범위는 이 메모리 문제를 제시하는 지점까지이며, 컨텍스트 병렬화의 구체적인 분할 알고리즘이나 통신 절차는 이어지는 내용에 포함되어 있지 않다.

🧾 핵심 주장 / 시사점

  • 병렬화 차수는 독립적인 숫자가 아니라 곱셈적으로 장치 배치를 형성하므로, 각 차수의 곱이 사용 가능한 월드 크기와 맞는지 먼저 확인해야 한다.
  • 완전 샤딩은 매개변수 상태의 메모리를 줄이고 텐서 병렬화는 계층 내부의 메모리와 계산을 정적으로 나누므로, 두 전략은 같은 문제를 중복 해결하기보다 서로 다른 메모리 병목을 보완한다.
  • 효율적인 다중 그래픽 처리 장치 학습의 핵심은 장치 수 자체보다 통신 계층에 맞춘 배치이며, 통신이 빈번한 텐서 병렬화는 노드 내부에 두고 상대적으로 상위 수준의 복제·샤딩 전략으로 여러 노드를 연결하는 구성이 중요하다.

✅ 액션 아이템

  • Accelerate의 ParallelismConfig와 Axolotl 설정으로 데이터·완전 샤딩·텐서·컨텍스트 병렬 차수를 선언적으로 묶어 실험 구성을 정한다.
  • 데이터 병렬화가 최상위 복제 계층으로 동작한다는 전제를 두고 모델 단일 장치 적합성을 확인한 뒤 조합 순서를 설계한다.
  • 완전 샤딩 데이터 병렬화는 메모리 절감과 통신 비용이 동시 발생하므로 순전파·역전파의 파라미터 집계·재분산 비용을 함께 점검한다.

❓ 열린 질문

  • 어떤 모델 크기에서 데이터 병렬화 단독 구성이 한계에 닿아 완전 샤딩 또는 텐서 병렬화를 추가해야 하는가?
  • 어떤 시퀀스 길이에서 어텐션 메모리의 제곱 증가가 커져 컨텍스트 병렬화를 먼저 늘려야 하는가?
  • 노드 구조와 장치 간 통신 속도를 반영해 텐서 병렬화의 활성값 동기화 비용 허용 범위는 어떻게 판단할 것인가?

관련 문서

공통 태그와 주제 흐름을 기준으로 같이 보면 좋은 문서를 이어서 제안합니다.