[flashinfer] FlashInfer, 통신 최적화를 통한 LLM 추론 속도 향상: Ulysses Head-Chunk Primitives 도입
PR 링크: flashinfer-ai/flashinfer#5027 상태: Merged | 변경: +2653 / -31
들어가며
대규모 언어 모델(LLM)의 추론 성능은 모델의 크기뿐만 아니라, 특히 분산 환경에서의 통신 오버헤드에 크게 좌우됩니다. 모델이 여러 GPU에 분산될 때, 각 GPU 간의 데이터 교환은 병목 현상의 주요 원인이 될 수 있습니다. FlashInfer는 이러한 문제를 해결하기 위해 Ulysses Head-Chunk Primitives를 도입하는 PR을 통해 통신 효율성을 극대화하는 새로운 방법을 제시합니다.
이 PR은 기존의 통신 방식을 개선하여, 특히 헤드(head) 단위로 데이터를 분할하고 재구성함으로써 GPU 간 데이터 전송량을 줄이고 계산과 통신을 중첩시키는 것을 목표로 합니다. 이는 LLM 추론 시 발생하는 지연 시간을 단축하고 전반적인 처리량을 향상시키는 데 기여합니다. 본 글에서는 이 PR의 주요 변경 사항, 코드 분석, 그리고 이러한 최적화가 왜 효과적인지에 대해 자세히 살펴보겠습니다.
코드 변경사항 분석
이번 PR은 주로 benchmarks/comm/bench_ulysses_head_chunk.py, flashinfer/triton/ulysses.py, tests/comm/test_ulysses_head_chunk.py 파일을 중심으로 변경되었습니다. 핵심은 Head-Chunked Ulysses Primitives의 도입과 이를 지원하는 재사용 가능한 워크스페이스(reusable workspaces) 및 헤드-청크 레이아웃(head-chunk layout) 입니다.
1. 재사용 가능한 워크스페이스 및 헤드-청크 레이아웃 도입 (benchmarks/comm/bench_ulysses_head_chunk.py)
이 파일은 새로운 Head-Chunk Primitives를 벤치마킹하고 참조하는 역할을 합니다. 주요 변경 사항은 다음과 같습니다.
-
ReferenceHeadChunkPipeline클래스: 이 클래스는 입력 A2A, 어텐션, 출력 A2A를 겹치는 3-스트림 파이프라인을 구현합니다. 이는 계산과 통신을 중첩시키는 핵심 메커니즘을 보여줍니다.__init__메서드:input_comm과output_comm을 받아 각각 입력과 출력에 대한 독립적인 통신 채널을 설정합니다. 이는 양방향 통신 오버랩을 가능하게 합니다.self.input_comm = input_comm self.output_comm = output_comm # ... self.input_workspace = input_comm.create_workspace(max_elems=input_capacity) self.output_workspace = output_comm.create_workspace(max_elems=output_capacity)create_workspace()메서드를 통해 재사용 가능한 워크스페이스를 생성하여, 매번 통신 시마다 메모리를 할당하고 해제하는 오버헤드를 줄입니다. 이는 성능 향상에 직접적으로 기여합니다.sequential메서드: Head-Chunk Primitives를 순차적으로 사용하는 방식을 구현합니다. 각 헤드 청크를 분산시키고(scatter), 어텐션을 계산한 후, 다시 모읍니다(gather).payload = self.input_comm.scatter_qkv_head_chunk( q, k, v, head_offset=offset, head_count=head_count, out=self.qkv_chunks[index], workspace=self.input_workspace, ) attention_out = _attention_from_fused(payload, self.head_dim) self.output_comm.gather_output_head_chunk( attention_out, local_heads=self.local_heads, head_offset=offset, out=self.output, workspace=self.output_workspace, )scatter_qkv_head_chunk와gather_output_head_chunkAPI는 헤드 청크 단위의 효율적인 데이터 전송을 담당합니다.workspace인자를 통해 미리 할당된 메모리 공간을 재사용합니다.overlap메서드: 계산과 통신을 중첩시키는 핵심 로직을 구현합니다. 별도의 CUDA 스트림을 사용하여 입력 A2A, 어텐션 계산, 출력 A2A를 동시에 진행합니다.이 코드는 입력 데이터의with torch.cuda.stream(self.input_stream): self.input_comm.scatter_qkv_head_chunk(...) # 입력 A2A self.input_ready[index].record(self.input_stream) compute_stream.wait_event(self.input_ready[index]) attention_out = _attention_from_fused(self.qkv_chunks[index], self.head_dim) # 어텐션 계산 self.compute_ready[index].record(compute_stream) with torch.cuda.stream(self.output_stream): self.output_stream.wait_event(self.compute_ready[index]) self.output_comm.gather_output_head_chunk(...) # 출력 A2Ascatter_qkv_head_chunk가 완료되기를 기다렸다가 어텐션 계산을 시작하고, 어텐션 계산이 완료되면 그 결과를gather_output_head_chunk로 보내는 과정을 별도의 스트림에서 비동기적으로 처리합니다. 이를 통해 GPU의 유휴 시간을 최소화합니다.
-
벤치마크 설정:
main함수에서는 분산 환경 초기화,UlyssesCommunicator생성, 그리고 다양한 통신 전략(ordinary, whole QKV fusion, sequential head chunks, overlapped head chunks)을 비교하는 벤치마킹 로직을 포함합니다.input_comm = UlyssesCommunicator( input_group, max_elems=whole_input_elems, dtype=dtype, backend=args.backend, device=device, ) output_comm = UlyssesCommunicator( output_group, max_elems=whole_output_elems, dtype=dtype, backend=args.backend, device=device, )UlyssesCommunicator는torch.distributed.ProcessGroup을 기반으로 하며, 최대 엘리먼트 수, 데이터 타입, 백엔드 등을 지정하여 초기화됩니다. 두 개의 독립적인ProcessGroup(input_group,output_group)을 사용하여 입력과 출력 통신을 분리하는 것이 핵심입니다.
2. Triton 커널 최적화 (flashinfer/triton/ulysses.py)
리뷰어 qsang-nv의 지적에 따라, Triton 커널의 동적 형태(dynamic shape) 및 스트라이드(stride) 처리가 최적화되었습니다. 이전에는 이러한 값들이 tl.constexpr로 컴파일되어 캐시 키에 포함되었기 때문에, 입력 크기나 스케줄이 조금만 달라져도 커널이 재컴파일되는 문제가 있었습니다.
-
이전 코드 (문제점):
# flashinfer/triton/ulysses.py (이전 버전의 개념적 표현) @triton.jit def pack_ulysses_qkv_head_chunk(...): # ... # head_count, head_offset, seq_len, batch 등 tl.constexpr로 사용 # ...위와 같이
tl.constexpr로 사용된 인자들은 컴파일 타임에 결정되어야 했고, 이는 다양한 입력에 대해 빈번한 재컴파일을 유발했습니다. -
개선된 코드 (
e782390a커밋 이후):# flashinfer/triton/ulysses.py (개선 후 개념적 표현) @triton.jit def pack_ulysses_qkv_head_chunk(..., head_count, seq_len, batch_idx, ...): # ... # head_count, seq_len, batch_idx 등 런타임 인자로 전달 # ... # stride 관련 값들도 런타임 인자로 처리 # ...e782390a커밋에서chunk_heads,seq_len/local_seq,head_offset, 그리고 모든 Q/K/V/소스/수신/출력 스트라이드 관련 값들이tl.constexpr대신 런타임 인자로 변경되었습니다. 이를 통해 컴파일 캐시 히트율이 높아지고 JIT 지연 시간이 감소했습니다.다만, 모델 및 토폴로지 불변량(
world_size,local_heads,head_dim)과 제어 상수(nccl_layout,BLOCK)는 여전히constexpr로 유지되어, 핫 인덱스 연산이 폴딩(folding)될 수 있도록 했습니다. 이 변경으로 인해 QKV-pack 및 output-merge 커널의 컴파일된 변형(variant) 수가 6개에서 2개로 감소했습니다.
3. 테스트 및 검증 개선 (tests/comm/test_ulysses_head_chunk.py)
리뷰어 qsang-nv가 지적한 사이드 스트림(side stream) 순서 및 워크스페이스 사용 관련 문제도 해결되었습니다.
- 사이드 스트림 순서 문제: 이전에는 사이드 스트림 테스트에서 입력과 참조를 기본 스트림에서 생성한 후, 의존성 없이 새로운 스트림에서 읽어오는 문제가 있었습니다. 개선 후에는
stream.wait_stream(torch.cuda.current_stream())을 호출하여 기본 스트림이 완료될 때까지 기다린 후 사이드 스트림 작업을 시작하도록 수정되었습니다.# tests/comm/test_ulysses_head_chunk.py (개선 후 개념적 표현) with torch.cuda.stream(self.input_stream): self.input_stream.wait_stream(torch.cuda.current_stream()) # ... 입력 데이터 처리 ... - 워크스페이스 사용:
ReferenceHeadChunkPipeline클래스에서sequential()과overlap()메서드가 동일한self.output버퍼를 반환하여sequential결과가overlap에 의해 덮어쓰이는 문제가 있었습니다. 이를 해결하기 위해sequential메서드가overlap메서드를 호출하기 전에self.output의 복사본을 생성하도록 수정되었습니다.(참고: 실제 커밋# tests/comm/test_ulysses_head_chunk.py (개선 후 개념적 표현) def sequential(self, q, k, v): # ... sequential_result = self.output.clone() # 복사본 생성 # ... return sequential_resulte782390a에서는sequential메서드가overlap메서드를 호출하기 전에self.output의 복사본을 생성하는 방식으로 수정되었습니다. 위 코드는 개념적인 이해를 돕기 위한 예시입니다.)
왜 이게 좋은가?
이번 PR에서 도입된 Head-Chunked Ulysses Primitives와 관련 최적화는 다음과 같은 이유로 LLM 추론 성능 향상에 크게 기여합니다.
- 통신량 감소: 기존에는 전체 QKV 데이터를 한 번에 주고받았다면, Head-Chunked 방식은 헤드 단위로 데이터를 분할하여 전송합니다. 이는 특히 헤드 수가 많은 모델에서 GPU 간 전송해야 하는 데이터 양을 크게 줄여 통신 병목 현상을 완화합니다. PR 설명의 벤치마크 결과는 아니지만, 관련 SM100 FA4 프로토타입에서 11~12%의 속도 향상을 보였다는 점은 이러한 잠재력을 시사합니다.
- 계산-통신 중첩 (Compute-Communication Overlap):
ReferenceHeadChunkPipeline의overlap메서드에서 볼 수 있듯이, 입력 데이터 수신(A2A), 어텐션 계산, 출력 데이터 전송(A2A)을 별도의 CUDA 스트림에서 동시에 수행합니다. 이를 통해 GPU가 데이터를 기다리는 시간을 최소화하고, 연산 장치를 최대한 활용할 수 있습니다. 이는 특히 NVLink와 같이 통신 속도가 빠른 환경에서 더욱 효과적입니다. - 메모리 할당 최적화:
create_workspace()를 통해 재사용 가능한 워크스페이스를 도입함으로써, 반복적인 통신 작업에서 발생하는 동적 메모리 할당 및 해제 오버헤드를 제거했습니다. 이는 특히 배치 크기나 시퀀스 길이가 가변적인 경우에 성능 안정성을 높이는 데 중요합니다. - Triton 커널 컴파일 최적화: 동적 형태 및 스트라이드 인자를 런타임 인자로 변경함으로써, 다양한 입력 크기에 대한 Triton 커널의 재컴파일 빈도를 줄였습니다. 이는 모델 로딩 시간 및 추론 시작 시간을 단축시키고, 컴파일 캐시 관리 효율성을 높입니다.
리뷰 과정에서 제기된 문제점들(사이드 스트림 순서, 워크스페이스 재사용, Triton 커널 최적화)이 해결되면서, 이 PR은 더욱 견고하고 효율적인 통신 프레임워크를 제공하게 되었습니다. 특히, qsang-nv 리뷰어는 Triton 커널의 constexpr 사용이 JIT 지연 시간과 캐시 크기 증가의 원인이 된다고 지적했고, tiffany940107님은 이를 런타임 인자로 변경하여 해결했습니다. 이는 LLM 추론 성능 최적화에서 커널 컴파일 전략의 중요성을 다시 한번 강조합니다.
결론
이번 FlashInfer PR은 Head-Chunked Ulysses Primitives 도입을 통해 LLM 추론 시 통신 효율성을 혁신적으로 개선했습니다. 재사용 가능한 워크스페이스, 계산-통신 중첩, 그리고 최적화된 Triton 커널 컴파일 전략은 전반적인 추론 속도 향상에 크게 기여할 것으로 기대됩니다. 이러한 기술은 대규모 모델을 효율적으로 서빙하기 위한 중요한 발걸음이며, 앞으로 LLM 인프라 발전에 긍정적인 영향을 미칠 것입니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.distributed.new_group.html
- https://docs.nvidia.com/deeplearning/triton-inference-server/main/triton-inference-server.html
- https://github.com/openai/flash-attention
- https://github.com/openai/triton
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [flashinfer] FlashInfer의 plan() 함수 최적화: Python max()에서 Tensor.max()로의 전환
- [flashinfer] FlashInfer BF16 KDA 성능 최적화: M64 Value Split 도입
- [flashinfer] FlashInfer의 Context-Parallel Decode 최적화: Fused A2A + LSE Reduce
- [flashinfer] Blackwell 아키텍처에서 FlashInfer Ragged Prefill 성능 3.3배 향상시키기
- [flashinfer] [FlashInfer] Paged Attention 최적화: 동일 Stride 구조에서의 주소 계산 오버헤드 제거
PR Analysis 의 다른글
- 이전글 [Liger-Kernel] Liger-Kernel: SM90 MoE 통신 최적화 및 안정성 개선
- 현재글 : [flashinfer] FlashInfer, 통신 최적화를 통한 LLM 추론 속도 향상: Ulysses Head-Chunk Primitives 도입
- 다음글 [cpython] Python 토크나이저 최적화: 불필요한 개행 문자 변환 건너뛰기로 성능 개선
댓글