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(N 2 N^2 N 2 )으로 긴 문맥에서 제한됨
Modern Recurrent Model
attentional bias를 최적화
고정 크기 메모리를 사용해서 KV cache를 관리
더 정교한 메모리 관리 필요
학습 규칙: DeltaNet → \rightarrow → Delta rule
단순 가산: M t = M t − 1 + v t k t ⊤ M_t = M_{t-1} + v_t k_t^\top M t = M t − 1 + v t k t ⊤
새로운 정보를 더함. 메모리가 누적되고 오래된 정보 제거 불가
DeltaNet: M t = M t − 1 + β t ( v t − M t − 1 k t ) k t ⊤ M_t = M_{t-1} + \beta_t(v_t-M_{t-1}k_t)k^\top_t M t = M t − 1 + β t ( v t − M t − 1 k t ) k t ⊤
forget gate: RetNet의 비의존적 gating → \rightarrow → 적응적 gating(Titans)
M t = α t M t − 1 + u p d a t e M_t = \alpha_tM_{t-1} + update M t = α t M t − 1 + u p d a t e
RetNet: α \alpha α 가 입력과 무관하게 동일한 방식으로 감소
Titans: α t \alpha_t α t 가 입력에 따라 변화
메모리 구조: 벡터 기반 메모리 → \rightarrow → deep-neural memory
벡터 행렬 기반: M t ∈ R d M_t \in \mathbb R^d M t ∈ R d or M t ∈ R d × d M_t \in \mathbb R^{d \times d} M t ∈ R d × d
Deep Memory Module: M ( k ) = M L P ( k ) M(k) = MLP(k) M ( k ) = M L P ( 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을 포함하는 더 일반적인 모델군
향상된 메모리 관리 구조
다양한 작업에서 성능 향상
2. Preliminaries
Notation
input: x ∈ R N × d i n x \in \mathbb R^{N \times d_{in}} x ∈ R N × d i n
시간 t에서의 메모리 상태: M t M_t M t
key, value, query: K , V , Q K, V, Q K , V , Q , 각 벡터 k t , v t , q t k_t, v_t, q_t k t , v t , q t
attention bias: l ( M t ; k t , v t ) l(M_t;k_t,v_t) l ( M t ; k t , v t ) : 메모리가 무엇을 잘 기억해야 하는가
Memory module: L M ≥ 1 L_M \geq 1 L M ≥ 1 인 MLP + residual connection
Memory module parameter: θ M : = { W 1 , . . . , W L M , . . . } \theta_M := \{W_1, ..., W_{L_M}, ...\} θ M : = { W 1 , . . . , W L M , . . . }
2.1 Background
Attention
입력 x에 대해
Q = x W Q Q = xW_Q Q = x W Q , K = x W K K=xW_K K = x W K , V = x W V V=xW_V V = x W V
causal attention: y i = ∑ j = 1 i e x p ( q i ⊤ k j d i n ) v j ∑ l = 1 i e x p ( q i ⊤ k l d i n ) 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}}}})}} y i = ∑ l = 1 i e x p ( d i n q i ⊤ k l ) ∑ j = 1 i e x p ( d i n q i ⊤ k j ) v j
W Q W_Q W Q , W K W_K W K , W V ∈ R d i n × d i n W_V \in \mathbb R^{d_{in} \times d_{in}} W V ∈ R d i n × d i n 은 학습 가능한 파라미터
분모는 normalization term
각 토큰마다 최소 N × d N \times d N × d 만큼의 연산이 필요하기에 긴 문맥에서 확장성 제한
Recurrent Models
장점
선형 시간 학습
병렬화 가능
Transformer와 유사한 성능
M t = A t ∗ M t − 1 + v t k t ⊤ M_t = A_t * M_{t-1} + v_t k_t^\top M t = A t ∗ M t − 1 + v t k t ⊤
M t ∈ R d × n M_t \in \mathbb R^{d \times n} M t ∈ R d × n : 메모리
k t , v t ∈ R d k_t,v_t \in \mathbb R^d k t , v t ∈ R d : 입력
A t A_t A t : 감쇠, 게이팅
∗ * ∗ : 임의의 연산자
단점
Deep Memory Module
메모리는 단순 저장이 아닌 학습되는 함수
M ∗ = a r g m i n M L ( M ( K ) ; V ) M^* = argmin_M L(M(K); V) M ∗ = a r g m i n M L ( M ( K ) ; V )
반복적인 알고리즘으로 수행: 메모리 업데이트 규칙
최적화 구조
inner Loop: 메모리 파라미터 θ M \theta_M θ 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 ( M t ; k t , v t ) = ∣ ∣ M t ( k t ) − v t ∣ ∣ 2 2 l(M_t; k_t, v_t) = ||M_t(k_t) - v_t||^2_2 l ( M t ; k t , v t ) = ∣ ∣ M t ( k t ) − v t ∣ ∣ 2 2
graident descent로 최적화
Capacity of l 2 l_2 l 2
Memory M: d v × d k d_v \times d_k d v × d k 크기 행렬
linear independent key를 가지는 (k, v)쌍을 O ( d k ) O(d_k) O ( d k ) 만큼 저장
메모리 파라미터 수에 비해 저장 가능한 정보량이 sub-linear
메모리 크기가 M이면 저장 가능한 패턴 수는 c × M c\times M c × M 보다 작음
Theorem 1: Effect of Deep Memory
L M ≥ 2 L_M \geq 2 L M ≥ 2 층의 MLP형태 메모리 M에 대해
input dimension: d k d_k d k , hidden dimension: d h d_h d h
메모리 저장 범위: O ( d k , d v ) 에서 O ( d k d v ∑ i = 1 L M m i n { d h ( j ) } j ≥ i d h ( 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) O ( d k , d v ) 에 서 O ( d k d v i = 1 ∑ L M m i n { d h ( j ) } j ≥ i d h ( j + 1 ) )
MLP 형태 메모리가 몇 개의 서로 다른 key를 linearly independent하게 value로 매핑할 수 있는가
MLP는 선형이 아닌 여러 층을 지나며 비선형적으로 변형됨
하한: 입력 차원 × \times × 출력 차원
상한
bottleneck: 가장 작은 차원의 layer가 전체 capacity를 제한함(m i n j ≥ i d h ( j ) \displaystyle min_{j \geq i}d^{(j)}_h m i n j ≥ i d h ( j )
각 레이어는 변환 능력을 추가함: d h ( j ) × d h ( j + 1 ) d^{(j)}_h \times d^{(j+1)}_h d h ( j ) × d h ( j + 1 )
표현력과 용량이 증가하지만, 여전히 super-linear하지 않음
단순히 d k d_k d k 의 차원을 늘리면
파라미터, 메모리 사용량, 계산 비용이 증가함
Kernel Attention Perspective를 이용
Polynomial Feature Mapping
입력 x에 대해: Φ p ( x ) = [ x β ] ∣ β ∣ ≤ p \Phi_p(x) = [x^{\beta}]_{|\beta| \leq p} Φ p ( x ) = [ x β ] ∣ β ∣ ≤ p
최대 차수가 p까지의 polynomial feature로 확장
메모리 학습: L ( M ( Φ ( K ) ) ; V ) L(M(\Phi(K));V) L ( M ( Φ ( K ) ) ; V )
Memory Capacity with Polynomial Mapping
다항 매핑 Φ p \Phi_p Φ p 를 사용하면 행렬 메모리는 최대 O ( d k p ) O(d^p_k) O ( d k p ) 개의 (k, v)쌍 저장
attentional bias를 gradient descent로 최적화할 때 메모리 업데이트
내적 기반: l ( 1 ) ( M t ; k t , v t ) = < M t k t , v t > l^{(1)}(M_t; k_t, v_t) = <M_tk_t, v_t> l ( 1 ) ( M t ; k t , v t ) = < M t k t , v t >
l 2 l2 l 2 오차 기반: l ( 2 ) ( M t ; k t , v t ) = ∣ ∣ M t Φ ( k t ) − v t ∣ ∣ 2 2 l^{(2)}(M_t; k_t, v_t) = ||M_t \Phi (k_t) - v_t||^2_2 l ( 2 ) ( M t ; k t , v t ) = ∣ ∣ M t Φ ( k t ) − v t ∣ ∣ 2 2
이를 gradient descent로 최적화
Hebbian Rule: M t = M t − 1 + η t v t Φ ( k t ) ⊤ M_t = M_{t-1} + \eta_tv_t\Phi(k_t)^\top M t = M t − 1 + η t v t Φ ( k t ) ⊤
새로운 정보 ( k t , v t ) (k_t,v_t) ( k t , v t ) 를 그대로 메모리에 추가
v t Φ ( k t ) T v_t\Phi(k_t)^T v t Φ ( k t ) T : key-value관계를 outer product로 저장
같이 나타난 것들은 연결해서 기억
Delta Rule: M t = ( I − η t Φ ( k t ) Φ ( k t ) ⊤ ) M t − 1 + η t v t Φ ( k t ) ⊤ M_t = (I-\eta_t\Phi(k_t)\Phi(k_t)^\top)M_{t-1} + \eta_t v_t \Phi(k_t)^\top M t = ( I − η t Φ ( k t ) Φ ( k t ) ⊤ ) M t − 1 + η t v t Φ ( k t ) ⊤
1항: 기존 메모리 수정(망각도 포함)
2항: 새로운 정보 추가
Hebbian Rule의 Kernel Attention Perspective
Attention의 exponential kernel: e x p ( q i ⊤ k j ) exp(q^\top_ik_j) e x p ( q i ⊤ k j )
테일러 전개: e x p ( q i ⊤ k j ) ≈ 1 + q i ⊤ k j + ( q i T k j ) 2 2 ! + ( q i ⊤ k j ) 3 3 ! + . . . 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!}}+... e x p ( q i ⊤ k j ) ≈ 1 + q i ⊤ k j + 2 ! ( q i T k j ) 2 + 3 ! ( q i ⊤ k j ) 3 + . . .
일반화: e x p ( q i ⊤ k j ) ≈ Φ p ( q i ) ⊤ Φ p ( k j ) exp(q^\top_ik_j) \approx \Phi_p(q_i)^\top\Phi_p(k_j) e x p ( q i ⊤ k j ) ≈ Φ p ( q i ) ⊤ Φ p ( k j )
다항 매핑은 softmax attention의 근사
해석
기존 attention의 exponential은 비선형이고 분리가 불가능함
e x p ( q i ⊤ k j ) ≠ f ( q i ) ⋅ g ( k j ) exp(q^\top_ik_j) \neq f(q_i) \cdot g(k_j) e x p ( q i ⊤ k j ) = f ( q i ) ⋅ g ( k j )
이로 인해 attention의 time complexity가 O ( N 2 ) O(N^2) O ( N 2 ) 이 됨
이를 Φ \Phi Φ 로 이루어진 두 식으로 변형하면서 고차원 공간에서의 단순 내적으로 변형시킴
즉 attention은 고차원 feature 공간에서의 내적임
Input gating
Φ p ( x ) = a 0 + a 1 x + a 2 x 2 + . . . + a p x p \Phi_p(x) = a_0 + a_1x + a_2x^2 +... + a_px^p Φ p ( x ) = a 0 + a 1 x + a 2 x 2 + . . . + a p x p
a i a_i a i 는 특정 차수의 특징을 켜거나 끔
3.2 Long-term Memory with Context Memorization
기존의 Recurrent Model은 현재 입력에 대해서만 attentional bias를 최적화하고 이전 메모리 상태는 유지함
m i n M l ( M ; k t , v t ) + R e t t ( M , M t − 1 ) min_M l(M;k_t,v_t)+Ret_t(M,M_{t-1}) m i n M l ( M ; k t , v t ) + R e t t ( M , M t − 1 )
이때 Ret은 보존 게이트
계산이 간단하고 빠르지만 문맥 전체를 제대로 기억하지 못함
매 시점마다 전체 입력 시퀀스에 대해 최적화
m i n M ∑ i = 1 t l ( M ; k i , v i ) min_M \displaystyle\sum^t_{i=1}l(M;k_i,v_i) m i n M i = 1 ∑ t l ( M ; k i , v i )
Sliding window recurrent model
m i n M ∑ i = t − c + 1 t γ i ( t ) l ( M ; k i , v i ) min_M \displaystyle \sum^t_{i=t-c+1}\gamma ^{(t)}_i l(M;k_i, v_i) m i n M i = t − c + 1 ∑ t γ i ( t ) l ( M ; k i , v i )
c: 문맥 길이
γ i ( t ) \gamma^{(t)}_i γ i ( t ) : 감쇠 게수
OmegaNet
Omega Rule
m i n M ∑ i = t − c + 1 t γ i ( t ) ∣ ∣ M ( k i ) − v i ∣ ∣ 2 2 min_M \displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i ||M(k_i) -v_i||^2_2 m i n M i = t − c + 1 ∑ t γ i ( t ) ∣ ∣ M ( k i ) − v i ∣ ∣ 2 2
문맥 전체 기준 업데이트
c = 1일때 기존 Delta Rule
OmegaNet
Update: M t = α t M t − 1 − ∇ ∑ i = t − c + 1 t γ i ( t ) ∣ ∣ M ( Φ ( k i ) ) − v i ∣ ∣ 2 2 M_t = \alpha_tM_{t-1} - ∇\displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i ||M(\Phi(k_i)) - v_i||^2_2 M t = α t M t − 1 − ∇ i = t − c + 1 ∑ t γ i ( t ) ∣ ∣ M ( Φ ( k i ) ) − v i ∣ ∣ 2 2
선형M t = ( d i a g ( a t ) − ∑ i = t − c + 1 t γ i ( t ) Φ ( k i ) Φ ( k i ) ⊤ ) M t − 1 − ∑ i = t − c + 1 t γ i ( t ) v i Φ ( k i ) ⊤ 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 M t = ( d i a g ( a t ) − i = t − c + 1 ∑ t γ i ( t ) Φ ( k i ) Φ ( k i ) ⊤ ) M t − 1 − i = t − c + 1 ∑ t γ i ( t ) v i Φ ( k i ) ⊤
d i a g ( a t ) diag(a_t) d i a g ( a t ) : 기존 메모리 유지 정도
∑ i = t − c + 1 t γ i ( t ) Φ ( k i ) Φ ( k i ) ⊤ \displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i \Phi (k_i) \Phi(k_i)^\top i = t − c + 1 ∑ t γ i ( t ) Φ ( k i ) Φ ( k i ) ⊤ : 특정 방향의 정보를 제거. 즉 현재 문맥과 충돌하는 기억을 지움
− ∑ i = t − c + 1 t γ i ( t ) v i Φ ( k i ) ⊤ -\displaystyle \sum^t_{i=t-c+1} \gamma^{(t)}_i v_i \Phi(k_i)^\top − i = t − c + 1 ∑ t γ i ( t ) v i Φ ( k i ) ⊤ : 새로운 key-value관계를 slide로 추가
Gradient Descent외의 예시
최적화 알고리즘
c = 1, momentum: S t = θ t S t − 1 − η t ∇ l ( M t − 1 ; k t , v t ) S_t = \theta_t S_{t-1} - \eta_t ∇l(M_{t-1};k_t,v_t) S t = θ t S t − 1 − η t ∇ l ( M t − 1 ; k t , v t ) : Titans
c = context length: M t = a r g m i n M ∑ i = 1 t ∣ ∣ M K i − v i ∣ ∣ 2 2 M_t=arg min_M \displaystyle \sum^t_{i=1} ||MK_i - v_i||^2_2 M t = a r g m i n M i = 1 ∑ t ∣ ∣ M K i − v i ∣ ∣ 2 2
이런 케이스들은 병렬화가 어렵고 계산 비용이 큼
3.3 Parallelizing Omega Rule
원래 Omega Rule
c개의 gradient ∇ l ∈ R d i n × d i n ∇l \in \mathbb R^{d_{in} \times d_{in}} ∇ l ∈ R d i n × d i n 을 모두 계산해야 함
청크 분할
입력 시퀀스 길이 L을 b ≥ 1 b \ge 1 b ≥ 1 인 청크로 나눔
S i = { x ( i − 1 ) b + 1 , . . . , x i b } S_i = \{x_{(i-1)b+1}, ..., x_{ib}\} S i = { x ( i − 1 ) b + 1 , . . . , x i b }
이전 청크의 마지막 메모리 상태를 기준으로 gradient 계산
γ i ( t ) = η t \gamma^{(t)}_i = \eta_t γ i ( t ) = η t 라고 가정
b = 1일때(청크 x): M t = α t M t − 1 − η t ∑ i = t − c + 1 t ∇ l ( M t − 1 ; k i , v i ) M_t = \alpha_tM_{t-1} - \eta_t \displaystyle \sum^t_{i=t-c+1} ∇l(M_{t-1}; k_i,v_i) M t = α t M t − 1 − η t i = t − c + 1 ∑ t ∇ l ( M t − 1 ; k i , v i )
b ≥ 1 b \ge 1 b ≥ 1
t ′ = t − m o d ( t , b ) t' = t - mod(t,b) t ′ = t − m o d ( t , b ) , t'은 현재 청크 시작 지점
M t = α t . . . α t ′ M t ′ − ∑ n = t ′ t ( α t . . . α n + 1 ) η n i = n − c + 1 n ∇ l ( M t ′ ; k i , v i ) 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) M t = α t . . . α t ′ M t ′ − n = t ′ ∑ t ( α t . . . α n + 1 ) η n i = n − c + 1 n ∇ l ( M t ′ ; k i , v i )
graident계산시 sliding window mask인 M s M_s M s 사용
효과
메모리 사용량 감소
병렬화 가능
계산 효율 유지
길어져서 렉걸리는관계로 다음 게시글로