어텐션 값 투영을 레이어별 토큰 메모리로 대체하는 Memory Attention
Memory Attention
무엇인가
이 논문은 언어모델이 어텐션 값을 만들 때 항상 현재 문맥의 은닉 상태에서 별도 값 투영 W_V를 거치는 관행을 문제로 삼는다. 일부 내용은 문맥과 무관하게 재사용될 수 있는데도 매번 밀집 행렬곱으로 값을 만든다는 것이다. 제한된 계산과 GPU 메모리에서 품질을 올리려면 조회 기반 표현으로 파라미터 용량을 늘리되 밀집 계산을 비례적으로 늘리지 않는 방법이 필요하다. 기존 Value Embedding, DeepEmbed, PLE, STEM, Engram 등은 값을 보강하거나 메모리를 추가하는 방향이었고, 이 논문은 한 걸음 더 나아가 기존 값 투영 자체를 명시적 메모리로 대체할 수 있는지 묻는다.
어떻게 동작하나
Memory Attention(MA)은 각 레이어마다 어휘 크기 N과 값 차원 d_v를 갖는 학습 가능한 메모리 테이블 E를 둔다. 입력 토큰 ID s로 E[s]를 조회하고, 각 키/값 헤드 안에서 토큰 벡터별로 RMSNorm을 적용해 M을 만든다. 그다음 Q=XW_Q, K=XW_K, V=K+M으로 값을 구성한다. 즉 별도 W_V를 제거하고 키 표현을 값의 문맥 성분으로 재사용하며, 메모리가 토큰 특화 성분을 더한다. RoPE를 쓰면 Q^R=RoPE(Q), K^R=RoPE(K), V=K+M으로, 위치 변환 전의 K로 값을 만들고 Q^R과 K^R로 어텐션 가중치를 계산한다. 어텐션의 가중치와 집계 규칙은 그대로 둔다. 같은 토큰이라도 레이어마다 다른 메모리 테이블을 쓰므로 깊이별로 다른 표현을 학습할 수 있다.
무엇과 다른가
저자들은 MLA와의 비교를 통해 MA의 위치를 설명한다. MLA는 압축 잠재 C를 어텐션 점수와 내용 집계에 공유하고, 집계 후 W_V를 적용한다. MA는 위치 변환 전 키 K를 공유 표현으로 쓰고, A_MA(K+M)=A_MA K + A_MA M처럼 같은 어텐션 가중치 아래 키와 토큰 메모리를 함께 집계한다. MLA가 학습된 선형 값 매핑을 적용하는 것과 달리 MA는 키를 직접 재사용하고 메모리를 덧셈으로 읽는다. 추론 시에는 정규화를 테이블에 미리 접어 넣을 수 있어 E_bar[i]=Norm(E[i]), V=K+E_bar[s]가 되고, 온라인 값 구성이 조회와 원소별 덧셈만 남는다. 파라미터는 표준 값 투영의 L d d_v 대신 L(N d_v + p_norm)이 되어 N>d이면 총 파라미터 저장량이 늘어난다. 학습 계산은 표준 값 투영의 약 6 L S d d_v FLOPs에 비해 O(L S d_v) 수준이고, 온라인 값 구성 감소량은 대략 ΔF_V ≈ L S d_v (2d-1)로 제시된다. 다만 이 산술 감소가 지연 시간 감소를 직접 보장하지는 않는다.
어떻게 쓰나
MA-Offload는 접힌 메모리 테이블을 CPU 메모리에 두고, 조회 주소가 토큰 ID와 레이어 인덱스에만 의존한다는 점을 이용해 해당 은닉 상태가 준비되기 전에 필요한 벡터를 GPU로 프리페치한다. 이는 b_w L N d_v 바이트의 메모리 테이블 파라미터를 GPU에서 CPU로 옮기지만, 총 파라미터 수와 값 구성 산술은 바꾸지 않는다. 캐싱이나 중복 제거가 없을 때 새로 처리하는 S개 토큰의 논리적 전송량은 D_Offload = b_w L S d_v이고, 일반 KV 캐시를 쓰면 디코딩 스텝마다 새 토큰에 대해서만 b_w L B d_v 바이트를 가져온다. 저자들은 오프로딩이 GPU 파라미터 상주량을 줄이지만 총 GPU 메모리 사용량이나 추론 지연 시간 감소를 보장하지는 않는다고 명시한다. MA-Recall은 값이 K+E_bar[s]로 표현된다는 점을 이용해 과거 값을 재구성하는 확장이다. 영구 값 캐시를 없애는 대신 키와 토큰 ID를 유지하고, 기존 콘텐츠 KV 캐시 2 b_c L B T d_v 바이트를 키 표현 하나만 남겨 b_c L B T d_v 바이트로 50% 줄일 수 있다고 분석한다. 그러나 재구성에는 대략 L B R d_v 번의 덧셈과, 메모리 테이블까지 CPU 오프로드할 경우 D_Recall = b_w L B R d_v의 전송량이 필요하다. MA-Recall은 이 논문의 추론 측정에는 포함되지 않았다.
전제와 한계
실험은 NVIDIA H800에서 flash-linear-attention 프레임워크로 진행됐고, 어텐션 블록과 gated MLP, RoPE, RMSNorm을 사용했다. Standard와 MA는 동일한 학습 토큰 예산을 쓰되 MA는 추가 메모리 파라미터를 갖는다. 제로샷 평가는 lm-evaluation-harness로 LAMBADA와 WikiText 퍼플렉시티, LAMBADA 단어 예측 정확도, ARC-Easy, ARC-Challenge, HellaSwag, PIQA, WinoGrande, OpenBookQA를 측정했다. Table 2에서 MA는 세 구성 모두에서 두 퍼플렉시티 지표와 평균 정확도를 개선했다. 작은 MHA 모델의 WikiText 퍼플렉시티는 31.55에서 28.64로, GQA의 LAMBADA 퍼플렉시티는 50.46에서 43.68로 낮아졌다. 평균 정확도 이득은 작은 MHA 0.71%p, GQA 0.61%p, 큰 MHA 1.16%p였다. 큰 모델에서는 OpenBookQA가 3.00점, ARC-Challenge가 2.65점 올랐지만 ARC-Easy와 PIQA는 약간 낮아졌다. 즉 모든 벤치마크에서 일관된 향상은 아니다.
검색 실험은 24층, 은닉 차원 1,024, 10B 토큰, 컨텍스트 길이 2,048로 학습한 모델에서 세 가지 single-needle NIAH 과제로 수행됐다. 평가 길이는 1,024, 2,048, 4,096 토큰이며 앞의 둘은 학습 창 안, 마지막은 두 배 외삽이다. Table 3에서 MA는 아홉 개 과제-길이 조합 모두에서 Standard를 앞섰다. 학습 창 안 평균은 1K에서 97.4%, 2K에서 96.8%로 Standard의 82.6%, 82.1%보다 높았다. niah_single_3 1K에서는 62.4%에서 98.2%로 가장 큰 개선이 있었다. 4K에서는 MA 평균 41.9% 대 Standard 25.9%로 우위를 유지했지만 두 모델 모두 학습 창을 넘으면 크게 저하됐다. 학습 효율은 동일 손실에 도달하는 데 필요한 토큰 수 비율로 측정됐고, Figure 2의 동작점에서 MA는 L24-D1024에서 1.42배, L24-D2048에서 1.16배의 토큰 효율을 보였다. 이는 각각 약 29.6%와 13.8% 더 적은 학습 토큰에 해당하지만, 벽시계 학습 속도를 측정한 것은 아니다.
추론 실험은 Standard, GPU에 메모리 테이블을 둔 MA, CPU로 오프로드한 MA-Offload를 비교했다. 두 MA 변형 모두 추론 전에 헤드별 RMSNorm을 테이블에 접어 넣어 온라인 값 구성을 조회와 덧셈으로 줄였고, MA-Offload는 프리페치 파이프라인으로 CPU 조회와 호스트-투-디바이스 전송을 모델 계산과 겹쳤다. 세 구성 모두 일반 KV 캐시를 유지하며 MA-Recall은 포함하지 않았다. 측정은 단일 H800, BF16, PyTorch 2.9.1, CUDA 12.6, FlashAttention 2.8.3에서 24층, 은닉 크기 2,048, 32개 쿼리 헤드, 32개 키/값 헤드, 어휘 32,000, 배치 8, 프리필 2,048, 디코딩 히스토리 2,048 토큰 조건으로 이뤄졌다. Table 4에서 MA는 프리필 지연을 90.970ms에서 88.202ms로 3.04% 줄였고, 디코드 지연은 17.636ms에서 17.788ms로 0.86% 늘었다. 총 파라미터 수가 Standard의 약 2.08배인데도 이 동작점에서는 전방 패스 지연이 비슷했다. MA-Offload는 1,572.864M개의 메모리 테이블 파라미터를 CPU로 옮겨 BF16 기준 3,000MiB를 확보했고, GPU 파라미터 저장량을 MA의 5,410.38MiB에서 2,410.38MiB로 55.45% 줄였다. Standard 대비로는 192MiB, 7.38% 감소다. 같은 조건에서 MA-Offload의 프리필과 디코드 지연은 각각 90.686ms와 17.189ms로 Standard의 90.970ms와 17.636ms보다 0.31%, 2.53% 낮게 측정됐다. 저자들은 이 수치가 값 투영 계산 감소, 메모리 접근, 전송, 스케줄링이 합쳐진 결과이며 통신이 공짜라거나 겹침의 기여를 분리한 것은 아니라고 밝힌다.
개발자 관점에서 MA는 값 투영을 없애고 토큰 메모리를 CPU로 오프로드해 GPU 파라미터 상주량을 줄이는 설계로, GPU 메모리가 병목인 LLM 서빙에서 검토할 만하다. 특히 조회 주소가 토큰 ID와 레이어 인덱스에만 의존해 프리페치가 가능하다는 점은 CPU-GPU 계층 분리에 유리하다. 그러나 저자들이 반복해서 강조하듯 품질 향상은 추가 메모리 파라미터 증가와 구조 변경이 분리되지 않았고, MA-Recall은 추론 측정에서 빠졌으며, 오프로딩은 총 GPU 메모리나 지연 시간 감소를 보장하지 않는다. 실험도 프로토타입이고, 타이밍은 테이블 초기화, 정규화 접기, 디코드 프리픽스 구성, 샘플링, GPU-CPU 토큰 ID 전송을 제외한 모델 전방 패스만 측정했다. 파라미터 저장량 추정은 KV 캐시, 활성값, 전송 버퍼를 제외하므로 총 또는 최대 GPU 메모리를 측정한 것이 아니다. 작은 타이밍 차이는 보고된 동작점의 측정값일 뿐 모든 워크로드에서 일관된 속도 우위를 뜻하지 않는다. 따라서 실무 도입 전에는 자신의 시퀀스 길이, 배치, 어텐션 패턴에서 토큰 효율, 지연, CPU 메모리와 전송 대역폭, KV 캐시 절감 가능성을 함께 검증해야 한다.
관련 논문
- 레이어별 활성값 기하를 보존하는 무학습 트랜스포머 압축, GeoPairGeoPair는 재학습 없이 트랜스포머 가중치를 압축하는 기법으로, 레이어마다 다른 활성값 화이트닝 공간을 그대로 두고 두 레이어가 사전을 공유하도록 최적 짝짓기와 닫힌 형태 해를 쓴다. Llama-3부터 영상 생성 모델까지 기준 정확도의 90% 이상을 유지한다고 주장한다.
- CompKV는 보상 잔차를 기준으로 KV 블록을 고르는 스파스 어텐션이다.장문맥 LLM 추론의 병목인 KV 캐시 메모리 트래픽을 줄이기 위해, 누락 토큰의 보상 기여를 미리 고려해 토큰을 고르는 스파스 어텐션 기법 CompKV를 제안한다. 기존 방식이 어텐션 질량만으로 토큰을 뽑고 나서 보상을 적용하던 것과 달리, 선택 단계에서 보상을 함께 설계한다.
- On-Demand Attention은 모델이 필요할 때만 전체 문맥을 다시 읽는다사전학습 모델의 디코딩 상태가 전역 어텐션의 필요성을 미리 알려준다는 관찰에서 출발한 On-Demand Attention을 소개한다. 로컬 우선 디코딩과 경량 recall 헤드로 필요한 순간에만 전체 문맥을 읽는 긴 문맥 추론 기법이다.
- dLLM 추론의 메모리 I/O 병목을 겨냥한 융합 KV 캐시와 자기 검증 병렬 디코딩확산 LLM(dLLM)의 추론 비효율을 KV 캐싱과 병렬 디코딩 결합으로 풀려는 연구다. 기존 가속 기법이 따로 다루던 두 축을 I/O 병목 관점에서 함께 설계하는 접근을 제안한다.
- JustFit은 24GiB 맥북에서 27B 256K 윈도우를 적시 상태 관리한다.오픈 웨이트 모델의 로컬 실행을 노트북 메모리 한계 안에 넣기 위해, JustFit은 압축 KV 실행과 컴포넌트 상주, 상태 보존 전환을 결합한 MLX 기반 추론 런타임을 제안한다. 가중치 양자화와 무관하게 상태를 적시에 만들고 해제하는 것이 핵심이다.