양자화된 KV 캐시

양자화된 KV 캐시 (Quantized KV Cache)

LLM을 돌리다 보면 KV 캐시(Key-Value cache)가 차지하는 메모리가 어느 순간 발목을 잡아요. 컨텍스트가 길수록 더 많은 토큰을 저장해야 하니까요. 이 KV 캐시를 FP8로 양자화하면 메모리 사용량을 크게 줄일 수 있고, 같은 메모리로 더 많은 토큰을 저장해서 처리량과 긴 컨텍스트 지원이 좋아져요. 이 페이지에서 vLLM의 양자화된 KV 캐시를 설정하는 방법을 살펴볼게요.

출처: vLLM 공식 문서 — Quantized KV Cache

기본 원리 (FP8 KV cache overview)

FP8 KV 캐시 양자화는 KV 캐시의 메모리 발자국을 유의미하게 줄여줘요. 참고로 Flash Attention 3 백엔드에 FP8 KV 캐시를 쓰면 어텐션 연산 자체도 양자화(FP8) 영역에서 수행돼요. 이 구성에서는 쿼리도 키·값과 함께 FP8로 양자화됩니다.

지원되는 FP8 KV 캐시 양자화 스킴 (Supported schemes)

vLLM은 FP8 KV 캐시에 두 가지 주요 양자화 전략을 지원해요.

  • 어텐션 헤드별(per-attention-head) 양자화: 각 헤드에 하나의 스케일. q_scale = [num_heads], k/v_scale = [num_kv_heads].
  • 텐서별(per-tensor) 양자화: 각 Q, K, V 텐서에 스케일 하나를 적용. q/k/v_scale = [1].

참고: 어텐션 헤드별 양자화는 현재 Flash Attention 백엔드에서만 사용할 수 있고, llm-compressor가 제공하는 캘리브레이션 경로가 필요해요.

스케일 캘리브레이션 방식 (Scale calibration approaches)

vLLM에서 양자화 스케일을 계산하는 방식은 세 가지를 고를 수 있어요.

  1. 캘리브레이션 없음(기본 스케일): 모든 양자화 스케일을 1.0으로 설정.
    kv_cache_dtype="fp8"
    calculate_kv_scales=False
    
  2. 랜덤 토큰 캘리브레이션(온더플라이): 워밍업 중 단일 랜덤 토큰 배치에서 스케일을 자동 추정 후 고정.
    kv_cache_dtype="fp8"
    calculate_kv_scales=True
    
  3. [권장] 데이터셋으로 캘리브레이션(llm-compressor 경유): 큐레이션된 캘리브레이션 데이터셋으로 최대 정확도를 얻는 방식. llm-compressor 라이브러리가 필요해요.

추가 kv_cache_dtype 옵션

  • kv_cache_dtype="fp8_e5m2": CUDA 11.8+ 지원
  • kv_cache_dtype="fp8_e4m3": CUDA 11.8+ 및 ROCm(AMD GPU) 지원
  • kv_cache_dtype="auto": 모델의 기본 데이터 타입 사용

예시 (Examples)

1. 캘리브레이션 없음

모든 양자화 스케일이 1.0으로 설정돼요.

from vllm import LLM, SamplingParams

sampling_params = SamplingParams(temperature=0.7, top_p=0.8)
llm = LLM(
    model="meta-llama/Llama-2-7b-chat-hf",
    kv_cache_dtype="fp8",
    calculate_kv_scales=False,
)

prompt = "London is the capital of"
out = llm.generate(prompt, sampling_params)[0].outputs[0].text
print(out)

2. 랜덤 토큰 캘리브레이션

워밍업 중 단일 토큰 배치에서 스케일을 자동 추정해요.

from vllm import LLM, SamplingParams

sampling_params = SamplingParams(temperature=0.7, top_p=0.8)
llm = LLM(
    model="meta-llama/Llama-2-7b-chat-hf",
    kv_cache_dtype="fp8",
    calculate_kv_scales=True,
)

prompt = "London is the capital of"
out = llm.generate(prompt, sampling_params)[0].outputs[0].text
print(out)

더 알아보기 (Learn more)