캐싱이 어떻게 동작하나요
캐싱이 어떻게 동작하나요 (How caching works)
누군가와 대화할 때, 상대가 여러분이 이전에 했던 말을 기억하지 못하고 매번 처음부터 다시 시작한다고 상상해 보세요. 느리고 비효율적이겠지요? 이 비유를 transformer 모델에 그대로 확장할 수 있어요. 자기회귀(autoregressive) 모델 생성은 한 번에 하나의 토큰씩 예측하기 때문에 느릴 수 있어요. 각 새 예측은 이전의 모든 컨텍스트에 의존하지요.
출처: 문서
본문
1000번째 토큰을 예측하려면 모델은 이전 999개 토큰의 정보가 필요해요. 그 정보는 토큰 표현들 간의 행렬 곱셈으로 표현돼요.
1001번째 토큰을 예측하려면 1000번째 토큰의 정보에 더해, 이전 999개 토큰의 동일한 정보가 또 필요해요. 모델이 토큰마다 계속해서 반복 계산해야 하는 행렬 곱셈이 정말 많지요!
키-값(KV) 캐시는 이전에 처리된 토큰의 attention 레이어에서 얻은 kv 쌍을 저장해서 이런 비효율을 없애 줘요. 저장된 kv 쌍은 캐시에서 꺼내어 이후 토큰에 재사용되므로, 다시 계산할 필요가 없어져요.
[!WARNING] 캐싱은 추론(inference) 용으로만 사용해야 해요. 학습 중에 켜두면 예상치 못한 오류가 발생할 수 있어요.
캐싱이 어떻게, 왜 동작하는지 이해하려면 attention 행렬의 구조를 자세히 살펴볼게요.
Attention 행렬 (Attention matrices)
scaled dot-product attention은 배치 크기 b, attention 헤드 수 h, 여기까지의 시퀀스 길이 T, 헤드당 차원 d_head에 대해 아래처럼 계산돼요.
$$ \text{Attention}(Q, K, V) = \text{softmax}\left( \frac{Q K^\top}{\sqrt{d_{\text{head}}}} \times \text{mask} \right) V $$
쿼리(Q), 키(K), 밸류(V) 행렬은 shape이 (b, h, T, d_head)인 입력 임베딩의 투영(projection)이에요.
인과(causal) attention에서 mask는 모델이 미래 토큰을 참조하지 못하게 해요. 토큰이 처리되면 그 표현은 미래 토큰에 대해 절대 변하지 않으므로, $ K_{\text{past}} $와 $ V_{\text{past}} $를 캐시에 저장해 마지막 토큰의 표현을 계산할 때 재사용할 수 있어요.
$$ \text{Attention}(q_t, [\underbrace{k_1, k_2, \dots, k_{t-1}}{\text{cached}}, k{t}], [\underbrace{v_1, v_2, \dots, v_{t-1}}{\text{cached}}, v{t}]) $$
추론 시점에는 다음 토큰 $ t+1 $을 예측하는 표현 $ x_t $를 계산하기 위해 마지막 토큰의 쿼리만 있으면 돼요. 각 단계에서 새 키와 밸류 벡터는 캐시에 저장되고 이전 키와 밸류에 추가(append) 돼요.
$$ K_{\text{cache}} \leftarrow \text{concat}(K_{\text{past}}, k_t), \quad V_{\text{cache}} \leftarrow \text{concat}(V_{\text{past}}, v_t) $$
attention은 모델의 각 레이어에서 독립적으로 계산되고, 캐싱도 레이어별로 이루어져요.
캐싱이 효율을 어떻게 개선하는지 아래 표로 비교해 볼게요.
| 캐싱 없이 | 캐싱 사용 시 |
|---|---|
매 단계마다 이전 모든 K와 V를 다시 계산 |
매 단계마다 현재 K와 V만 계산 |
| 단계당 attention 비용은 시퀀스 길이에 대해 이차(quadratic) | 단계당 attention 비용은 시퀀스 길이에 대해 선형(linear) (메모리는 선형으로 증가하지만 토큰당 연산은 낮게 유지) |
Cache 클래스
기본 KV 캐시 인터페이스는 현재 토큰의 키·밸류 텐서를 받아 업데이트된 K와 V 텐서를 반환해요. 이는 모델의 forward 메서드가 내부적으로 관리해요.
new_K, new_V = cache.update(k_t, v_t, layer_idx)
attn_output = attn_layer_idx_fn(q_t, new_K, new_V)
Transformers의 Cache 클래스를 사용하면, self-attention 모듈이 과거와 현재 정보를 통합하기 위해 몇 가지 중요한 단계를 수행해요.
-
attention 모듈은 현재 kv 쌍을 캐시에 저장된 과거 kv 쌍과 연결(concatenate)해요. 이렇게 하면 shape이
(new_tokens_length, past_kv_length + new_tokens_length)인 attention 가중치가 만들어져요. 현재와 과거 kv 쌍은 본질적으로 결합되어 attention 점수를 계산하며, 모델이 이전 컨텍스트와 현재 입력을 모두 인식하게 해요. -
forward메서드를 반복적으로 호출할 때 attention mask shape이 과거·현재 kv 쌍의 결합 길이와 일치하는 게 매우 중요해요. attention mask의 shape은(batch_size, past_kv_length + new_tokens_length)여야 해요. 이는 보통 generate() 안에서 내부적으로 처리되지만, Cache로 나만의 생성 루프를 구현하려면 이 점을 꼭 기억하세요! attention mask는 과거와 현재 토큰 값을 모두 담아야 해요.
Cache 저장 구현 (Cache storage implementation)
캐시는 레이어의 리스트로 구성되며, 각 레이어는 키와 밸류 캐시를 담고 있어요. 키·밸류 캐시는 shape이 [batch_size, num_heads, seq_len, head_dim]인 텐서예요.
레이어는 서로 다른 타입일 수 있어요(예: DynamicLayer, StaticLayer, StaticSlidingWindowLayer). 이는 주로 시퀀스 길이를 어떻게 처리하고 캐시를 어떻게 업데이트하는지에 영향을 줘요.
가장 단순한 건 DynamicLayer로, 더 많은 토큰이 처리됨에 따라 커져요. 새 토큰마다 시퀀스 길이 차원(seq_len)이 증가해요:
cache.layers[idx].keys = torch.cat([cache.layers[idx].keys, key_states], dim=-2)
cache.layers[idx].values = torch.cat([cache.layers[idx].values, value_states], dim=-2)
StaticLayer와 StaticSlidingWindowLayer 같은 다른 레이어 타입은 캐시가 생성될 때 설정되는 고정 시퀀스 길이를 가져요. 그래서 torch.compile과 호환돼요. StaticSlidingWindowLayer의 경우 새 토큰이 추가되면 기존 토큰이 캐시 밖으로 밀려나요.
아래 예시는 DynamicCache로 생성 루프를 만드는 방법을 보여줘요. 앞서 얘기했듯이 attention mask는 과거와 현재 토큰 값의 연결이에요.
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, DynamicCache
from accelerate import Accelerator
device = Accelerator().device
model_id = "meta-llama/Llama-2-7b-chat-hf"
model = AutoModelForCausalLM.from_pretrained(model_id, dtype=torch.bfloat16, device_map=device)
tokenizer = AutoTokenizer.from_pretrained(model_id)
past_key_values = DynamicCache(config=model.config)
messages = [{"role": "user", "content": "Hello, what's your name."}]
inputs = tokenizer.apply_chat_template(messages, add_generation_prompt=True, return_tensors="pt", return_dict=True).to(model.device)
generated_ids = inputs.input_ids
max_new_tokens = 10
for _ in range(max_new_tokens):
outputs = model(**inputs, past_key_values=past_key_values, use_cache=True)
# Greedily sample one next token
next_token_ids = outputs.logits[:, -1:].argmax(-1)
generated_ids = torch.cat([generated_ids, next_token_ids], dim=-1)
# Prepare inputs for the next generation step by leaving unprocessed tokens, in our case we have only one new token
# and expanding attn mask for the new token, as explained above
attention_mask = inputs["attention_mask"]
attention_mask = torch.cat([attention_mask, attention_mask.new_ones((attention_mask.shape[0], 1))], dim=-1)
inputs = {"input_ids": next_token_ids, "attention_mask": attention_mask}
print(tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0])
"[INST] Hello, what's your name. [/INST] Hello! My name is LLaMA,"