[논문리뷰]Gated Delta Netoworks: Improving MAMBA2 with Delta Rule (ICLR, 2025)

Date:     Updated:

카테고리:

Songlin Yang, Jan Kautz, and Ali Hatamizadeh. 2025. Gated delta networks: Improving mamba2 with delta rule. In International Conference on Learning Representations, Vol. 2025. 29687–29707.

1. Problem Formulation

이 논문은 standard Transformer의 quadratic self-attention을 대체할 efficient sequence modeling / language modeling architecture를 개발하면서, 기존 linear recurrent model의 retrieval 및 long-context memory management 문제를 해결하는 것을 목표로 한다. Transformer의 self-attention은 정확한 sequence modeling과 GPU parallelism에는 강하지만 sequence length에 대해 quadratic cost를 가지며, linear Transformer는 recurrent matrix state를 사용해 효율성을 높이는 대신 long sequence에서 key-value association이 충돌해 retrieval 성능이 약해지는 문제가 있다. Gated DeltaNet은 이를 위해 adaptive memory decay와 selective key-value update를 동시에 수행하는 gated delta rule을 제안한다.

  • 입력: 각 time step t에서 생성되는 \(q_t, k_t, v_t\)와 data-dependent gating parameters \(\alpha_t, \beta_t\)
  • 출력: recurrent memory state \(S_t\)와 이를 query \(q_t\)로 읽은 token-mixing output \(o_t = S_tq_t\)
  • 최종 목표: Mamba2보다 정밀한 key-value association learning을 수행하고 DeltaNet보다 효과적으로 memory를 clear하면서, hardware-efficient parallel training이 가능한 recurrent architecture를 만드는 것이다.



2. Limitations of Existing Works

[Quadratic Attention] Standard Transformer는 query가 이전 token들의 key/value를 직접 참조할 수 있기 때문에 precise sequence modeling에 강하지만, self-attention computation이 sequence length에 대해 quadratic하게 증가한다. 긴 sequence에서 training과 inference cost가 커지므로, 논문은 linear Transformer와 linear RNN 계열을 효율적인 대안으로 다룬다.

[Memory Collision] Linear Transformer는 과거 정보를 matrix-valued state \(S_t\)에 outer-product 형태의 key-value association으로 압축한다. 그러나 state dimension이 고정되어 있으므로 저장 가능한 orthogonal key-value pair 수가 제한되고, sequence가 길어지면 여러 정보가 같은 state에 superposition되는 memory collision이 발생한다. 이 때문에 exact retrieval과 in-context retrieval에서 Transformer보다 약한 문제가 발생한다.



3. Methodology

1

전체 pipeline은

  1. input representation
  2. \(q, k, v, \alpha, \beta\) 생성
  3. gated delta rule을 통한 recurrent state \(S_t\) 업데이트
  4. \(S_t\)를 query \(q_t\)로 읽어 output 생성

으로 구성된다. Figure 1의 오른쪽 block에서 \(q, k, v\)는 linear projection, short convolution, SiLU를 거쳐 생성되며 \(q, k\)에는 L2 normalization이 적용된다. \(\alpha, \beta\)는 gating/update strength를 결정하며, recurrent computation 자체는 hardware-efficient chunkwise form으로 병렬화된다. 기본 Gated DeltaNet은 self-attention을 이 token mixer로 대체하고, 저자들은 추가로 SWA와 Mamba2를 섞은 hybrid model도 제안한다.

3.1. Linear Attention as Associative Memory

GDN을 이해하려면 먼저 linear attention의 state \(S_t\)key-value memory matrix라고 보는 것이 중요하다.

Vanilla linear attention은 다음 recurrence로 표현된다.

$$S_t = S_{t-1}+v_tk_t^\top, \quad o_t = S_tq_t$$

즉 \(v_tk_t^\top\)를 state에 계속 추가한다. \(k_t\)는 memory에 정보를 쓰는 주소와 비슷하고 \(v_t\)는 그 주소에 저장되는 내용이며, \(q_t\)는 현재 state에서 정보를 읽는 query에 해당한다. 전체 과거 토큰에 attention을 다시 계산하는 대신 \(S_t\)만 유지하면 되므로 recurrent inference가 가능하다. 그러나 계속 새로운 outer product를 덧붙이기만 하면 state에 많은 association이 겹치게 된다.

3.2. Mamba2: Global Forgetting

Mamba2는 vanilla linear attention에 data-dependent decay \(\alpha_t \in (0, 1)\)를 추가한다.

$$S_t = \alpha_tS_{t-1} + v_tk_t^\top, \quad o_t = S_tq_t$$

여기서 \(\alpha_t\)는 이전 memory 전체의 retention 정도를 결정한다. \(\alpha_t \approx 1\)이면 이전 state를 거의 유지하고, \(\alpha_t \approx 0\)이면 이전 state를 거의 제거한다. 따라서 context가 바뀌었을 때 state를 빠르게 비울 수 있다는 장점이 있다.

그러나 핵심 문제가 있다.

$$S_{t-1} \rightarrow \alpha_tS_{t-1}$$

이므로 state 내부의 모든 key-value association이 동일한 \(\alpha_t\)로 감소한다. 현재 필요 없는 memory 하나만 선택적으로 제거하는 것이 아니라 중요한 memory도 같이 약해진다. 이 점이 Mamba2의 장기 retrieval에서 문제가 된다.

3.3. DeltaNet: Selective Memory Update

DeltaNet은 global decay 대신 delta rule을 사용한다. 현재 key \(k_t\)를 state에서 읽으면

$$v_t^\text{old} = S_{t-1}k_t$$

가 된다. 즉 기존 state가 현재 key에 대해 이미 기억하고 있는 value이다. Delta rule은 이 association을 새로운 \(v_t\) 방향으로 수정하며 최종 recurrence는 다음과 같다.

$$S_t = S_{t-1} (I - \beta_t k_t k_t^\top) + \beta_t v_t k_t^\top$$

여기서 \(\beta_t \in (0, 1)\)는 writing strength이다. 이를 전개하면 직관적으로 다음과 같다.

$$S_t = S_{t-1} + \beta_t(v_t-S_{t-1}k_t)k_t^\top$$

즉 새로운 \(v_t\)를 무조건 state에 더하는 것이 아니라, \(v_t - S_{t-1}k_t\)이라는 현재 memory prediction과 실제 incoming value 사이의 error, 즉 delta만 기록하는 구조이다. 따라서 이미 \(S_{t-1}k_t \approx v_t\)라면 수정량이 작고, 현재 key에 잘못된 value가 연결되어 있다면 해당 key direction을 강하게 수정한다. 논문은 이를 generalized Householder transition \(I -\beta_t k_t k_t^\top\)으로 설명한다. 핵심은 Mamba2는 state 전체를 지우고, DeltaNet은 특정 kt direction을 수정한다는 차이이다.

3.4. Gated Delta Rule

Gated DeltaNet의 핵심 수식은 하나이다.

$$S_t = S_{t-1}\big( \alpha_t (I - \beta_tk_tk_t^\top) \big) + \beta_tv_tk_t^\top$$

여기서

  • \(S_t\): recurrent memory state
  • \(k_t\): 현재 memory address 역할의 key
  • \(v_t\): 현재 기록할 value
  • \(\beta_t\): selective writing strength
  • \(\alpha_t\): data-dependent global state decay, global memory control

이다. 여기서 가장 중요한 부분은 transition matrix가 \(\alpha_t (I - \beta_tk_tk_t^\top)\)라는 점이다. 즉 두 종류의 memory control이 동시에 존재한다.

\(\alpha_t\) - global memory control

  • \(\alpha_t \rightarrow 0\)이면 이전 \(S_{t-1}\)의 영향이 빠르게 사라진다. 따라서 context switch가 발생했거나 기존 memory가 대부분 불필요한 경우 빠르게 state를 reset할 수 있다.
  • 반대로 \(\alpha_t = 0\)이면 global forgetting이 거의 없어지고 위의 수식은 사실상 DeltaNet의 update로 돌아간다. 저자들이 Gating과 Delta Rule을 complementary하다고 설명하는 이유가 이것이다

\((I - \beta_tk_tk_t^\top)\) - targeted memory control

  • 이 항은 현재 \(k_t\)와 관련된 state component에 선택적으로 영향을 준다. 따라서 전체 state를 동일하게 decay시키는 Mamba2와 달리 특정 key-value association을 수정할 수 있다.

결과적으로 GDN은 개념적으로

Global forgetting + Local/targeted correction

을 하나의 transition에 결합한 구조이다.

3.5. Online Learning / Fast Weight Interpretation

논문은 Gated DeltaNet을 단순한 RNN update가 아니라 online learning 또는 fast weight programming 관점에서도 해석한다. DeltaNet은 state \(S_t\)를 하나의 fast weight matrix로 보고 다음 regression objective를 최적화한다고 볼 수 있다.

$$\mathcal L(S_t) = \frac{1}{2}\vert\vert S_tk_t - v_t \vert\vert^2$$

이를 stochastic gradient descent로 한 step 업데이트하면

$$S_{t+1} = S_t - \beta_t \nabla \mathcal{L}(S_t) = S_t - \beta_t (S_t k_t - v_t) k_t^{\top} = S_t \left(I - \beta_t k_t k_t^{\top}\right) + \beta_t v_t k_t^{\top}$$

이 되어 DeltaNet update가 그대로 나온다. 여기서 \(\beta_t\)는 adaptive learning rate로 해석된다. 논문은 이 관점에서 Gated DeltaNet의 \(\alpha_t\)를 adaptive weight decay로 해석한다. 즉 GDN의 recurrent state는 고정된 memory가 아니라 sequence를 읽는 동안 계속 online regression으로 수정되는 fast weight이고, \(\alpha_t\)가 그 fast weight에 얼마나 많은 과거 정보를 남길지를 조절한다. 이 해석이 GDN을 이해하는 가장 중요한 두 번째 관점이다.

GDN≈test-time online regression+adaptive weight decay

3.6. 왜 Gate와 Delta Rule이 둘 다 필요한가?

논문은 S-NIAH를 이용해 이 부분을 직접 분석한다.

1

[Table 2] S-NIAH-1은 단순한 pass-key를 장기간 유지해야 하는 setting이다. 여기서는 과도한 decay가 오히려 memory retention을 해친다. 8K에서 DeltaNet은 98.8을 유지하지만 Mamba2는 30.4까지 감소하며, Gated DeltaNet은 91.8을 기록한다. 즉 Delta Rule이 long-term memorization을 보존하는 데 중요한 역할을 한다.

반대로 S-NIAH-2/3처럼 real-world essay가 포함되어 state에 많은 정보가 유입되는 경우에는 forgetting이 없는 DeltaNet이 memory collision을 겪는다. Gating을 가진 Mamba2와 Gated DeltaNet은 불필요한 정보를 제거할 수 있으며, Gated DeltaNet은 동시에 Delta Rule의 memorization 능력도 유지한다.

특히 S-NIAH-3의 2K에서는 Gated DeltaNet이 84.2, DeltaNet이 47.0, Mamba2가 47.6을 기록한다. 즉 단순히 Mamba2 또는 DeltaNet 중 하나를 선택하는 것이 아니라, “기억해야 할 것은 정밀하게 유지하면서 버릴 것은 빠르게 버리는 것”이 성능 차이를 만든다는 것이 논문의 핵심 주장이다.

3.7. Hardware-Efficient Chunkwise Training

GDN의 recurrence는 token을 하나씩 처리하면 본질적으로 sequential하다. 따라서 좋은 recurrent rule을 만들었다고 해도 그대로 구현하면 GPU의 tensor core를 충분히 활용하기 어렵다. 논문은 DeltaNet에서 사용된 WY representation과 UT transform을 gating까지 확장하여 sequence를 chunk 단위로 계산한다.

Gated DeltaNet에서 chunk 내부의 transformed value는 다음과 같이 계산된다.

$$U^{g}_{[t]} = \left[I + \operatorname{strictLower}\left(\operatorname{diag}(\beta_{[t]})\left(\Gamma_{[t]} \odot K_{[t]}K_{[t]}^{\top}\right)\right)\right]^{-1}\operatorname{diag}(\beta_{[t]})V_{[t]}$$

여기서 \(\Gamma_{[t]}\)가 \(\alpha\)에 따른 cumulative decay를 chunk 내부 연산에 포함한다. 이후 chunk state와 output은 matrix multiplication 중심으로 계산된다.

$$S_{[t+1]} = \overrightarrow{S}_{[t]} + \left(U^{g}_{[t]} - \overleftarrow{W}_{[t]}S_{[t]}^{\top}\right)^{\top}\overrightarrow{K}_{[t]}$$
$$O_{[t]} = \overleftarrow{Q}_{[t]}S_{[t]}^{\top} + \left(Q_{[t]}K_{[t]}^{\top} \odot M\right)\left(U^{g}_{[t]} - \overleftarrow{W}_{[t]}S_{[t]}^{\top}\right)$$

이를 통해 token-by-token recurrence를 그대로 실행하는 대신 chunk 안에서는 큰 matmul로 계산할 수 있고, tensor-core-friendly training이 가능해진다. 이것이 GDN이 단순히 이론적인 recurrent memory rule에 그치지 않고 실제 대규모 모델에 적용 가능한 핵심 이유이다.

3.8 Gated DeltaNet Block

1

Figure 1 기본 Gated DeltaNet은 Llama의 macro architecture처럼 token mixer와 SwiGLU MLP를 반복하지만, self-attention token mixer를 gated delta rule로 대체한다.

각 path는 다음과 같다.

$$\mathrm{Input} \rightarrow\begin{cases}q : \mathrm{Linear} \rightarrow \mathrm{ShortConv} \rightarrow \mathrm{SiLU} \rightarrow L2 \\k : \mathrm{Linear} \rightarrow \mathrm{ShortConv} \rightarrow \mathrm{SiLU} \rightarrow L2 \\v : \mathrm{Linear} \rightarrow \mathrm{ShortConv} \rightarrow \mathrm{SiLU} \\\alpha : \mathrm{Linear} \\\beta : \mathrm{Linear}\end{cases}$$

그 뒤 gated delta rule을 수행하고, output에 normalization과 gating을 적용한 후 output projection을 수행한다. \(q, k\)에 L2 normalization을 적용하는 것은 training stability를 위한 설계이다.

Hybrid Models 논문의 기본 Gated DeltaNet은 pure recurrent architecture이지만, 저자들도 fixed-size recurrent state가 retrieval 및 local comparison에서 여전히 한계를 가진다고 명시한다.

따라서 다음 hybrid들을 제안한다.

  • GatedDeltaNet-H1: Gated DeltaNet + SWA
  • GatedDeltaNet-H2: Mamba2 + Gated DeltaNet + SWA

linear recurrence만으로 모든 self-attention을 완전히 대체하는 것보다 attention과 결합하는 것이 유리할 수 있음을 인정한다.



4. Experiments

4.1. Main Results

1

동일한 1.3B parameter, FineWeb-Edu 100B-token training 조건에서 Gated DeltaNet은 recurrent baselines 중 가장 높은 overall commonsense reasoning 성능을 보인다. Avg. accuracy는 55.32로 Mamba2의 54.89와 DeltaNet의 52.14를 상회하며, LAMBADA perplexity도 12.17로 Mamba2의 12.56보다 낮다. 즉 gated delta rule이 Mamba2의 forgetting과 DeltaNet의 associative update를 결합했을 때 language modeling에서도 이점이 유지된다.

1

[Table 4] Real-world in-context retrieval에서는 Gated DeltaNet이 pure recurrent model 중 Avg. 30.6으로 Mamba2 29.8과 DeltaNet 26.2를 모두 상회한다. Hybrid architecture에서는 GatedDeltaNet-H2가 40.1로 Transformer++ 37.0과 Samba 37.3보다 높다. 저자들은 pure linear recurrent model과 attention 사이의 retrieval gap이 존재하지만, hybridization을 통해 이를 추가로 줄일 수 있다고 해석한다.

1

[Table 5] LongBench에서도 pure Gated DeltaNet은 Avg. 16.6으로 Mamba2 13.5와 DeltaNet 13.6보다 높으며, GatedDeltaNet-H2는 18.4를 기록한다. 특히 저자들은 single-doc QA, few-shot in-context learning, Code task에서 Gated DeltaNet의 이점을 강조한다.

4.2. Ablation Study

1

[Table S.1] Gated DeltaNet block을 naive Delta Rule로 변경하면 Avg-PPL이 27.35에서 30.87로 악화되고 Avg-Acc도 47.26에서 45.12로 감소한다. 이는 단순 DeltaNet이 아니라 gating이 결합된 update가 실제 성능 향상에 중요함을 보여준다.

또한 Short Conv 제거 시 Avg-Acc가 46.16, Output Gate 제거 시 45.46으로 감소한다. L2 normalization 역시 중요한 구성요소이며, 저자들은 head dimension 128을 performance와 computational efficiency 사이의 적절한 trade-off로 선택한다.



5. Conclusion

Contribution

  • [Gated Delta Rule] Mamba2의 data-dependent state decay와 DeltaNet의 selective key-value update를 하나의 recurrent transition으로 결합한다. 이를 통해 Mamba2보다 정밀한 key-value association learning과 DeltaNet보다 빠른 adaptive memory clearance를 동시에 제공한다.
  • [Hardware-efficient Training] DeltaNet의 WY-based parallelization을 gating까지 확장하여 recurrent gated delta rule을 chunkwise matrix computation으로 계산한다. 이에 따라 현대 GPU에서 효율적인 parallel training이 가능하다.
  • [Hybrid Architecture] Gated DeltaNet을 SWA 또는 Mamba2와 결합한 hybrid architecture를 구성하고, pure recurrent model보다 높은 retrieval, long-context performance와 training throughput을 확인한다.

Limitations

  • [Fixed-state Retrieval] 저자들은 linear Transformer의 fixed-size state가 retrieval에서 여전히 한계를 가지며 local shifts와 comparison modeling에도 약하다고 명시한다. 이것이 GatedDeltaNet-H1/H2에서 attention을 추가하는 직접적인 이유이다.
  • [Longer Context] Length extrapolation은 최대 20K sequence까지 평가되며, 저자들은 더 긴 sequence에서의 동작을 future work로 남긴다.



6. Gated DeltaNet은 정확히 무엇이 다른가?

가장 간단하게 세 모델을 비교하면 다음과 같다.

Mamba2

$$S_t = \alpha_t S_{t-1} + v_t k_t^{\top}$$
  • “전체 memory를 얼마나 남길 것인가?”를 잘 결정한다.
  • 하지만 \(S_{t-1} \rightarrow \alpha_tS_{t-1}\)이므로 모든 memory가 함께 decay한다.

DeltaNet

$$S_t = S_{t-1}\left(I - \beta_t k_t k_t^{\top}\right) + \beta_t v_t k_t^{\top}$$
  • “현재 key에 연결된 memory를 어떻게 고칠 것인가?”를 잘 결정한다.
  • 하지만 unrelated stale memory까지 한 번에 clear하지는 못한다.

Gated DeltaNet

$$S_t = S_{t-1}\underbrace{\alpha_t}_{\text{global forgetting}}\underbrace{\left(I - \beta_t k_t k_t^{\top}\right)}_{\text{selective update}} + \underbrace{\beta_t v_t k_t^{\top}}_{\text{new write}}$$

따라서 GDN의 핵심은 단순히 “Linear Attention에 Gate 하나를 추가했다”가 아니다.

  • Memory Retention+Targeted Correction+New Information Write를 하나의 recurrent update 안에서 학습하도록 만든 구조이다.

NR 카테고리 내 다른 글 보러가기

댓글 남기기