[논문 리뷰] Google ATLAS

마계닭·2026년 4월 23일

논문 리뷰

목록 보기
17/18

ATLAS: Learning to Optimally Memorize the Context at Test Time

Behrouz, A., Li, Z., Kacham, P., Daliri, M., Deng, Y., Zhong, P., Razaviyayn, M., & Mirrokni, V. (2025). Atlas: Learning to optimally memorize the context at test time. arXiv. https://arxiv.org/abs/2505.23735

Titans 논문을 읽고 오는 것이 좋다.

Introduction(기존 모델의 한계)

Transformer(Attention module)

  • associative memory처럼 작동
  • 토큰 간 쌍별 의존성 -> K-V mapping, QK 유사도로 검색
    • 시간 복잡도가 O(N2N^2)으로 긴 문맥에서 제한됨

Modern Recurrent Model

  • attentional bias를 최적화
    • 고정 크기 메모리를 사용해서 KV cache를 관리
  • 더 정교한 메모리 관리 필요
    • 학습 규칙: DeltaNet \rightarrow Delta rule
      • 단순 가산: Mt=Mt1+vtktM_t = M_{t-1} + v_t k_t^\top
        • 새로운 정보를 더함. 메모리가 누적되고 오래된 정보 제거 불가
      • DeltaNet: Mt=Mt1+βt(vtMt1kt)ktM_t = M_{t-1} + \beta_t(v_t-M_{t-1}k_t)k^\top_t
        • 예측과 실제 값 차이를 줄이도록 업데이트
    • forget gate: RetNet의 비의존적 gating \rightarrow 적응적 gating(Titans)
      • Mt=αtMt1+updateM_t = \alpha_tM_{t-1} + update
      • RetNet: α\alpha가 입력과 무관하게 동일한 방식으로 감소
      • Titans: αt\alpha_t가 입력에 따라 변화
    • 메모리 구조: 벡터 기반 메모리 \rightarrow deep-neural memory
      • 벡터 행렬 기반: MtRdM_t \in \mathbb R^d or MtRd×dM_t \in \mathbb R^{d \times d}
        • 표현력이 부족하고 저장가능한 패턴 수 제한
      • Deep Memory Module: M(k)=MLP(k)M(k) = MLP(k)
        • 추상화된 정보를 저장할 수 있지만 문맥 단위 학습 부족
  • 한계
    • online update: 메모리가 현재 토큰에 대해서만 최적화
    • 제한된 메모리 용량: KV 매핑에 의해 저장 가능한 정보량 제한
    • 메모리 관리 표현력 부족: 1차 정보에 기반한 graident descent로 비효율적 mapping

Memory Perspective

Associative Memory: 서로 다른 개체, 사건을 연결하는 능력

  • 메모리는 입력에 의해 갱신되는 신경적 상태
  • surprise가 클수록 더 큰 영향을 미침
  • 개별 토큰 중심의 surprise라는 한계점이 존재함
    • 토큰이 아닌 문맥의 surprise를 측정하는 방식

Test time memorization

  • 모델의 학습된 파라미터를 업데이트 하지 않음
  • 주어진 입력 문맥 내에서만 정보를 저장, 검색
  • 즉, 메모리가 초기화되면 새로운 문맥에는 이전 학습이 유지되지 않음

Contribution

  • 메모리 용량에 대한 이해, 개선
    • 입력에 대해 higher-order mapping을 사용해 메모리 용량 증가
  • 새로운 학습 규칙: Omega rule
    • 문맥 전체를 고려해 메모리 업데이트
  • Transforemr의 일반화
    • DeepTransformers를 제시해 기존 Transformer을 포함하는 더 일반적인 모델군
  • 향상된 메모리 관리 구조
    • OmegaNet과 Atlas
  • 다양한 작업에서 성능 향상

2. Preliminaries

Notation

  • input: xRN×dinx \in \mathbb R^{N \times d_{in}}
  • 시간 t에서의 메모리 상태: MtM_t
  • key, value, query: K,V,QK, V, Q, 각 벡터 kt,vt,qtk_t, v_t, q_t
  • attention bias: l(Mt;kt,vt)l(M_t;k_t,v_t): 메모리가 무엇을 잘 기억해야 하는가
  • Memory module: LM1L_M \geq 1인 MLP + residual connection
  • Memory module parameter: θM:={W1,...,WLM,...}\theta_M := \{W_1, ..., W_{L_M}, ...\}

2.1 Background

Attention

입력 x에 대해

  • Q=xWQQ = xW_Q, K=xWKK=xW_K, V=xWVV=xW_V
  • causal attention: yi=j=1iexp(qikjdin)vjl=1iexp(qikldin)y_i = {{\sum^i_{j=1}exp({{q^\top_ik_j}\over{\sqrt{d_{in}}}})v_j}\over{\sum^i_{l=1} exp({{q^\top_ik_l}\over{\sqrt{d_{in}}}})}}
    • WQW_Q, WKW_K, WVRdin×dinW_V \in \mathbb R^{d_{in} \times d_{in}}은 학습 가능한 파라미터
    • 분모는 normalization term
  • 각 토큰마다 최소 N×dN \times d만큼의 연산이 필요하기에 긴 문맥에서 확장성 제한

Recurrent Models

  • 장점
    • 선형 시간 학습
    • 병렬화 가능
    • Transformer와 유사한 성능
  • Mt=AtMt1+vtktM_t = A_t * M_{t-1} + v_t k_t^\top
    • MtRd×nM_t \in \mathbb R^{d \times n}: 메모리
    • kt,vtRdk_t,v_t \in \mathbb R^d: 입력
    • AtA_t: 감쇠, 게이팅
    • *: 임의의 연산자
  • 단점
    • 메모리 계속 누적

Deep Memory Module

  • 메모리는 단순 저장이 아닌 학습되는 함수
  • M=argminML(M(K);V)M^* = argmin_M L(M(K); V)
  • 반복적인 알고리즘으로 수행: 메모리 업데이트 규칙
  • 최적화 구조
    • inner Loop: 메모리 파라미터 θM\theta_M 최적화
    • Outer Loop: 나머지 파라미터 학습

3. Learning to Memorize the Context at Test Time

3.1 Associataive Memory with Super Linear Capacity

모델은 서로 상관이 없는 KV 쌍을 최대 몇 개 까지 저장할 수 있는가

  • matrix memory + l2 loss를 사용
    • l(Mt;kt,vt)=Mt(kt)vt22l(M_t; k_t, v_t) = ||M_t(k_t) - v_t||^2_2
    • graident descent로 최적화

Capacity of l2l_2

Memory M: dv×dkd_v \times d_k크기 행렬

  • linear independent key를 가지는 (k, v)쌍을 O(dk)O(d_k)만큼 저장
  • 메모리 파라미터 수에 비해 저장 가능한 정보량이 sub-linear
    • 메모리 크기가 M이면 저장 가능한 패턴 수는 c×Mc\times M보다 작음

Theorem 1: Effect of Deep Memory

LM2L_M \geq 2층의 MLP형태 메모리 M에 대해

  • input dimension: dkd_k, hidden dimension: dhd_h
  • 메모리 저장 범위: O(dk,dv)에서O(dkdvi=1LMmin{dh(j)}jidh(j+1))O(d_k, d_v)에서 O(d_kd_v \displaystyle \sum^{L_M}_{i=1} min\{d^{(j)}_h\}_{j \geq i}d^{(j+1)}_h)
    • MLP 형태 메모리가 몇 개의 서로 다른 key를 linearly independent하게 value로 매핑할 수 있는가
    • MLP는 선형이 아닌 여러 층을 지나며 비선형적으로 변형됨
    • 하한: 입력 차원 ×\times 출력 차원
    • 상한
      • bottleneck: 가장 작은 차원의 layer가 전체 capacity를 제한함(minjidh(j)\displaystyle min_{j \geq i}d^{(j)}_h
      • 각 레이어는 변환 능력을 추가함: dh(j)×dh(j+1)d^{(j)}_h \times d^{(j+1)}_h
  • 표현력과 용량이 증가하지만, 여전히 super-linear하지 않음

단순히 dkd_k의 차원을 늘리면

  • 파라미터, 메모리 사용량, 계산 비용이 증가함
  • Kernel Attention Perspective를 이용

Polynomial Feature Mapping
입력 x에 대해: Φp(x)=[xβ]βp\Phi_p(x) = [x^{\beta}]_{|\beta| \leq p}

  • 최대 차수가 p까지의 polynomial feature로 확장
    메모리 학습: L(M(Φ(K));V)L(M(\Phi(K));V)

Memory Capacity with Polynomial Mapping

다항 매핑 Φp\Phi_p를 사용하면 행렬 메모리는 최대 O(dkp)O(d^p_k)개의 (k, v)쌍 저장

  • 차수 p에 따라 용량이 급격히 증가

attentional bias를 gradient descent로 최적화할 때 메모리 업데이트

  • 내적 기반: l(1)(Mt;kt,vt)=<Mtkt,vt>l^{(1)}(M_t; k_t, v_t) = <M_tk_t, v_t>
  • l2l2오차 기반: l(2)(Mt;kt,vt)=MtΦ(kt)vt22l^{(2)}(M_t; k_t, v_t) = ||M_t \Phi (k_t) - v_t||^2_2
    이를 gradient descent로 최적화

Hebbian Rule: Mt=Mt1+ηtvtΦ(kt)M_t = M_{t-1} + \eta_tv_t\Phi(k_t)^\top

  • 새로운 정보 (kt,vt)(k_t,v_t)를 그대로 메모리에 추가
  • vtΦ(kt)Tv_t\Phi(k_t)^T: key-value관계를 outer product로 저장
  • 같이 나타난 것들은 연결해서 기억

Delta Rule: Mt=(IηtΦ(kt)Φ(kt))Mt1+ηtvtΦ(kt)M_t = (I-\eta_t\Phi(k_t)\Phi(k_t)^\top)M_{t-1} + \eta_t v_t \Phi(k_t)^\top

  • 1항: 기존 메모리 수정(망각도 포함)
  • 2항: 새로운 정보 추가

Hebbian Rule의 Kernel Attention Perspective
Attention의 exponential kernel: exp(qikj)exp(q^\top_ik_j)

  • 테일러 전개: exp(qikj)1+qikj+(qiTkj)22!+(qikj)33!+...exp(q_i^\top k_j) \approx 1+q_i^\top k_j + {{(q^T_ik_j)^2}\over{2!}} + {{(q^\top_i k_j)^3}\over{3!}}+...
  • 일반화: exp(qikj)Φp(qi)Φp(kj)exp(q^\top_ik_j) \approx \Phi_p(q_i)^\top\Phi_p(k_j)
  • 다항 매핑은 softmax attention의 근사

해석
기존 attention의 exponential은 비선형이고 분리가 불가능함

  • exp(qikj)f(qi)g(kj)exp(q^\top_ik_j) \neq f(q_i) \cdot g(k_j)
  • 이로 인해 attention의 time complexity가 O(N2)O(N^2)이 됨
  • 이를 Φ\Phi로 이루어진 두 식으로 변형하면서 고차원 공간에서의 단순 내적으로 변형시킴
  • 즉 attention은 고차원 feature 공간에서의 내적임

Input gating
Φp(x)=a0+a1x+a2x2+...+apxp\Phi_p(x) = a_0 + a_1x + a_2x^2 +... + a_px^p

  • aia_i는 특정 차수의 특징을 켜거나 끔

3.2 Long-term Memory with Context Memorization

기존의 Recurrent Model은 현재 입력에 대해서만 attentional bias를 최적화하고 이전 메모리 상태는 유지함

  • minMl(M;kt,vt)+Rett(M,Mt1)min_M l(M;k_t,v_t)+Ret_t(M,M_{t-1})
    • 이때 Ret은 보존 게이트
    • 계산이 간단하고 빠르지만 문맥 전체를 제대로 기억하지 못함
  • 매 시점마다 전체 입력 시퀀스에 대해 최적화
    • minMi=1tl(M;ki,vi)min_M \displaystyle\sum^t_{i=1}l(M;k_i,v_i)
      • 모든 과거 토큰을 보기에 비용 증가

Sliding window recurrent model
minMi=tc+1tγi(t)l(M;ki,vi)min_M \displaystyle \sum^t_{i=t-c+1}\gamma ^{(t)}_i l(M;k_i, v_i)

  • c: 문맥 길이
  • γi(t)\gamma^{(t)}_i: 감쇠 게수

OmegaNet

Omega Rule

  • minMi=tc+1tγi(t)M(ki)vi22min_M \displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i ||M(k_i) -v_i||^2_2
    • 문맥 전체 기준 업데이트
      c = 1일때 기존 Delta Rule

OmegaNet

  • Update: Mt=αtMt1i=tc+1tγi(t)M(Φ(ki))vi22M_t = \alpha_tM_{t-1} - ∇\displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i ||M(\Phi(k_i)) - v_i||^2_2
  • 선형Mt=(diag(at)i=tc+1tγi(t)Φ(ki)Φ(ki))Mt1i=tc+1tγi(t)viΦ(ki)M_t = (diag(a_t) - \displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i \Phi (k_i) \Phi(k_i)^\top )M_{t-1} -\displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i v_i \Phi(k_i)^\top
    • diag(at)diag(a_t): 기존 메모리 유지 정도
    • i=tc+1tγi(t)Φ(ki)Φ(ki)\displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i \Phi (k_i) \Phi(k_i)^\top: 특정 방향의 정보를 제거. 즉 현재 문맥과 충돌하는 기억을 지움
    • i=tc+1tγi(t)viΦ(ki)-\displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i v_i \Phi(k_i)^\top: 새로운 key-value관계를 slide로 추가

Gradient Descent외의 예시

최적화 알고리즘

  • c = 1, momentum: St=θtSt1ηtl(Mt1;kt,vt)S_t = \theta_t S_{t-1} - \eta_t ∇l(M_{t-1};k_t,v_t): Titans
  • c = context length: Mt=argminMi=1tMKivi22M_t=arg min_M \displaystyle \sum^t_{i=1} ||MK_i - v_i||^2_2
  • 이런 케이스들은 병렬화가 어렵고 계산 비용이 큼

3.3 Parallelizing Omega Rule

원래 Omega Rule

  • c개의 gradient lRdin×din∇l \in \mathbb R^{d_{in} \times d_{in}}을 모두 계산해야 함
    • 메모리 사용량과 I/O 비용 모두 증가

청크 분할

  • 입력 시퀀스 길이 L을 b1b \ge 1인 청크로 나눔
    • Si={x(i1)b+1,...,xib}S_i = \{x_{(i-1)b+1}, ..., x_{ib}\}
      • 이전 청크의 마지막 메모리 상태를 기준으로 gradient 계산
  • γi(t)=ηt\gamma^{(t)}_i = \eta_t라고 가정
    • b = 1일때(청크 x): Mt=αtMt1ηti=tc+1tl(Mt1;ki,vi)M_t = \alpha_tM_{t-1} - \eta_t \displaystyle \sum^t_{i=t-c+1} ∇l(M_{t-1}; k_i,v_i)
    • b1b \ge 1
      • t=tmod(t,b)t' = t - mod(t,b), t'은 현재 청크 시작 지점
      • Mt=αt...αtMtn=tt(αt...αn+1)ηni=nc+1nl(Mt;ki,vi)M_t = \alpha_t ... \alpha_{t'}M_{t'} - \displaystyle \sum^t_{n=t'}(\alpha_t ... \alpha_{n+1})\eta_n \displaystyle^n_{i=n-c+1} ∇l(M_{t'}; k_i, v_i)
  • graident계산시 sliding window mask인 MsM_s 사용
  • 효과
    • 메모리 사용량 감소
    • 병렬화 가능
    • 계산 효율 유지

길어져서 렉걸리는관계로 다음 게시글로

profile
뉴비

0개의 댓글