KV 캐시와 생성 속도
KV 캐시를 켰을 때와 껐을 때의 언어 모형 생성 속도를 비교한다.
KV 캐시
- KV 캐시(KV cache): 앞 토큰들의 주의 계산에 쓴 Key와 Value를 저장해 다음 토큰 생성에 재사용
- 재사용이 가능한 이유: 새 토큰이 추가되어도 앞 토큰들의 Key와 Value는 바뀌지 않음
use_cache=False: 다음 토큰을 생성할 때 앞선 토큰의 Key와 Value를 다시 계산use_cache=True: 저장한 Key와 Value를 사용해 반복 계산을 줄임- 특징:
- 입력과 생성 길이가 길수록 속도 향상 효과가 커짐
- 과거 토큰의 Key와 Value를 저장하므로 메모리 사용량이 증가
실습 준비
!pip install -q "transformers==4.55.4" accelerate jinja2
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
model_name = "LGAI-EXAONE/EXAONE-4.0-1.2B" # 허브 저장소 식별자
tokenizer = AutoTokenizer.from_pretrained(model_name) # 토크나이저 로드
model_dtype = torch.float16 if torch.cuda.is_available() else torch.float32
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=model_dtype, # GPU는 float16, CPU는 float32 사용
device_map="auto", # 사용 가능한 장치에 자동 배치
)
비교할 입력
- 비교 조건: 탐욕 탐색으로 새 토큰을 64개씩 생성하고 KV 캐시 사용 여부만 변경
prompt = "인공지능은 데이터를 학습하여 패턴을 찾는다. 이를 활용하면" # 공통 입력
inputs = tokenizer(
prompt,
return_tensors="pt", # 토큰 ID를 PyTorch 텐서로 반환
).to(model.device) # 입력을 모형과 같은 장치로 이동
시간 측정
perf_counter(): 코드 실행 시간 측정용 타이머torch.cuda.synchronize(): 예약된 GPU 연산이 끝날 때까지 기다림- GPU 연산은 비동기로 실행되므로 측정 시작과 종료 직전에 호출
from time import perf_counter
def measure_generation(use_cache, new_tokens):
if torch.cuda.is_available():
torch.cuda.synchronize() # 이전 GPU 연산이 끝난 뒤 측정 시작
start = perf_counter()
with torch.inference_mode(): # 기울기 계산을 끄고 생성 시간만 측정
output_ids = model.generate(
**inputs, # 같은 입력 사용
min_new_tokens=new_tokens, # 종료 토큰이 나와도 정한 길이까지 생성
max_new_tokens=new_tokens, # 캐시 조건별 생성 길이 고정
do_sample=False, # 탐욕 탐색 사용
num_beams=1, # 문장 후보 하나만 유지
use_cache=use_cache, # KV 캐시 사용 여부
)
if torch.cuda.is_available():
torch.cuda.synchronize() # 예약된 GPU 연산이 끝난 뒤 측정 종료
elapsed_seconds = perf_counter() - start
return output_ids, elapsed_seconds
캐시를 켜고 끄기
- 워밍업(warm-up): 연산 초기화 시간을 측정에서 제외하기 위한 짧은 사전 실행
_ = measure_generation(use_cache=False, new_tokens=4) # 캐시를 끈 워밍업
_ = measure_generation(use_cache=True, new_tokens=4) # 캐시를 켠 워밍업
uncached_ids, uncached_seconds = measure_generation(
use_cache=False,
new_tokens=64,
)
cached_ids, cached_seconds = measure_generation(
use_cache=True,
new_tokens=64,
)
uncached_seconds # 캐시를 끈 생성 시간(초)
실행 결과
7.945843660999969
cached_seconds # 캐시를 켠 생성 시간(초)
실행 결과
5.175696586000015
uncached_seconds / cached_seconds # 시간 비율(끔/켬)
실행 결과
1.5352220766752547
- 실행 결과:
- 캐시를 끈 경우: 약 7.95초
- 캐시를 켠 경우: 약 5.18초
- 시간 비율: 캐시를 켠 쪽이 약 1.54배 빠른 결과
- 차이가 크지 않은 이유: 입력 16개, 생성 64개 토큰으로 문맥이 짧아 다시 계산하는 양이 적음
- 입력이나 생성이 길어질수록 캐시를 껐을 때 다시 계산하는 양이 늘어 차이가 커짐
생성 결과 비교
- 캐시의 역할: 중복 계산을 줄이는 최적화
- 토큰 선택 규칙: KV 캐시 사용 여부와 무관
torch.equal(uncached_ids, cached_ids) # 생성 토큰 일치 여부
실행 결과
True
- 생성 토큰: 두 조건에서 모두 일치
- 해석: KV 캐시는 탐욕 탐색의 선택 규칙을 바꾸지 않고 생성의 반복 계산을 줄임