RCCLX: AMD 플랫폼에서 GPU 통신 혁신
Meta가 AMD용 RCCL 확장판 RCCLX를 공개하고 DDA와 저정밀 집합연산으로 성능을 끌어올렸다.
Meta가 RCCLX의 초기 버전을 오픈소스로 공개했다. RCCLX는 RCCL을 확장한 버전으로, Torchcomms와 완전히 통합돼 AMD 플랫폼에서도 통신 혁신을 빠르게 실험할 수 있게 한다.
또한 CTran을 AMD 플랫폼에 통합해 AllToAllvDynamic 같은 GPU-resident collective를 지원한다. 다만 현재 오픈소스 버전에 모든 CTran 기능이 들어간 것은 아니며, 향후 몇 달 안에 추가할 계획이다.
핵심 기능으로는 **Direct Data Access (DDA)**와 Low Precision Collectives가 소개된다. 두 기능 모두 AMD 플랫폼에서 통신 병목을 줄이고, 대규모 AI 워크로드의 효율을 높이는 데 초점을 맞춘다.
DDA는 LLM 추론의 두 단계, 즉 프롬프트를 처리해 KV cache를 만드는 prefill과 토큰을 하나씩 생성하는 decode의 특성이 다르다는 점에 맞춰 설계됐다.
- Prefill은 attention 계산이 시퀀스 길이에 따라 급격히 커져 compute-bound 성격이 강하다.
- Decode는 KV cache와 모델 가중치 읽기가 지배적인 memory-bound 단계다.
- Tensor parallelism 환경에서는 AllReduce가 E2E latency의 최대 **30%**까지 차지할 수 있다.
이를 줄이기 위해 Meta는 두 가지 DDA 알고리즘을 만들었다.
- DDA flat: 작은 메시지 크기에서 각 rank가 다른 rank의 메모리를 직접 읽어 로컬 reduce를 수행한다. 지연은 **O(N)**에서 **O(1)**로 줄이고, 대신 데이터 교환량은 **O(n)**에서 **O(n²)**로 늘어난다.
- DDA tree: AllReduce를 reduce-scatter와 all-gather 두 단계로 나눈 뒤 각 단계에서 direct data access를 사용한다. ring algorithm과 같은 데이터량을 이동하면서도 조금 더 큰 메시지 크기에서 지연을 상수 수준으로 낮춘다.
AMD MI300X 기준으로 DDA는 RCCL baseline 대비 decode 구간에서 10~50%, prefill 구간에서 10~30% 성능 향상을 보였다. 이 결과로 **TTIT(time-to-incremental-token)**가 약 10% 줄어들어, 실제 사용자 체감이 큰 decode 구간이 개선됐다.
Low-precision collectives는 AllReduce, AllGather, AlltoAll, ReduceScatter를 AMD Instinct MI300/MI350 GPU에 맞게 최적화한 분산 통신 알고리즘이다. FP32와 BF16을 지원하며, FP8 quantization으로 최대 4:1 압축을 적용해 특히 16MB 이상의 큰 메시지에서 통신 오버헤드를 줄인다.
이 알고리즘은 P2P mesh communication을 사용해 AMD Infinity Fabric의 대역폭과 지연 특성을 적극 활용한다. 계산 단계는 안정성을 위해 고정밀(FP32)로 수행되며, 정밀도 손실은 주로 collective당 1~2회 정도의 quantization 횟수와 FP8 범위 내 표현 가능 여부에 의해 결정된다.
내부 실험에서는 다음과 같은 결과가 관찰됐다.
- GSM8K 평가에서 약 0.3% 수준의 변화
- E2E latency 9~10% 감소
- throughput 약 7% 증가
측정은 param-bench rccl-tests로 수행됐고, MI300은 ROCm 6.4, MI350은 ROCm 7.0 기반 RCCLX에서 테스트했다. 각 테스트는 10회 warmup과 100회 측정으로 진행됐으며, 그래프의 수치는 측정 구간 평균 throughput이다.
RCCLX는 Torchcomms API의 커스텀 backend로 통합돼 있어, 사용자는 플랫폼이 바뀌어도 같은 API를 유지한 채 애플리케이션을 옮길 수 있다. Meta는 이 backend를 NVIDIA용 NCCLX와 기능적으로 맞추는 것을 목표로 하며, CTran이 제공하는 새 기능도 같은 API 아래에서 확장할 계획이다.
마지막으로, 사용자는 Torchcomms 설치 후 torchcomms.new_comm("rcclx", torch.device("hip"), ...)처럼 communicator를 초기화하고, 환경 변수 RCCL_LOW_PRECISION_ENABLE=1로 저정밀 집합연산을 활성화할 수 있다.
이 요약은 원문 이해를 돕기 위한 큐레이션입니다. 저작권은 원저작자에게 있으며, 정확한 내용과 맥락은 원문을 확인하세요.
요약 오류, 출처 표기 문제, 삭제 요청은 문의 · 건의로 알려주세요.