nanoVLM에서 KV 캐시 구현

Hugging Face는 비전 언어 모델을 학습하기 위한 간결한 PyTorch 코드베이스인 nanoVLM에서 KV(키-값) 캐싱을 처음부터 구현했습니다. 이 최적화는 자동회귀 추론 과정에서 중복 계산을 제거함으로써 생성 속도를 38% 향상시켰습니다.

자동회귀 생성에서의 계산 중복

자동회귀 언어 모델은 한 번에 하나의 토큰씩 텍스트를 생성합니다. 캐싱이 없는 표준 트랜스포머 구현에서는 다음 토큰을 예측하기 위해 전체 시퀀스—이전에 생성된 모든 토큰을 포함—를 처리해야 합니다.

트랜스포머는 내부적으로 병렬 처리되기 때문에, 새로운 토큰을 예측할 때마다 모든 레이어를 통한 전체 순방향 패스가 필요합니다. 이는 시퀀스 길이에 비례하여 메모리와 연산 요구량이 2차적으로 증가하게 합니다. 구체적으로, 모델은 매 단계마다 이전 토큰 전체에 대해 키(K)와 값(V) 텐서를 다시 계산하는데, 이 토큰들과 해당 투영은 변하지 않음에도 불구하고 재계산됩니다.

KV 캐싱이 추론을 최적화하는 방법

KV 캐싱은 초기 프롬프트가 처리된 후 각 레이어에 대해 계산된 키와 값을 저장함으로써 이 비효율성을 완화합니다. 전체 시퀀스를 다시 처리하는 대신, 모델은 다음과 같은 점진적인 워크플로를 따릅니다:

  1. 초기 상태 캐시: 첫 번째 패스 후, 각 레이어에 대해 계산된 $K$와 $V$가 캐시됩니다.
  2. 점진적 계산: 이후 생성 단계에서는 모델이 최신 토큰에 대해서만 $K$와 $V$를 계산합니다.
  3. 캐시 업데이트: 새로운 $K$와 $V$가 기존 캐시에 추가됩니다.
  4. 어텐션 계산: 현재 토큰에 대한 Query($Q$)가 캐시된 $K$와 $V$와 함께 사용되어 출력이 생성됩니다.

실제로 이 캐시는 각 레이어별 사전(dictionary) 형태로 유지되며, "key"와 "value" 텐서는 (batch_size, num_heads, seq_len_cached, head_dim) 형태를 가집니다.

nanoVLM에서의 기술 구현

nanoVLM의 구현은 전체 시퀀스 재계산에서 점진적 업데이트 시스템으로 전환하기 위해 세 가지 주요 구성 요소에 걸친 수정이 포함됩니다.

1. 어텐션 블록 업데이트

LanguageModelGroupedAttention 클래스에서 forward 함수가 block_kv_cache를 받도록 수정되었습니다. 캐시가 존재하면(모델이 프리필 단계가 아님을 의미) 현재 토큰에 대해 $K_{new}$와 $V_{new}$를 계산하고 이를 캐시된 텐서와 연결합니다. 캐시가 없으면 프롬프트에 대한 초기 계산을 수행합니다.

2. 레이어별 캐시 추적

LanguageModel 클래스가 이제 레이어별 캐시 추적을 구현합니다. start_pos 인자를 사용하여 로터리 위치 인코딩이 현재 생성 인덱스와 정확히 맞춰지도록 하여, 모델이 새로 생성된 토큰의 절대 위치를 시퀀스에 대해 알 수 있게 합니다.

3. 생성 루프 분리

VisionLanguageModelgenerate() 메서드가 두 개의 뚜렷한 단계로 분리되었습니다:

  • 프리필 단계: 모델이 전체 입력 프롬프트를 인코딩하고 모든 레이어에 대한 초기 KV 캐시를 구성합니다.
  • 디코드 단계: 모델이 토큰을 순차적으로 생성하며, 캐시된 키와 값을 사용해 프롬프트와 이전에 생성된 토큰을 다시 처리하지 않습니다.

아키텍처 변경 요약

모듈 기존 동작 새로운 동작
LanguageModelGroupedAttention.forward 매 단계마다 $Q$, $K$, $V$를 재계산 KV 캐시를 사용하고 업데이트
LanguageModel.forward 이전 상태를 기억하지 않음 레이어별 KV 캐시를 추적하고 start_pos를 처리
VisionLanguageModel.generate 단일 단계 생성 루프 프리필디코드 단계로 분리

트레이드오프 및 영향

KV 캐싱은 토큰당 추론 복잡도를 2차에서 $O(\text{seq len})$으로 감소시켜, 더 빠른 추론과 소비자 하드웨어에서 대형 모델을 실행할 수 있게 합니다. 그러나 이러한 효율성은 트레이드오프를 동반합니다: 캐시를 저장하기 위한 메모리 사용량이 증가하고 코드 복잡도가 높아집니다. 또한 빔 서치와 같이 더 복잡한 캐시 관리가 필요한 일부 추론 방식은 제한될 수 있습니다.

Sources