Full text not available for this paper

Summary (Overview)

  • Triadic linear attention generalizes linear attention by writing a triadic outer product of a key, a second key, and a value into a third-order (3D) tensor state, increasing state capacity from d2d^2 to d2Ed^2 E entries with minimal parameter overhead (only two extra projections).
  • The construction is compatible with data-dependent forgetting (per-slice gates along the second-key axis), the delta rule (joint erase before write), and chunkwise-parallel training via state tiling.
  • Applied to Gated DeltaNet (GDN) and scalar-gated linear attention (sGLA), triadic attention substantially improves long-context language modeling and recall at 400M and 1.3B parameter scales, outperforming alternatives that enlarge state via larger heads, wider values, more heads, or grouped values.
  • A pretrained dyadic model can be "upcycled" to triadic form during long-context extension, recovering 50–75% of the gains of training from scratch.
  • In GDN/Transformer hybrids, triadic states outperform larger key-value caches in perplexity while using significantly less memory at long contexts.

Introduction and Theoretical Foundation

Motivation

Recurrent neural networks (RNNs) compress history into a fixed-size memory state, allowing constant-time inference. The state size is crucial for performance—linear attention (Katharopoulos et al., 2020) extended vector-valued hidden states to matrix-valued states via a dyadic outer product of key and value vectors, enabling modern RNNs to rival Transformers. However, state size is fundamentally limited: a model with a fixed-size state cannot even copy sequences beyond a certain length (Jelassi et al., 2024), and modern linear attention models still degrade on recall-intensive and long-context tasks (Arora et al., 2024; Hsieh et al., 2024).

Naively increasing state size (more heads, larger value dimensions) grows parameter count significantly. The paper asks: How can we increase state size parameter-efficiently?

Theoretical Foundation: Tensor Product Representations

The paper frames linear attention as a tensor product representation (Smolensky, 1990), where:

  • The value is a filler, bound to a role (the key) via an outer product
  • The state StS_t is a superposition of bindings
  • Readout with the query is unbinding via contraction

Standard linear attention:

St=St−1+ktvt⊤,ot=St⊤qt=∑s≤t(qt⊤ks)vs(1)S_t = S_{t-1} + k_t v_t^\top, \quad o_t = S_t^\top q_t = \sum_{s \leq t} (q_t^\top k_s) v_s \tag{1}

This coincides with the fast-weight programmer of Schmidhuber (1992). The capacity is determined by the state shape: with dd-dimensional keys/values, StS_t has d2d^2 entries and can separate at most dd mutually orthogonal keys.

Key insight: A binding itself is a vector (of dimension d2d^2), which can be bound to a further role. Each nesting raises the tensor order by one. With orthonormal roles, an nn-th order tensor product representation can store dn−1d^{n-1} associations exactly. Binding each value to a second key increases capacity from dd to d2d^2 associations.


Methodology

Triadic Linear Attention

Each position produces, besides qt,kt,vt∈Rdq_t, k_t, v_t \in \mathbb{R}^d, a second key kt′k'_t and second query qt′∈REq'_t \in \mathbb{R}^E. The state is a third-order tensor St∈Rd×E×dS_t \in \mathbb{R}^{d \times E \times d}:

St=St−1+kt⊗kt′⊗vt,ot=St×1qt×2qt′=∑s≤t(qt⊤ks)(qt′⊤ks′)vs(2)S_t = S_{t-1} + k_t \otimes k'_t \otimes v_t, \quad o_t = S_t \times_1 q_t \times_2 q'_t = \sum_{s \leq t} (q_t^\top k_s)(q'^\top_t k'_s) v_s \tag{2}

where ⊗\otimes is the outer product and ×n\times_n is contraction along axis nn. For E=1E=1 with kt′=qt′=1k'_t = q'_t = 1, this reduces to standard linear attention (Equation 1).

The state has d2Ed^2 E entries (cubic in dd for E=dE=d), while parameters grow only by the two projections for kt′k'_t and qt′q'_t.

Forgetting and Delta Rule

Forgetting: Each slice St[:,e,:]S_t[:, e, :] gets its own decay gate αt,e\alpha_{t,e}:

St=St−1×2diag(αt)+kt⊗kt′⊗vt(3)S_t = S_{t-1} \times_2 \text{diag}(\alpha_t) + k_t \otimes k'_t \otimes v_t \tag{3}

Delta rule: The stored association is removed before writing:

St=St−1+βt(kt⊗kt′⊗vt−St−1×1kt×2kt′)(4)S_t = S_{t-1} + \beta_t \left( k_t \otimes k'_t \otimes v_t - S_{t-1} \times_1 k_t \times_2 k'_t \right) \tag{4}

By flattening the joint key κt=kt⊗kt′∈Rd⋅E\kappa_t = k_t \otimes k'_t \in \mathbb{R}^{d \cdot E}, this can be rewritten as Gated DeltaNet with key dimension d⋅Ed \cdot E:

St=(I−βtκtκt⊤)DtSt−1+βtκtvt⊤,ot=St⊤(qt⊗qt′)(5)S_t = (I - \beta_t \kappa_t \kappa_t^\top) D_t S_{t-1} + \beta_t \kappa_t v_t^\top, \quad o_t = S_t^\top (q_t \otimes q'_t) \tag{5}

Efficient Implementation

Chunkwise-parallel form: The sequence is split into chunks of CC positions. The masked attention within a chunk factorizes:

(qr⊗qr′)⊤(ks⊗ks′)=(qr⊤ks)(qr′⊤ks′)(q_r \otimes q'_r)^\top (k_s \otimes k'_s) = (q_r^\top k_s)(q'^\top_r k'_s)

so the attention matrix becomes QK⊤⊙R′QK^\top \odot R' where R′=tril(Q′K′⊤)R' = \text{tril}(Q'K'^\top), reducing cost from C2(d⋅E)C^2(d \cdot E) to C2(d+E)C^2(d + E).

Tiling the state: The state is split along the value axis into blocks of 32 columns per thread block, so no streaming multiprocessor ever holds the full third-order state. At E=8E=8, one head's state is 512 KiB (twice the register file of a Hopper SM), but one block is 128 KiB.

Experimental Setup

  • Models: 400M (24 layers, dmodel=1024d_{model}=1024, 8 heads) and 1.3B (24 layers, dmodel=2048d_{model}=2048, 16 heads) parameters, head dimension d=128d=128.
  • Training: 50 tokens per parameter (2.5× Chinchilla-optimal), pretrained on Fineweb-Edu with 4k context, then long-context extended to 64k on 5 tokens/parameter.
  • Baselines: Larger heads, wider values, grouped values, more heads—all matched for state size and parameter count. Transformer with RoPE, QK-norm, GQA-8.
  • Evaluation: PG19 perplexity by context position, WikiText perplexity, recall suite from Arora et al. (2024), RULER needle-in-a-haystack.

Empirical Validation / Results

Capacity in Isolation (MQAR)

On multi-query associative recall (MQAR), every doubling of EE shifts the accuracy curve right by roughly a doubling of NN. For E=16E=16, the model stores ~16× more associations than ordinary linear attention, with only a 1.08× increase in non-embedding parameters.

Main Results

Figure 2 shows that at 400M and 1.3B scales, Triadic GDN with E=8E=8 predicts the next token better than the Transformer even at 64k context, and substantially improves recall on tasks requiring copying from long documents (FDA, SWDE). RULER needle-in-a-haystack accuracy stays high up to longer contexts with larger states.

State-Matched Comparison

Table 1 compares triadic attention to alternative state-enlargement methods:

ModelState (MB)Params (M)Wiki. ↓PG19 ≤4k ↓PG19 4k–16k ↓PG19 16k–64k ↓Recall ↑
GDN (base)6.3380.911.2515.0314.3814.1526.2
Larger heads (dd=256)12.6380.711.2215.1014.4214.1726.6
Grouped values (2/key)12.6378.211.1815.0514.3714.1227.1
Wider values (dvd_v=256)12.6377.911.2015.0714.3914.1427.4
More heads (16)12.6381.611.3215.2514.5414.2827.5
Triadic (EE=2)12.6381.911.0814.9014.2113.9528.4
Triadic (EE=4)25.2383.010.9314.8014.0813.7931.1

Triadic attention achieves consistently lower perplexity on every range and highest recall. Alternative methods barely improve at 2× and degrade at 4× due to MLP width reduction needed to match parameters.

Upcycling

A pretrained model can be expanded from E=1E=1 to E=8E=8 by copying forget gates to every slice and initializing second key/query projections from scratch, then performing long-context extension. Upcycled models recover 50–75% of the gains of training from scratch.

Hybrid Models

In a 3:1 GDN/GQA-8 hybrid, triadic state (E=4E=4) achieves the lowest perplexity at every context range and best NIAH average, while using less memory than enlarging the key-value cache beyond ~4.6k tokens (almost half the memory at 64k).

Training Efficiency

Triadic GDN with E=8E=8 is 3% slower than Transformer at 4k tokens and 5.1× faster at 64k. Overhead over vanilla GDN: 28–30% for E=8E=8, 14–15% for E=4E=4, 9–11% for E=2E=2.

Ablations

  • Key dimensions: At fixed joint key dimension dk⋅E=1024d_k \cdot E = 1024, keeping dkd_k large (128) and EE small (8) slightly outperforms balanced splits.
  • Activation: Non-negative activations (softplus, sigmoid) outperform signed activations (SiLU, linear) for the second key/query, since signed entries allow cancellation across slices.

Theoretical and Practical Implications

Theoretical Significance

  1. Tensor product view of memory: The paper provides a principled framework for increasing RNN state capacity by raising the tensor order, connecting to Smolensky's tensor product representations and higher-order associative memories (Baldi & Venkatesh, 1987), where storable patterns scale as dn−1d^{n-1}.

  2. Parameter-efficiency: The ratio of state size to parameter count increases asymptotically—an EE-dimensional second key gives an EE-fold state increase with only two extra projections.

  3. Memory dynamics analysis: Triadic GDN learns substantially larger delta-rule write strengths (β\beta mean increases from 0.31 to 0.50) and develops a broad range of memory timescales: median half-lives range from 0.20 tokens (shortest-lived slice) to 1080 tokens (longest-lived), versus a single 5.1-token half-life for vanilla GDN.

Practical Implications

  1. Long-context modeling: Triadic linear attention enables pure linear RNNs to match or exceed Transformer perplexity at 64k context with a smaller state than the Transformer's key-value cache.

  2. Hybrid architectures: Enlarging linear-attention states is more parameter- and memory-efficient than enlarging softmax attention key-value caches in hybrid models.

  3. Upcycling: Pretrained dyadic models can be converted to triadic during long-context extension, suggesting a staged training recipe where state size grows with context length.

  4. Training cost: The 15–30% training overhead for E=4E=4/E=8E=8 is a practical tradeoff, but the models remain faster than Transformers beyond a few thousand tokens.


Conclusion

The paper introduces triadic linear attention, which updates a third-order tensor state with the outer product of three vectors (key, second key, value), increasing the capacity of constant-time sequence mixers with minimal parameter overhead. Key findings:

  • Triadic Gated DeltaNet consistently lowers perplexity and improves recall at 400M and 1.3B scales, outperforming all alternative state-enlargement methods at matched state size.
  • The construction admits efficient chunkwise-parallel kernels (via Kronecker structure factorization and state tiling), keeping training within 1.3× of vanilla GDN for E=8E=8.
  • The larger state is genuinely used for distant context, with memory slices developing diverse timescales and stronger write strengths.

Limitations: Training slowdown (15–30%), performance on some recall tasks still trails Transformers, and the construction is only applied to GDN and sGLA.

Future directions: The third-order state opens possibilities for entirely new linear attention variants; revisiting other matrix-state developments (e.g., test-time training) for 3D states; applying triadic construction to mixture-of-experts architectures where smaller key/query projections matter.


Key Formulas

EquationDescription
(1)Standard linear attention: St=St−1+ktvt⊤S_t = S_{t-1} + k_t v_t^\top
(2)Triadic linear attention: St=St−1+kt⊗kt′⊗vtS_t = S_{t-1} + k_t \otimes k'_t \otimes v_t
(3)With per-slice forgetting: St=St−1×2diag(αt)+kt⊗kt′⊗vtS_t = S_{t-1} \times_2 \text{diag}(\alpha_t) + k_t \otimes k'_t \otimes v_t
(4)With delta rule: St=St−1+βt(kt⊗kt′⊗vt−St−1×1kt×2kt′)S_t = S_{t-1} + \beta_t(k_t \otimes k'_t \otimes v_t - S_{t-1} \times_1 k_t \times_2 k'_t)
(5)Joint-key form: St=(I−βtκtκt⊤)DtSt−1+βtκtvt⊤S_t = (I - \beta_t \kappa_t \kappa_t^\top) D_t S_{t-1} + \beta_t \kappa_t v_t^\top
(6)Chunkwise parallel form (standard)
(7)Factorized masked attention: QK⊤⊙R′QK^\top \odot R'
(8)Tiled state chunkwise form

Related papers