인공지능/논문 리뷰 or 진행

Which Heads Matter for Reasoning? RL-Guided KV Cache Compression

이게될까 2026. 8. 4. 20:46
728x90
728x90

https://arxiv.org/abs/2510.08525

 

Which Heads Matter for Reasoning? RL-Guided KV Cache Compression

Reasoning large language models exhibit complex reasoning behaviors via extended chain-of-thought generation that are highly fragile to information loss during decoding, creating critical challenges for KV cache compression. Existing token-dropping methods

arxiv.org

 

ICLR은 떨어지고, ICML 2026에 붙은 것 같네요

 

Reasoning LLM에서 KV cache는 계산 확인하거나, 중간 결과를 유지, 풀이 탐색 등의 모든 토큰이 저장되어 출력이 길어질 수록 캐시 사용량도 선형적으로 증가함. 

KV캐시는 단순 이전 문장을 기억하는 것이 아니라 현재 까지 추론 상태와 이전에 세운 가정과 중간 결과, 추론을 계속 할지, 종료할지에 대한 흐름, 자기 수정과 이전 단계 참조에 필요한 정보를 모두 담는다고 합니다. 

그래서 일반 LLM에서 잘 작동하던 KV 캐시 압축을 reasoning model에 적용하면 추론 과정 자체가 붕괴될 수 있다고 합니다. 

 

기존 방법들의 문제를 보여주네요 

토큰 드랍하는 방법은 지금은 중요하지 않아도 나중에 다시 필요한 토큰이 존재할 수 있는데 그 것을 없애버리고, Repetitive error가 발생한다고 합니다. 

모든 head에서 토큰을 제거하는 대신 일부 중요한 head에선 full KC cache를 할당하고 나머지 해드에는 작은 캐시만 할당하는 방법은 token dropping보다는 전체 sequence정보를 더 잘 보존할 수 있으나, Retrieval head는 긴 입력에서 특정 정보만 다시 찾아오는데 중요한 헤드지 reasoning에는 단순 정보 검색 외의 기능도 필요하기에 reasoning에 중요한 헤드와 일치하지 않게 된다. 

이 논문에서는 어떤 attneiton head가 실제 reasoning behavior를 유지하는데 중요한지 확인하려고 합니다. 

그러기 위해 특정 헤드의 kv cache를 압축했을 때 실제 autoregressive generation의 최종 정답과 추론 행동이 어떻게 달라지는지 확인합니다. 

각 attention head에 gate를 추가하고, full attention와 local attention의 gate 합으로 더하여 진행한다. 1은 과거 전체 KV가 필요한 head고, 0은 최근 토큰만 있어도 동작하는 head다. 

기존 하이퍼 파라미터는 고정하고, L * H 개의 gate만 학습하게 된다. 

수학 문제에 대해 여러 응답을 생성하고, 최종 정답을 검사해 reward를 계산해 캐쉬 압축으로 일어나는 작은 오류가 이후 어떻게 누적되는지 리워드에 반영함. 

이제 GRPO를 통해 최대한 많은 헤드를 0으로 만들게 하고, 정답은 맞출 수 있도록 유지함. 단순히 RL reward와 L1 penalty(헤드 0으로 만드는 긋)만으로는 학습이 쉽게 붕괴함

붕괴가 발생하는 이유

  1. L1 penalty가 gate 값을 감소시킴
  2. 너무 많은 head가 local cache로 전환됨
  3. reasoning 성능이 저하됨
  4. 정답 reward가 거의 나오지 않음
  5. reward 신호는 약해지지만 L1 penalty는 계속 모든 gate에 적용됨
  6. gate가 더 작아짐
  7. 모델이 회복하지 못함

이를 위해 두 가지 방법을 사용함

1. 문제를 그대로 사용하기 보단 맞춘 문제만 선별하여 출력 길이에 따라 3000개를 구성하여 출게함. 

2. 평균 리워드가 낮아니면 L1 penalty를 자동으로 줄여 모델이 회복할 수 있도록 함. 

추론할 때는 이제 이진화 시켜 gate 값이 높은 상위 head만 full KV 진행하고, 나머지는 오래된 KV를 저장하지 않는 등 진행함. 

 

 

  R-KV RLKV
압축 단위 각 head 내부의 토큰 attention head
중요도 기준 attention 및 token redundancy 실제 reasoning rollout의 정답 reward
모든 head 압축 여부 대부분의 head에서 토큰 제거 중요한 head는 전혀 압축하지 않음
보존 대상 중요하다고 판단된 토큰 reasoning-critical head의 전체 history
대표 오류 반복 루프, 추론 일관성 붕괴 중간 sparsity에서는 상대적으로 안정적
학습 필요 기본적으로 heuristic 중심 모델별 gate RL 학습 필요
장점 별도 full-cache head를 두지 않아 단순함 reasoning 성능을 더 안정적으로 보존
단점 중요한 reasoning token을 제거할 위험 사전 학습 비용과 정적 head selection 필요

 

LLama와 Qwen2.5는 모든 헤드의 중요도가 높게 나타났다. 

Qwen 3는 중요한 헤드와 압축 가능한 헤드가 섞여 있었다. 

RL 학습이 붕괴하는 모습을 보여주며 그냥 진행하면 gate평균 값이 낮아질수록 모델이 망가지지만, 파란색은 gate도 낮추며 모델 성능도 최대한 유지했다. 

압축률이 높을수록 성능이 낮아지긴 하지만 이 논문의 방식이 잘 버티는 것을 볼 수 있었다. 

대부분 모델에서 20 ~ 40%가 안정적이었지만 일부 테스크에서는 50 ~ 60%까지 성공하는 부분도 있었다. 

Ablation에서는 Adaptibe penalty Weighting, Self-distillation sampling, L1 weight를 실험하여 RL 자체 뿐이 아니라 다른 요인에 크게 의존하는 것을 보여준다. 

Table3는 속도 향상과 정확도를 부여주며 단일 요청이 아니라 KV메모리 절약을 통해 동시 실행할 수 있는 요청 수가 증가하여 얻은 개선이다. 

table4는 Sink와 Local Window 크기의 영향을 보여주며, 적게 썼을 때와 크게 썼을 때의 차이를 보여주며 클 수록 높은 압축률에도 더 좋은 성능을 보여준다. 

문제 정의 Reasoning LLM은 긴 Chain-of-Thought를 생성하므로 KV cache 메모리와 추론 비용이 크게 증가한다.
그러나 기존 KV cache 압축을 적용하면 중간 추론 정보가 손실되어 반복 생성, 오답, 지나치게 긴 추론이 발생한다.
기존 Token-dropping의 한계 H2O, R-KV처럼 각 head 내부에서 중요도가 낮은 토큰을 제거하면, 현재는 중요하지 않아 보이지만 이후 추론에서 다시 필요한 중간 상태까지 삭제될 수 있다.
이 경우 추론 흐름이 끊기고 동일 문장이나 계산을 반복하는 repetitive error가 주로 발생한다.
기존 Head-reallocation의 한계 DuoAttention, KVZip은 일부 head에 full KV cache를 할당하지만, 주로 long-context retrieval에 중요한 retrieval head를 기준으로 선택한다.
Retrieval head는 정보 검색에는 중요하지만 CoT 일관성, 추론 진행 및 종료를 보존하는 head와 일치하지 않을 수 있다.
핵심 가설 모든 attention head가 전체 과거 토큰의 KV cache를 필요로 하는 것은 아니다.
일부 reasoning-critical head만 full history를 필요로 하며, 나머지 head는 초기 sink token과 최근 token만 유지해도 된다.
Reasoning-critical head의 정의 Full KV cache 대신 local KV cache를 사용했을 때 reasoning 성능이 크게 감소하는 head.
논문은 이 head들이 CoT consistency와 generation termination에 중요한 것으로 해석한다.
핵심 방법: RLKV 각 layer와 KV head에 학습 가능한 gate α_{l,h}를 추가한다.
각 head의 출력은 α⋅Full Attention+(1−α)⋅Local Attention으로 계산된다.
α가 높을수록 해당 head가 full KV cache에 의존한다는 의미다.
왜 강화학습을 사용하는가? Attention score나 next-token loss 같은 정적 proxy 대신, 압축된 상태에서 모델이 실제로 생성한 autoregressive CoT의 최종 정답 여부를 직접 관찰하기 위해서다.
이를 통해 초기의 작은 압축 오류가 긴 생성 과정에서 누적되는 영향까지 반영한다.
RL 학습 방식 LLM 본체는 고정하고 L × H개의 gate만 GRPO로 학습한다.
정답을 생성한 rollout에는 높은 reward를 주어 필요한 head의 gate를 유지하고, L1 penalty는 불필요한 gate를 0에 가깝게 만들어 full-cache head 수를 줄인다.
학습 목적식의 의미 .
Reward는 추론 능력을 보존하고, L1 regularization은 full KV cache 사용 head를 최소화한다.
두 신호의 경쟁을 통해 reasoning-critical head가 선택된다.
Self-distillation sampling 모델이 full KV cache 상태에서 이미 맞힐 수 있는 DeepScaleR 문제만 선별하고, 출력 길이에 따라 3,000개를 구성한다.
이 연구의 목표는 새로운 추론 능력 학습이 아니라 기존 능력을 압축 후에도 보존할 head를 찾는 것이기 때문이다.
Adaptive penalty weighting 압축이 과도해져 reward가 떨어지면 L1 penalty를 약화하거나 제거한다.
이는 sparse reward가 사라진 상태에서 dense L1 penalty만 계속 gate를 0으로 밀어 학습이 붕괴하는 것을 방지한다.
추론 시 적용 방식 학습된 gate를 기준으로 상위 head를 reasoning-critical head로 선택한다.
이 head들은 전체 KV cache를 유지하고, 나머지 head들은 기본 설정에서 첫 16개 sink token과 최근 64개 local token만 저장한다.
R-KV와의 핵심 차이 R-KV는 각 head에서 어떤 토큰을 삭제할지 결정한다.
RLKV는 어떤 head에서는 과거 토큰을 전부 보존해야 하는지 결정한다.
따라서 중요한 reasoning head 내부의 중간 정보를 임의로 삭제하지 않는다.
평가 모델 Llama-3.1-8B-R1, Qwen-2.5-7B-R1, Qwen-3-4B-Thinking.
평가 태스크 수학 추론: GSM8K, Math500, AIME24 / 코드: MBPP / 지식 추론: MMLU-Pro의 Chemistry, CS, Law, Physics / 장문 추론: 최대 70K context의 LongReason.
비교 방법 Token-dropping: H2O, R-KV / Head-reallocation: DuoAttention, KVZip.
주요 정확도 결과 모델과 태스크에 따라 20–60% KV cache budget sparsity에서 near-lossless 성능을 달성했다.
다수 설정에서 기존 방법보다 높은 정확도를 보였으며, 일부 결과는 full KV cache baseline과 같거나 약간 높았다.
대표 결과 Llama-3.1-8B-R1 Math500에서 40% sparsity로 정확도 84.6%를 기록해 full baseline보다 1.6%p 높았다.
Qwen-2.5-7B-R1 GSM8K에서는 40% sparsity에서 90.1%, Qwen-3-4B-Thinking Math500에서는 60% sparsity에서 75.6%로 full 대비 2.0%p 감소에 그쳤다.
Long-context 일반화 Gate는 최대 8K-token rollout으로 학습했지만 70K context의 LongReason에서도 기존 방법보다 우수했다.
이는 특정 토큰 위치보다는 full history가 필요한 head의 특성을 학습했을 가능성을 보여준다.
시스템 효율성 SGLang에 full-head용 paged KV pool과 compressed-head용 고정 크기 circular buffer를 구현했다.
절약한 메모리를 더 많은 동시 요청에 사용해 continuous batching 처리량을 높였다.
속도 결과 Llama-3.1-8B-R1 Math500에서 40% sparsity는 정확도를 유지하면서 1.56× end-to-end speedup, 60% sparsity는 2.06× speedup을 달성했다.
다만 60%에서는 정확도가 79.4%에서 73.8%로 감소했다.
속도 결과의 올바른 해석 2.06배는 단일 요청 latency가 그대로 절반이 되었다는 뜻이 아니다.
KV cache 절감으로 동시 처리 요청 수를 150개에서 375개로 늘려 얻은 serving-level throughput 개선이다.
Head sensitivity 분석 RLKV가 높은 점수를 부여한 head부터 압축하면 retrieval head나 random head를 압축할 때보다 정확도가 더 빠르게 감소했다.
이는 선택된 head가 실제 reasoning 성능에 민감하다는 근거다.
오류 유형 분석 Token-dropping 및 reasoning-critical head 압축은 주로 repetitive/incorrect error를 유발했다.
Retrieval head 기반 압축은 문장은 유창하지만 결론에 도달하지 못하는 overlength error가 상대적으로 많았다.
논문의 주요 발견 Reasoning-critical head는 retrieval head와 기능적으로 다르다.
전자는 단순 정보 검색보다 CoT 일관성 유지, 추론 진행, 반복 방지, 생성 종료와 더 밀접한 것으로 나타났다.
주요 기여 ① RL을 reasoning head 탐색용 on-policy probe로 사용
② retrieval head와 reasoning-critical head의 차이를 실험적으로 제시
③ head-level cache 압축을 실제 SGLang serving speedup으로 연결.
한계 1 Gate와 head 선택이 학습 이후 고정되는 static allocation이다.
문제 유형이나 query에 따라 중요한 head가 달라질 수 있지만 query-adaptive gating은 다루지 않았다.
한계 2 모델마다 reasoning-critical head 분포가 달라 별도의 RL 학습이 필요하다.
모델별 학습 비용은 약 22–40 GPU-hours이며, 학습된 head mask의 모델 간 전이 가능성은 검증하지 않았다.
한계 3 80% 수준의 극단적인 sparsity에서는 대부분 성능이 크게 붕괴한다.
더 높은 압축률을 달성하려면 KV quantization 등 다른 압축 기법과의 결합이 필요하다.
한계 4 수학 문제처럼 최종 정답을 자동 검증할 수 있는 reward를 사용했다.
Open-ended generation, agent, 주관적 평가 태스크에서 어떤 reward를 사용할지는 추가 연구가 필요하다.
최종 핵심 메시지 Reasoning LLM의 KV cache를 안전하게 압축하려면 모든 head에서 토큰을 조금씩 삭제하기보다, 실제 생성 결과를 통해 전체 history가 반드시 필요한 소수의 head를 찾아 보호해야 한다.
RLKV는 이를 통해 중간 수준의 cache 절감에서 reasoning 성능과 serving 효율을 함께 확보한다.
728x90