본문으로 건너뛰기

[flashinfer] FlashInfer의 Context-Parallel Decode 최적화: Fused A2A + LSE Reduce

PR 링크: flashinfer-ai/flashinfer#4929 상태: Merged | 변경: +1627 / -0

들어가며

Context-Parallel(CP) 환경에서 LLM 디코딩을 수행할 때, KV 캐시는 시퀀스 단위로 여러 랭크(Rank)에 분산됩니다. 각 랭크는 로컬 KV 샤드에 대해 어텐션을 수행한 후, 부분적인 출력(o_r)과 정규화를 위한 lse_r(Log-Sum-Exp) 값을 생성합니다. 기존에는 이들을 합치기 위해 NCCL All-to-All 통신을 수행한 뒤, 별도의 커널에서 LSE 가중치를 계산하여 병합하는 2단계 과정을 거쳤습니다. 이 과정에서 발생하는 고정적인 오버헤드가 디코딩 성능의 병목이 되곤 합니다. 이번 PR은 이 두 과정을 하나의 cooperative CUDA 커널로 융합하여 성능을 획기적으로 개선했습니다.

코드 분석

1. Fused Kernel 구현 (csrc/dcp_lse_reduce.cu)

핵심은 NCCL의 SymmetricMemory를 활용하여 통신과 연산을 융합한 것입니다. 각 랭크는 자신의 데이터를 상대방의 메모리 영역에 직접 쓰고(LSA/NVLink), 모든 데이터가 준비되면 즉시 LSE 가중치를 계산하여 최종 출력을 산출합니다.

// Before: 별도의 NCCL 통신 후 별도 연산
// dist.all_to_all_single(recv_o, send_o);
// dist.all_to_all_single(recv_lse, send_lse);
// output = (peer_o.float() * weights).sum(dim=-2) / denom;

// After: Fused Kernel 내부에서 직접 수행
__global__ void FusedKernel(...) {
  // 1. LSA를 통한 데이터 교환
  // 2. LSE 기반 Softmax 가중치 계산
  // 3. FP32 누적 후 FP16/BF16 변환
}

2. 벤치마크 및 검증 (benchmarks/bench_dcp_lse_reduce.py)

기존 방식은 통신과 연산이 분리되어 있어, 데이터 크기가 작을 때 고정 오버헤드가 지배적이었습니다. 새로운 구현은 이를 제거하여 성능을 크게 높였습니다.

# 벤치마크 코드 중 일부: Eager 모드와 Graph 모드 비교
def eager_call() -> torch.Tensor:
    return decode_cp_a2a_lse_reduce(partial_o, partial_lse, workspace, ...)

왜 이게 좋은가

이 최적화는 단순히 연산 속도를 높이는 것이 아니라, 분산 환경에서의 고정 오버헤드(Launch Overhead)를 제거하는 데 초점을 맞췄습니다. 벤치마크 결과에 따르면, 기존 방식 대비 최대 14.7배의 성능 향상을 보였습니다.

Shape Fused eager Legacy A2A + merge Speedup
batch 1, heads 2, dim 64 12.14 µs 178.91 µs 14.7×
batch 1, heads 8, dim 128 15.04 µs 179.22 µs 11.9×

교훈

  1. 통신과 연산의 융합: 분산 시스템에서 작은 크기의 데이터를 빈번하게 주고받을 때는, 통신과 연산을 별도로 수행하는 것보다 커널 융합을 통해 오버헤드를 줄이는 것이 훨씬 효율적입니다.
  2. NCCL LSA 활용: NCCL의 SymmetricMemory와 같은 저수준 API를 직접 활용하면 프레임워크 수준의 추상화 오버헤드를 제거할 수 있습니다.
  3. CUDA Graph: 반복적인 커널 호출 시 CUDA Graph를 활용하면 런타임 오버헤드를 추가로 절감할 수 있음을 확인했습니다.

결론

이번 PR은 FlashInfer가 고성능 분산 추론을 위해 얼마나 저수준의 최적화까지 고려하는지 보여주는 좋은 사례입니다. 특히 NCCL Device API를 직접 다루는 방식은 향후 다른 분산 연산 최적화에도 큰 참고가 될 것입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글