본문으로 건너뛰기

[vllm] vLLM의 작은 배치 사이즈를 위한 Triton 기반 Split-row Top-p 샘플링 최적화

PR 링크: vllm-project/vllm#54651 상태: Merged | 변경: +627 / -13

들어가며

LLM 추론 엔진인 vLLM에서 샘플링 단계는 특히 작은 배치 사이즈(small-batch) 환경에서 병목이 될 수 있습니다. 기존의 모놀리식(monolithic) Triton 커널은 배치 사이즈가 작을 때 GPU의 SM(Streaming Multiprocessor) 활용도가 낮고, 전체 어휘 사전(vocab)을 순회하는 방식 때문에 지연 시간이 컸습니다. 본 PR은 이러한 문제를 해결하기 위해 split-row 파이프라인을 도입하여, 각 행을 여러 프로그램으로 나누어 병렬 처리함으로써 작은 배치에서의 샘플링 성능을 획기적으로 개선했습니다.

코드 분석

1. vllm/v1/sample/ops/topk_topp_sampler.py

기존에는 배치 사이즈가 8 미만일 경우 무조건 PyTorch 구현체로 폴백(fallback)했으나, 이제는 Triton이 사용 가능한 환경이라면 배치 사이즈와 관계없이 최적화된 Triton 커널을 사용하도록 변경되었습니다.

# Before
if HAS_TRITON and logits.shape[0] >= 8:
    return apply_top_k_top_p_triton(logits, k, p)
return apply_top_k_top_p_pytorch(logits, k, p)

# After
if HAS_TRITON:
    return apply_top_k_top_p_triton(logits, k, p)

2. vllm/v1/sample/ops/topk_topp_triton.py

핵심 변경 사항은 _topp_sb_stats_kernel과 같은 분할 처리 로직입니다. 각 행을 S개의 슬라이스로 나누어 병렬로 통계량을 계산하고, 이를 결합하는 방식을 채택했습니다.

# 각 슬라이스별로 부분 통계량 계산
@triton.jit
def _topp_sb_stats_kernel(...):
    slice_len = (VOCAB_SIZE + S - 1) // S
    start = slice_id * slice_len
    # ... (부분 합 및 최대값 계산)
    tl.store(STATS + base + 0, m) # max
    tl.store(STATS + base + 1, exp_sum) # sum_exp

왜 이게 좋은가

이번 최적화는 특히 GB200과 같은 최신 하드웨어에서 큰 성능 향상을 보여줍니다.

  • 성능 향상: 배치 사이즈 1~64 구간에서 기존 대비 1.5배에서 최대 3.7배 이상의 속도 향상을 기록했습니다. 특히 top-p 전용 경로에서 지연 시간이 크게 감소했습니다.
  • 결정론적 동작: 리뷰 과정에서 제기된 tie-breaking(동점 처리) 문제에 대해, 인덱스 순서를 보존하는 로직을 추가하여 PyTorch 구현체와 동일한 결정론적 결과를 보장하도록 수정했습니다.
  • 일반적 교훈: GPU 커널 최적화 시, 단순히 전체를 하나의 커널로 처리하는 것보다, 데이터의 특성(배치 사이즈, 연산 밀도)에 따라 작업을 분할(split-row)하고 병렬화하는 전략이 SM 활용도를 높이는 데 매우 효과적임을 보여줍니다.

리뷰어 피드백 반영

리뷰어 cakeng은 커널 간 데이터 전달을 최소화하고 로직을 통합할 것을 제안했으나, 실험 결과 현재의 다단계 파이프라인 방식이 성능과 코드 복잡도 측면에서 가장 균형 잡힌 결과를 보였습니다. 이는 성능 최적화가 항상 이론적인 코드 단순화와 일치하지 않으며, 하드웨어 특성에 맞춘 실측이 중요함을 시사합니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글