본문으로 건너뛰기

[vllm] vLLM의 KV 캐시 로드 경로 최적화: NumPy를 활용한 Vectorization

PR 링크: vllm-project/vllm#48531 상태: Merged | 변경: +211 / -21

들어가며

vLLM은 LLM 추론 속도를 높이기 위해 다양한 최적화 기법을 적용하고 있습니다. 특히 긴 컨텍스트를 처리할 때 발생하는 KV 캐시(Key-Value Cache) 관련 연산은 성능에 큰 영향을 미칠 수 있습니다. 이번 PR(#48531)은 vLLM의 KV 캐시 로드 경로에서 ChunkedTokenDatabase.prepare_value 함수가 차지하는 상당한 시간을 개선하는 데 초점을 맞추고 있습니다. 기존 구현은 요청당 약 35ms의 GIL(Global Interpreter Lock) 점유 시간을 소요했는데, 이는 특히 여러 스레드가 동시에 KV 캐시를 처리할 때 병목 현상을 더욱 심화시켰습니다. 이 PR은 NumPy를 활용한 Vectorization을 통해 이 문제를 해결하고, 성능을 획기적으로 개선했습니다.

코드 분석

1. vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/data.py

이 파일은 KV 캐시 데이터를 관리하는 핵심 로직을 담고 있습니다. 이번 PR의 핵심 변경 사항은 prepare_valueprepare_values 함수의 구현 방식에 있습니다.

prepare_value (기존 방식)

기존 prepare_value 함수는 단일 토큰 범위에 대해 반복문을 사용하여 각 KV 캐시 영역별로 주소와 크기를 계산했습니다. 이는 Python의 스칼라 연산에 의존하므로, 특히 많은 KV 캐시 영역을 가진 경우 비효율적이었습니다.

def prepare_value(
    self, start: int, end: int, block_ids: list[int]
) -> tuple[list[int], list[int], int]:
    """Compute memory addresses and sizes for a token range.

    Returns:
        (addr_list, size_list, block_id)
    """
    addr_list = []
    size_list = []
    block_id = block_ids[start // self.block_size]
    length = len(self.block_len)
    for index, base_addr in enumerate(self.kv_caches_base_addr):
        addr = base_addr + block_id * self.block_len[index % length]
        assert (end - start) % self.block_size == 0
        size = self.block_len[index % length] * cdiv(end - start, self.block_size)
        addr_list.append(addr)
        size_list.append(size)
    return addr_list, size_list, block_id

prepare_values (새로운 방식)

PR 이후, prepare_valueprepare_values 함수의 래퍼(wrapper)로 변경되었으며, prepare_values 함수는 NumPy를 사용하여 여러 토큰 범위에 대한 주소 및 크기 계산을 벡터화했습니다. 이를 통해 Python의 반복문을 NumPy의 최적화된 배열 연산으로 대체하여 성능을 크게 향상시켰습니다.

import numpy as np

# ... (이전 코드) ...

def prepare_values(
    self,
    chunks: Sequence[tuple[int, int]],
    block_ids: list[int],
) -> tuple[list[list[int]], list[list[int]], list[int]]:
    """Compute memory addresses and sizes for multiple token ranges.

    Returns:
        (addr_lists, size_lists, chunk_block_ids), one entry per chunk.
    """
    if not chunks:
        return [], [], []
    base = np.asarray(self.kv_caches_base_addr, dtype=np.int64)
    length = len(self.block_len)
    blen = np.asarray(
        [self.block_len[i % length] for i in range(base.shape[0])],
        dtype=np.int64,
    )
    n = len(chunks)
    starts = np.fromiter((c[0] for c in chunks), dtype=np.int64, count=n)
    spans = np.fromiter((c[1] for c in chunks), dtype=np.int64, count=n) - starts
    assert not (spans % self.block_size).any()
    bids = np.fromiter(
        (block_ids[i] for i in (starts // self.block_size).tolist()),
        dtype=np.int64,
        count=n,
    )
    addrs = base[None, :] + bids[:, None] * blen[None, :]
    sizes = blen[None, :] * (spans // self.block_size)[:, None]
    return addrs.tolist(), sizes.tolist(), bids.tolist()
  • base = np.asarray(self.kv_caches_base_addr, dtype=np.int64): KV 캐시의 기본 주소를 NumPy 배열로 변환합니다.
  • blen = np.asarray(...): 각 KV 캐시 영역의 블록 길이를 NumPy 배열로 변환합니다. i % length를 사용하여 블록 길이가 순환하도록 처리합니다.
  • starts, spans: 입력된 chunks에서 시작 위치와 길이를 추출하여 NumPy 배열로 만듭니다.
  • bids: 각 청크에 해당하는 블록 ID를 계산합니다. starts // self.block_size를 통해 블록 인덱스를 얻고, block_ids 리스트에서 해당 값을 가져옵니다.
  • addrs = base[None, :] + bids[:, None] * blen[None, :]: 이 부분이 핵심적인 벡터화 연산입니다. NumPy의 브로드캐스팅 기능을 활용하여 모든 KV 캐시 영역에 대해 모든 청크의 주소를 한 번에 계산합니다. bids[:, None]bids를 열 벡터로 만들어 브로드캐스팅을 가능하게 합니다.
  • sizes = blen[None, :] * (spans // self.block_size)[:, None]: 마찬가지로, 모든 KV 캐시 영역에 대해 모든 청크의 크기를 한 번에 계산합니다.

2. vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/worker.py

이 파일은 KV 캐시 데이터를 저장하고 전송하는 워커 스레드의 로직을 담당합니다. 이번 PR에서는 _handle_request 메소드에서 prepare_values 함수를 활용하도록 수정되었습니다.

_handle_request (수정 전/후)

기존에는 각 청크마다 db.prepare_value를 호출했지만, 이제는 그룹별로 모은 청크들을 db.prepare_values로 한 번에 처리하도록 변경되었습니다.

# ... (이전 코드) ...

def _handle_request(self, req_meta: ReqMeta):
    # ... (중략) ...
    addrs: list[list[int]] = []
    sizes: list[list[int]] = []
    stored_events: list[BlockStored] = []
    chunks_per_group: list[list[tuple[int, int]]] = [
        [] for _ in self.token_databases
    ]
    for start, end, g_idx in zip(starts, ends, group_indices, strict=True):
        chunks_per_group[g_idx].append((start, end))
    for g_idx, chunks in enumerate(chunks_per_group):
        if not chunks:
            continue
        db = self.token_databases[g_idx]
        group_addrs, group_sizes, _ = db.prepare_values(
            chunks, block_ids_per_group[g_idx]
        )
        addrs.extend(group_addrs)
        sizes.extend(group_sizes)

    # parent_block_hash chains live within a group, not across.
    if self.enable_kv_event:
        prev_key_per_group: dict[int, Any] = {}
        for s, e, g_idx in zip(starts, ends, group_indices, strict=True):
            db = self.token_databases[g_idx]
            # addr, size, _ = db.prepare_value(
            #     start, end, req_meta.block_ids[g_idx]
            # )
            # addrs.append(addr)
            # sizes.append(size)

            if self.enable_kv_event:
                token_ids = (
                    req_meta.token_ids[s:e]
                )
                # ... (이하 생략) ...

# ... (이후 코드) ...

리뷰어 ivanium이 지적한 것처럼, 이 부분도 벡터화될 수 있는지 질문이 있었고, GirasoleY가 이를 prepare_values 함수를 사용하여 해결했음을 확인했습니다. 즉, _handle_request 내에서 개별 prepare_value 호출을 제거하고, 그룹별로 prepare_values를 호출하여 결과를 통합하는 방식으로 변경되었습니다.

3. 테스트 코드 추가

tests/v1/kv_connector/unit/test_mooncake_store_prepare_values.py 파일이 새로 추가되어, prepare_values 함수의 정확성과 성능 개선을 검증합니다. 특히 test_prepare_values_matches_reference 함수는 다양한 num_regionsnum_block_lens 설정에 대해 새로운 벡터화된 구현이 기존 스칼라 구현과 동일한 결과를 반환하는지 확인합니다.

왜 이게 좋은가

성능 향상

PR 설명에 따르면, 289K 토큰과 같이 긴 컨텍스트 요청의 경우, 2,258개의 블록을 처리할 때 prepare_value 함수의 측정 시간이 34.2ms에서 7.7ms로 약 4.4배 감소했습니다. 이는 스칼라 연산에서 NumPy를 이용한 벡터화 연산으로 전환함으로써 달성된 상당한 성능 향상입니다.

기존의 스칼라 방식은 Python 인터프리터의 오버헤드와 GIL로 인해 병렬 처리에 제약이 있었습니다. NumPy의 벡터화 연산은 C로 구현된 최적화된 루프를 사용하며, GIL을 해제하고 병렬 처리를 활용할 수 있어 CPU 집약적인 연산에서 훨씬 효율적입니다.

일반적인 교훈

  1. Python 스칼라 연산의 한계 인식: 반복문 기반의 Python 코드는 특히 대규모 데이터셋이나 반복적인 계산에서 성능 병목이 될 수 있습니다. NumPy와 같은 라이브러리를 활용하여 벡터화하는 것이 필수적입니다.
  2. NumPy의 브로드캐스팅 활용: NumPy의 브로드캐스팅 기능은 여러 배열 간의 연산을 효율적으로 수행하는 강력한 도구입니다. 이를 잘 활용하면 복잡한 루프를 단순화하고 성능을 극대화할 수 있습니다.
  3. 병렬 처리 및 GIL: Python의 GIL은 멀티스레딩을 통한 CPU 바운드 작업의 병렬화를 제한합니다. NumPy와 같이 GIL을 해제하는 라이브러리를 사용하거나 멀티프로세싱을 고려하는 것이 중요합니다.
  4. 테스트의 중요성: 성능 최적화는 정확성을 해치지 않아야 합니다. 새로운 테스트 케이스를 추가하여 변경 사항이 기존 동작과 일치하는지, 그리고 성능 향상이 실제로 달성되었는지 검증하는 것이 중요합니다.
  5. 리뷰 피드백 반영: 리뷰어의 질문(

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글