본문으로 건너뛰기

[sglang] SGLang Diffusion: 2-rank Ulysses를 위한 CUDA-IPC 기반 Zero-Staging All-to-All 최적화

PR 링크: sgl-project/sglang#31854 상태: Merged | 변경: +779 / -8

들어가며

SGLang의 Diffusion 모델 추론 파이프라인에서 2-rank Ulysses 병렬화(TP=2)를 사용할 때, NCCL All-to-All 통신은 병목 현상의 주범이었습니다. 각 Transformer 레이어마다 4번의 NCCL All-to-All이 발생하며, 이 과정에서 발생하는 rendezvous 오버헤드와 반복적인 데이터 staging(transpose, contiguous, cat)이 실제 데이터 전송 시간보다 더 큰 비용을 차지했습니다. 본 PR은 이를 해결하기 위해 CUDA-IPC(Inter-Process Communication)를 활용한 커스텀 전송 계층을 도입하여 성능을 최적화했습니다.

코드 분석

1. ipc_a2a.py: CUDA-IPC 전송 계층 구현

핵심은 각 rank가 상대방의 staging 버퍼를 자신의 device context에 직접 매핑하여 NVLink를 통해 접근하는 것입니다. NCCL의 rendezvous 대신 GPU-side sequence counter(bump_signal, spin_wait)를 사용하여 동기화 오버헤드를 최소화했습니다.

# ipc_a2a.py: GPU-side 동기화 커널 예시
__global__ void bump_signal_kernel(int* seq, volatile int* peer_flag) {
    int v = *seq + 1;
    *seq = v;
    __threadfence_system();
    *peer_flag = v;
}

2. base_device_communicators.py: 통신 경로 최적화

기존 NCCL 경로를 유지하면서, 2-rank 환경에서 조건이 충족될 경우(ipc_a2a_ready) IPC 경로를 우선 사용하도록 분기 처리했습니다.

+            if world_size == 2 and scatter_dim in (1, 2):
+                fast = _ipc_all_to_all_4d(group, input_, scatter_dim)
+                if fast is not None:
+                    return fast

왜 이게 좋은가

  1. Zero-Staging: Qwen joint attention의 경우, projection 출력을 상대방의 소비 버퍼로 직접 이동시켜 중간 gather 복사본을 완전히 제거했습니다. 데이터는 projection 출력에서 소비 버퍼로 단 한 번만 이동합니다.
  2. CUDA-Graph Capturable: NCCL collective는 CUDA Graph 캡처 시 시퀀스 번호 문제로 데드락이 발생하기 쉽지만, 본 구현은 로컬 메모리 기반 커널을 사용하여 전체 forward 패스를 CUDA Graph로 캡처할 수 있게 합니다.
  3. 성능 수치: Qwen-Image-2512 모델 기준, 기존 NCCL 방식 대비 약 10%의 성능 향상(87ms/step → 79ms/step)을 달성했으며, 전체 forward 그래프 캡처 적용 시 63ms/step까지 단축되었습니다.

교훈: 2-rank와 같은 소규모 노드 통신에서는 NCCL의 범용적인 오버헤드보다, NVLink를 활용한 P2P 직접 접근과 커스텀 동기화가 훨씬 효율적입니다. 다만, rank가 증가할수록 NCCL의 커널 퓨전 효율이 더 높아지므로, 이 최적화는 2-rank 환경에 한정하는 것이 적절합니다.

리뷰어 피드백 반영

  • Timeout 처리: spin_waitclock64 기반 타임아웃을 추가하여 데드락을 방지하고, 실패 시 NCCL로 안전하게 fallback하도록 설계했습니다.
  • 플랫폼 호환성: ROCm 등 CUDA를 지원하지 않는 환경에서는 IPC를 비활성화하고 NCCL을 사용하도록 가드를 추가했습니다.
  • 정확성 검증: staging 버퍼 캐시 eviction 로직을 추가하여 메모리 누수를 방지하고, CI에서 bitwise-identical한 결과를 보장하도록 테스트를 강화했습니다.

참고 자료

⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.

댓글

관련 포스트

PR Analysis 의 다른글