Recurrent Memory: Extending Transformer Context

Michael BrenndoerferUpdated July 6, 202568 min read

Part of Language AI Handbook

Explains how Transformer-XL uses segment-level recurrence to extend effective context length.

Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.

Article links

Make inline references clickable

Recurrent Memory

Transformers process sequences in fixed-length segments. When a document exceeds the context window, the standard approach is to truncate or split it, processing each piece independently. But this independence comes at a cost: information from earlier segments vanishes entirely. A pronoun in segment 3 cannot resolve to its antecedent in segment 1 because segment 1 no longer exists in the model's view. The transformer has, in effect, amnesia: it wakes up at the start of each segment with no memory of what came before.

This limitation becomes painfully apparent when working with real-world text. A legal brief references a clause defined pages earlier. A novel's character acts on motivations established chapters back. A scientific paper's conclusion depends on definitions in its introduction. When we split these documents into segments and process each independently, we destroy the long-range connections that give language its meaning. The model sees each fragment in isolation and cannot form the kind of coherent, document-level understanding that these texts demand.

Transformer-XL introduced a solution that seems almost obvious in retrospect: what if we kept the hidden states from the previous segment and let the current segment attend to them? This segment-level recurrence creates a form of memory that extends effective context far beyond the training sequence length. The model processes sequences one segment at a time, but each segment can "remember" what came before through cached hidden states. Think of it like reading a book with a bookmark that marks your place and carries forward a compressed summary of everything you have read. You do not reread from the beginning, but that summary shapes how you understand each new page.

Transformer-XL extends context without recomputing attention over the entire history. Storing and recomputing attention over all past tokens would cost quadratic time in the total sequence length, which is exactly what we wanted to avoid by using segments in the first place. Instead, the cached hidden states serve as a fixed-length, information-dense memory that can be attended to with only linear overhead relative to the segment size. The tradeoff is that information is compressed into a fixed-dimensional representation rather than preserved exactly, but this compression is what makes the approach computationally tractable.

This chapter explores how Transformer-XL implements recurrent memory, why it requires relative positional encodings, and what limitations remain. Understanding this approach illuminates a key tension in long-context modeling: the tradeoff between computational efficiency and true bidirectional context. We will also see that solving the memory problem immediately creates a new problem: positions. When hidden states from one segment are concatenated with hidden states from another, the model must understand where each token lies in the combined sequence. Standard absolute position encodings fail here in a subtle but important way, leading to one of Transformer-XL's most elegant contributions.

Historical Context

Transformer-XL was published by Dai et al. from Carnegie Mellon University and Google Brain in 2019, in a paper titled "Transformer-XL: Attentive Language Models Beyond a Fixed-Length Context." At the time, the dominant approach to long-sequence modeling was either truncation or independent segment processing. Recurrent neural networks had long handled variable-length sequences through their hidden state, but they suffered from vanishing gradients and could not parallelize across time steps during training. Transformer-XL combined the parallelism of transformers with an RNN-like recurrence mechanism, achieving state-of-the-art results on language modeling benchmarks including enwiki8, text8, WikiText-103, and One Billion Word. The paper also introduced relative position encodings, which have since become influential in their own right and were adopted by subsequent models including XLNet. Transformer-XL demonstrated that transformers could surpass RNNs in quality and effective context length, setting the stage for the long-context modeling research that followed.

The Segment Boundary Problem

Standard transformers suffer from what the Transformer-XL paper calls "context fragmentation." When you split a long document into fixed-length segments, each segment is processed without any knowledge of its neighbors. The model sees each chunk as an independent sequence. This limitation affects predictions directly: the model's predictions for the first few tokens of each new segment are systematically worse than predictions in the middle of a segment, because those early tokens have no context to draw on. The model is perpetually starting cold.

The severity of context fragmentation depends on both the segment length and the nature of the text. For documents with frequent long-range dependencies, every segment boundary is a potential failure point. Consider coreference resolution, the task of linking pronouns to their referents. If "the researcher" appears in segment 1 and "she" appears in segment 2, a model with context fragmentation cannot resolve the pronoun, even if a human reader would find the reference obvious. The same problem arises with discourse connectives ("As mentioned earlier..."), ellipsis resolution, and any form of reasoning that requires holding multiple pieces of information in mind simultaneously.

Think of context fragmentation as reading a story through a narrow window that reveals only one page at a time. Each page you read, you must discard the previous one. When a character on page 15 says "You know what I mean," you cannot, because the context that would let you understand was on page 12, which has already been discarded. The model faces this exact situation at every segment boundary.

The quantitative impact of context fragmentation is stark. Consider a 4096-token document processed with a 512-token window. Only a small fraction of all possible token-to-token dependencies are visible within any single segment. The vast majority of potentially informative connections between tokens are simply cut off.

Context Fragmentation

Context fragmentation occurs when a transformer processes a long sequence in fixed-length segments without information flow between them. Each segment starts fresh, losing all context from previous segments regardless of semantic continuity.

Consider processing a 4096-token document with a 512-token context window. The naive approach splits this into 8 segments, processing each independently. Token 513 (the first token of segment 2) cannot attend to token 512 (the last token of segment 1). They're as disconnected as tokens from completely different documents.

In[3]:
Code
# Simulate context fragmentation
document_length = 4096
segment_length = 512
num_segments = document_length // segment_length

# For each segment, the maximum dependency distance is limited
# to within-segment connections only
max_within_segment_distance = segment_length - 1
max_possible_distance = document_length - 1

# Calculate the fraction of potential dependencies that are visible
visible_dependencies_per_segment = segment_length * (segment_length - 1) // 2
total_possible_dependencies = document_length * (document_length - 1) // 2
visibility_fraction = (
    visible_dependencies_per_segment * num_segments
) / total_possible_dependencies
Out[4]:
Console
Document length: 4,096 tokens
Segment length: 512 tokens
Number of segments: 8

Within each segment:
  Maximum attention distance: 511 tokens
  Visible dependencies per segment: 130,816

Across entire document:
  Total possible dependencies: 8,386,560
  Fraction visible with fragmentation: 12.48%

Less than 13% of potential token-to-token dependencies are visible when processing with context fragmentation. Cross-segment dependencies, which may carry information such as coreference chains or long-range discourse structure, are completely invisible.

Out[5]:
Visualization
Block diagonal attention pattern showing 4 isolated segments with no cross-segment connections.
Context fragmentation in standard transformer processing. Each segment is processed independently, with no information flow across segment boundaries. The attention matrix shows isolated blocks along the diagonal, indicating that tokens can only attend within their own segment. The red dashed lines mark the hard boundaries where attention cannot cross.

The attention pattern reveals the fundamental limitation: each segment is an island. No matter how important a reference in segment 1 might be for understanding segment 4, that information cannot flow through the attention mechanism. The red boundaries are walls, not suggestions.

Transformer-XL: Segment-Level Recurrence

Now that we understand the problem, let's explore how Transformer-XL solves it. The solution is remarkably simple: instead of discarding the previous segment entirely, cache its hidden states and make them available during attention computation for the current segment.

Think of it like this: when you read a new paragraph, you don't forget the previous one. You carry forward a mental summary of what came before. That summary isn't the original words themselves, but your processed understanding of them. Transformer-XL does exactly this, but with hidden states instead of mental summaries. The cached states are not raw token embeddings. They are processed representations that already encode relationships between tokens, contextual information from within that segment, and whatever higher-level abstractions the network has learned to extract.

The key insight is that we do not need to preserve the original tokens from previous segments. We need to preserve the information they contained in a form that is useful for processing the current segment. Hidden states serve this purpose because they encode meaning rather than surface form alone. A token's hidden state at layer 3 is less like "the word 'cat'" and more like "a specific cat being discussed in a context involving a mat and some sitting." That richer representation is what gets cached and carried forward.

This approach also fits naturally with the autoregressive nature of language modeling. When generating text one token at a time, we process segments sequentially anyway. Caching the hidden states of the just-processed segment and feeding them to the next segment adds minimal overhead to the already-sequential processing. There is no need to maintain an exponentially growing memory buffer; we only ever need the most recent segment's hidden states.

Segment-Level Recurrence

Segment-level recurrence is a technique where hidden states from the previous segment are cached and concatenated with the current segment's keys and values during attention computation. This allows information to flow across segment boundaries without recomputing attention over the entire history.

The Mechanism Step by Step

The recurrence operates at each layer independently. When processing layer nn of segment τ\tau, the model:

  1. Retrieves the cached hidden states hτ−1n−1\mathbf{h}_{\tau-1}^{n-1} from processing the previous segment at layer n−1n-1
  2. Concatenates these cached states with the current segment's hidden states
  3. Computes attention where queries come only from the current segment, but keys and values span both the cached and current states
  4. Caches the current segment's output hidden states for use when processing the next segment

This creates an asymmetric attention pattern: current tokens can "look back" at cached tokens, but we never recompute outputs for the cached tokens. They're frozen representations from the previous forward pass. This asymmetry is not a limitation but a design choice. Recomputing the outputs of cached tokens would require backpropagating gradients through an unbounded history, which is exactly the computational cost we are trying to avoid. By treating cached states as read-only context, Transformer-XL achieves information flow across segments at constant computational cost per segment.

Another subtle but important point is that the cached states are from layer n−1n-1, not layer nn. When processing layer nn, we concatenate the previous segment's layer n−1n-1 output with the current segment's layer n−1n-1 output and compute attention. This is not an implementation detail but a consequence of how information flows through a multi-layer network. At layer nn, the current segment's representation already incorporates information from lower layers. The cached representation we need to extend must be at the same "processing depth" to combine coherently with the current representation.

The Mathematical Formulation

Let's formalize this mechanism. Suppose we're processing segment τ\tau, which contains LL tokens. The previous segment's hidden states, also of length LL (or some memory length MM), have been cached. At layer nn, we want to compute attention over an extended context that includes both segments.

The extended context for attention at layer nn concatenates the stop-gradient of the previous segment's hidden states with the current segment's hidden states. Concretely:

h~τn−1=StopGrad(hτ−1n−1)∘hτn−1\tilde{\mathbf{h}}_\tau^{n-1} = \text{StopGrad}(\mathbf{h}_{\tau-1}^{n-1}) \circ \mathbf{h}_\tau^{n-1}

where:

  • h~τn−1\tilde{\mathbf{h}}_\tau^{n-1}: the extended hidden states combining previous and current segments, used as keys and values
  • hτ−1n−1\mathbf{h}_{\tau-1}^{n-1}: cached hidden states from the previous segment at layer n−1n-1
  • hτn−1\mathbf{h}_\tau^{n-1}: current segment's hidden states at layer n−1n-1
  • StopGrad(⋅)\text{StopGrad}(\cdot): stops gradient flow to prevent backpropagating through the cached states
  • ∘\circ: concatenation along the sequence dimension, resulting in a tensor of length M+LM + L

Why does StopGrad matter? If we allowed gradients to flow back through the cached states, we would need to store the computation graph for the previous segment's forward pass, then the segment before that, and so on, going back all the way to the beginning of the document. This would make the memory cost proportional to the total document length, defeating the purpose of segmentation. By stopping gradients at the cached states, each segment's backward pass touches only the current segment's computation, keeping memory usage constant regardless of document length.

The attention computation then becomes:

qτn=hτn−1Wqn,kτn=h~τn−1Wkn,vτn=h~τn−1Wvn\mathbf{q}_\tau^n = \mathbf{h}_\tau^{n-1} W_q^n, \quad \mathbf{k}_\tau^n = \tilde{\mathbf{h}}_\tau^{n-1} W_k^n, \quad \mathbf{v}_\tau^n = \tilde{\mathbf{h}}_\tau^{n-1} W_v^n hτn=Attention(qτn,kτn,vτn)\mathbf{h}_\tau^n = \text{Attention}(\mathbf{q}_\tau^n, \mathbf{k}_\tau^n, \mathbf{v}_\tau^n)

where:

  • qτn∈RL×d\mathbf{q}_\tau^n \in \mathbb{R}^{L \times d}: queries from the current segment only
  • kτn,vτn∈R(L+M)×d\mathbf{k}_\tau^n, \mathbf{v}_\tau^n \in \mathbb{R}^{(L + M) \times d}: keys and values from the extended context (current + cached)
  • Wqn,Wkn,WvnW_q^n, W_k^n, W_v^n: learnable projection matrices for layer nn
  • MM: the length of the cached memory (typically equal to segment length LL)

The critical detail is that queries come only from the current segment while keys and values include the cached previous segment. This asymmetry is intentional: we want current tokens to attend to past context, but we don't want to regenerate outputs for past tokens. The cached states are read-only: they provide context but don't receive updates.

Why does this formula make sense? Notice that the query dimension remains LL (only current tokens ask questions) while the key and value dimensions are L+ML + M (both cached and current tokens can be referenced). The attention matrix has shape (L,L+M)(L, L+M), which means each current token computes a score against every cached token and every preceding current token. The resulting weighted sum draws from both the cached and current portions of the extended context, exactly what we want.

Implementing the Mechanism

Let's translate this into code. The core operation is straightforward: concatenate cached and current hidden states, project to queries/keys/values, and compute attention with appropriate masking.

In[6]:
Code
import numpy as np


def transformer_xl_attention(
    current_hidden: np.ndarray,
    cached_hidden: np.ndarray,
    W_q: np.ndarray,
    W_k: np.ndarray,
    W_v: np.ndarray,
    d_k: int,
) -> tuple[np.ndarray, np.ndarray]:
    """
    Compute Transformer-XL style attention with segment-level recurrence.

    Args:
        current_hidden: Current segment hidden states, shape (L, d)
        cached_hidden: Cached previous segment states, shape (M, d)
        W_q, W_k, W_v: Projection matrices, shape (d, d)
        d_k: Key dimension for scaling

    Returns:
        attention_output: Output for current segment, shape (L, d)
        attention_weights: Attention pattern, shape (L, L+M)
    """
    # Concatenate for extended context
    extended_hidden = np.concatenate([cached_hidden, current_hidden], axis=0)

    # Project to queries, keys, values
    # Queries: only from current segment
    queries = current_hidden @ W_q  # (L, d)
    # Keys and values: from extended context
    keys = extended_hidden @ W_k  # (L+M, d)
    values = extended_hidden @ W_v  # (L+M, d)

    # Compute attention scores
    scores = queries @ keys.T / np.sqrt(d_k)  # (L, L+M)

    # Apply causal masking
    # Current segment tokens can attend to all of cache + their own past
    L = current_hidden.shape[0]
    M = cached_hidden.shape[0]
    mask = np.ones_like(scores) * float("-inf")
    for i in range(L):
        # Can attend to all cached tokens + tokens up to and including position i
        mask[i, : M + i + 1] = 0
    scores = scores + mask

    # Softmax
    exp_scores = np.exp(scores - np.max(scores, axis=-1, keepdims=True))
    attention_weights = exp_scores / exp_scores.sum(axis=-1, keepdims=True)

    # Weighted sum of values
    output = attention_weights @ values  # (L, d)

    return output, attention_weights
In[7]:
Code
# Demonstrate the mechanism with a simple example
rng = np.random.default_rng(42)

segment_length = 8
memory_length = 8
hidden_dim = 16
d_k = hidden_dim

# Simulate hidden states
current_segment = rng.standard_normal((segment_length, hidden_dim)) * 0.1
cached_segment = rng.standard_normal((memory_length, hidden_dim)) * 0.1

# Random projection matrices (in practice, these are learned)
W_q = rng.standard_normal((hidden_dim, hidden_dim)) * 0.1
W_k = rng.standard_normal((hidden_dim, hidden_dim)) * 0.1
W_v = rng.standard_normal((hidden_dim, hidden_dim)) * 0.1

output, attention_weights = transformer_xl_attention(
    current_segment, cached_segment, W_q, W_k, W_v, d_k
)
Out[8]:
Console
Current segment shape: (8, 16)
Cached segment shape: (8, 16)
Attention weights shape: (8, 16)
Output shape: (8, 16)

Attention weight distribution for token 0 of current segment:
  Attention to cached tokens (0-7): 0.889
  Attention to current token (8): 0.111

The output shows that token 0 of the current segment allocates substantial attention to the cached memory. This is expected: with no preceding tokens in the current segment, the memory provides all available context. The attention weights reveal how information flows. The first token in the current segment can attend to all 8 cached tokens plus itself. Later tokens in the current segment can attend to even more context: all cached tokens plus all preceding tokens in the current segment.

Out[9]:
Visualization
Heatmap showing attention weights with a vertical segment boundary, where current tokens attend to both cached and current positions.
Transformer-XL attention pattern showing segment-level recurrence. Current segment tokens (positions 8-15) can attend to cached previous segment tokens (positions 0-7) as well as their own causal context. The triangular pattern in the right half reflects causal self-attention, while the rectangular left half shows uniform access to cached memory.

The attention pattern shows the distinctive Transformer-XL signature: a triangular pattern in the current segment (causal self-attention) combined with a rectangular region on the left (attention to cached memory). Every token in the current segment can see the entire cached memory, and each token can additionally see all the preceding tokens in the current segment.

Effective Context Length

The recurrence mechanism creates a dependency chain across segments. Information from segment 1 flows to segment 2 through the cached hidden states. Segment 2's hidden states then carry forward information to segment 3. This chain means the effective context length grows beyond a single segment, bounded by how far information can propagate through the hidden states.

For an NN-layer transformer with segment length LL, the maximum effective context length is O(N×L)O(N \times L). To understand why, consider how information propagates layer by layer:

  • Layer 1 of segment τ\tau receives cached states from layer 0 of segment τ−1\tau - 1, giving it access to 1 previous segment
  • Layer 2 of segment τ\tau receives cached states from layer 1 of segment τ−1\tau - 1, which already incorporated information from segment τ−2\tau - 2 at its own layer 1
  • Layer nn can potentially access information from nn segments back through this chain of cached representations

The depth of the network multiplies the effective context window: the same depth that provides representational power also provides temporal reach. A deeper network can represent more complex patterns and carry information farther back.

Think of this layer-by-layer propagation as a game of telephone played upward through the layers. At layer 1, each segment's representation contains information from at most the immediately preceding segment. At layer 2, each representation has been informed by a layer-1 representation that already incorporated one segment's worth of history. By the time information reaches the top layer, it has been aggregated across NN segments, even though no single attention computation ever looked more than one segment back.

The key insight here is that depth and context are not independent in Transformer-XL. Every extra layer you add to the network also extends the effective context window by one segment length. This means increasing model depth is simultaneously an investment in representational capacity and in temporal reach. Conversely, it also means that shallow models will have limited effective context even when using segment-level recurrence.

In[10]:
Code
def compute_effective_context(num_layers: int, segment_length: int) -> dict:
    """
    Compute the effective context length for Transformer-XL.

    The effective context grows with depth because information
    propagates one segment further back at each layer.
    """
    # At layer n, information can potentially come from n segments back
    # This is because cached states at layer n contain information
    # that was already aggregated from previous segments at layer n-1

    # Direct attention span (what a single layer can see)
    direct_span = 2 * segment_length  # current + one cached segment

    # Maximum theoretical span considering all layers
    max_theoretical_span = num_layers * segment_length + segment_length

    # Information decay means effective span is somewhat less
    # Upper layers have access to more distant information but it's diluted

    return {
        "num_layers": num_layers,
        "segment_length": segment_length,
        "direct_attention_span": direct_span,
        "max_theoretical_span": max_theoretical_span,
        "context_multiplier": max_theoretical_span / segment_length,
    }


# Compare different configurations
configs = [
    (6, 128),  # Small model
    (12, 256),  # Medium model
    (24, 512),  # Large model
]

results = [compute_effective_context(n, l) for n, l in configs]
Out[11]:
Console
Effective Context Length in Transformer-XL

Configuration         Direct Span Max Theoretical   Multiplier
--------------------------------------------------------------
6L × 128L                     256             896          7.0x
12L × 256L                    512           3,328         13.0x
24L × 512L                  1,024          12,800         25.0x

A 24-layer model with 512-token segments can theoretically access information from over 13,000 tokens ago, even though it only directly attends to 1,024 tokens per layer. The depth of the network amplifies the effective context. This explains why Transformer-XL's original paper used configurations with many layers: the authors were adding expressiveness and extending reach at the same time.

It is worth being clear about what "theoretical" means here. The numbers above represent the maximum possible reach under the assumption that information propagates perfectly and without loss through each cached representation. In practice, information degrades as it passes through the hidden state bottleneck. We will examine this degradation in the limitations section. The theoretical maximum sets an upper bound; the practical reach depends on the nature of the information and how much the network has learned to preserve it during training.

Out[12]:
Visualization
Line plot showing linear growth of effective context with layer depth, compared to constant direct attention span.
Effective context growth with layer depth in Transformer-XL. The maximum theoretical context (solid line) grows linearly with the number of layers, while direct attention span (dashed line) remains constant at twice the segment length. For a 24-layer model with 512-token segments, the effective context exceeds 12,000 tokens. This shows how network depth multiplies context reach.
Out[13]:
Visualization
Diagram showing expanding context reach at higher transformer layers, with layer 1 seeing 1 previous segment and layer 6 seeing 6 previous segments.
Information reach by layer in Transformer-XL. Each colored row represents a transformer layer, and the width of the colored band shows how many prior segments that layer can access through the recurrence chain. Layer 1 can only see one segment back through direct cached attention. Layer 6 can access information from up to six segments back, because its cached input already contains aggregated information from the prior five segments.

The visualization shows how context reach expands with network depth. Layer 1 can only see the immediately preceding segment through cached states. Layer 6, however, has access to information from 6 segments back because the cached states at layer 5 already incorporated information that propagated through the previous segment's entire network.

The Position Encoding Problem

Standard absolute position encodings break under segment-level recurrence. If we use learned or sinusoidal position embeddings based on absolute positions within each segment, we encounter a fundamental inconsistency.

Consider a token at position 5 in segment τ\tau. In the previous segment τ−1\tau - 1, there was also a token at position 5. Both receive the same absolute position encoding. But from the perspective of the current segment, these tokens are at very different distances: position 5 in the current segment is "here," while position 5 in the cached segment is 8 positions back (if segments have length 8).

The problem is even deeper than this. Absolute positions within a segment reset at the start of each segment. Position 0 in segment 1 and position 0 in segment 2 both have the same position encoding, but they are separated by the entire length of segment 1. When we concatenate their hidden states and ask the attention mechanism to use position information to understand their relationship, we give it the same encoding for tokens that are far apart. The attention mechanism cannot distinguish between "this is my first token" and "this is the first token in some previous segment."

Think of it like a book where every chapter starts its page numbers over from 1. If you're on page 5 of chapter 3 and want to cite page 5 of chapter 1, both references look identical in the page numbering system, but they point to completely different locations. The reader needs additional information (the chapter number) to disambiguate. Transformer-XL solves this analogous problem by replacing absolute position numbers with relative distance measures, so the position encoding always says "this token is kk steps before me" rather than "this token is at absolute position nn."

In[14]:
Code
def demonstrate_position_conflict():
    """
    Show how absolute position encodings create confusion
    in segment-level recurrence.
    """
    segment_length = 8

    # Absolute positions assigned during each segment's processing
    segment_tau_minus_1_positions = list(
        range(segment_length)
    )  # [0, 1, 2, ..., 7]
    segment_tau_positions = list(range(segment_length))  # [0, 1, 2, ..., 7]

    # True temporal distances from perspective of segment tau
    # Cached tokens are at positions -8 to -1 relative to segment tau's start
    true_distances_cached = list(range(-segment_length, 0))  # [-8, -7, ..., -1]
    true_distances_current = list(range(segment_length))  # [0, 1, ..., 7]

    return {
        "cached_absolute_pos": segment_tau_minus_1_positions,
        "current_absolute_pos": segment_tau_positions,
        "cached_true_distance": true_distances_cached,
        "current_true_distance": true_distances_current,
    }


pos_info = demonstrate_position_conflict()
Out[15]:
Console
Position Encoding Conflict in Segment-Level Recurrence

Cached Segment (τ-1):
  Absolute positions (as encoded): [0, 1, 2, 3, 4, 5, 6, 7]
  True distance from current τ:    [-8, -7, -6, -5, -4, -3, -2, -1]

Current Segment (τ):
  Absolute positions (as encoded): [0, 1, 2, 3, 4, 5, 6, 7]
  True distance from current τ:    [0, 1, 2, 3, 4, 5, 6, 7]

Problem: Position 5 in cached segment and position 5 in current segment
both have the same absolute encoding, but their true distances differ by 8!

If absolute position encodings were used, the model would receive conflicting signals. Two tokens with identical position encodings would be at different temporal distances. The attention mechanism, which relies on position information to understand sequence structure, would be confused.

Transformer-XL solves this with relative position encodings. Instead of encoding absolute positions and adding them to token embeddings, the model directly encodes the relative distance between query and key positions in the attention computation itself.

Relative Position Encoding in Transformer-XL

We've established that segment-level recurrence breaks absolute position encodings. A token at position 5 in the cached segment and position 5 in the current segment have the same absolute encoding, yet they are 8 positions apart from the perspective of the current segment. How do we fix this? The answer lies in rethinking what position information attention needs.

Attention computes a score between every query and every key to decide how much information should flow from each key to each query. Position information enters this computation to reflect the fact that two tokens at different positions should interact differently than two tokens at the same position. But the original formulation bakes in a specific assumption: that absolute position matters. Transformer-XL's relative position encoding challenges this assumption and replaces it with a more fundamental one: what matters is not where each token is, but how far apart the tokens are.

This changes the role of position information. The same pair of query and key content at a distance of 3 positions should produce the same position-influenced attention score regardless of whether that query is at absolute position 10 or absolute position 1000. The distance is always 3, so the position contribution to attention should always be the same. This property is called translation equivariance, and it is exactly what relative position encodings provide.

Think of relative position encoding as teaching the model about spatial relationships rather than absolute coordinates. Just as you naturally say "the store is two blocks north of here" rather than "the store is at GPS coordinates 37.7749, -122.4194," relative encodings describe relationships, not locations. Two tokens that are adjacent will always have a relative distance of 1, regardless of where they appear in the document.

Why Relative Distance Matters

When you read a sentence like "The cat sat on the mat because it was tired," you understand that "it" refers to "cat" not because of their absolute positions in the document, but because of their relative proximity. The pronoun comes shortly after its antecedent. This observation is the key insight: attention cares about how far apart tokens are, not where they are in absolute terms.

Consider two identical queries, one at position 10 and one at position 100, both attending to keys that are 3 positions before them. If the tokens involved have the same content, shouldn't these attention computations behave similarly? With absolute position encodings, they don't, because positions 7 and 97 have completely different encodings. With relative position encodings, they do, because "3 positions back" always means the same thing.

The implications for generalization are significant. If a model trained on sequences of length 512 learns that pronouns tend to refer to nouns within the previous 20 tokens, that pattern should generalize to any position in any sequence. Absolute position encodings undermine this generalization by making the same pattern look different at different absolute positions. Relative position encodings make the pattern position-agnostic, which is exactly what we want.

Decomposing Standard Attention

To understand how Transformer-XL achieves relative position encoding, we need to first dissect how position information enters standard attention. In the original transformer, the attention score between a query at position ii and a key at position jj starts as a simple dot product:

scoreij=qi⊤kj\text{score}_{ij} = \mathbf{q}_i^\top \mathbf{k}_j

where:

  • scoreij\text{score}_{ij}: the attention score determining how much position ii attends to position jj
  • qi\mathbf{q}_i: the query vector at position ii
  • kj\mathbf{k}_j: the key vector at position jj

But where does position come in? The original transformer adds position embeddings to token embeddings before projecting to queries and keys. So the query at position ii is Wq(xi+pi)W_q(\mathbf{x}_i + \mathbf{p}_i) and the key at position jj is Wk(xj+pj)W_k(\mathbf{x}_j + \mathbf{p}_j). Substituting these into the dot product and using the distributive property of matrix multiplication:

scoreij=(xi+pi)⊤Wq⊤Wk(xj+pj)\text{score}_{ij} = (\mathbf{x}_i + \mathbf{p}_i)^\top W_q^\top W_k (\mathbf{x}_j + \mathbf{p}_j)

where:

  • xi,xj\mathbf{x}_i, \mathbf{x}_j: token embeddings at positions ii and jj
  • pi,pj\mathbf{p}_i, \mathbf{p}_j: absolute position embeddings for positions ii and jj
  • Wq,WkW_q, W_k: learnable query and key projection matrices

This is where the magic of algebra reveals hidden structure. When we expand this product using the distributive property, we get four distinct terms:

scoreij=xi⊤Wq⊤Wkxj⏟content-content+xi⊤Wq⊤Wkpj⏟content-position+pi⊤Wq⊤Wkxj⏟position-content+pi⊤Wq⊤Wkpj⏟position-position\text{score}_{ij} = \underbrace{\mathbf{x}_i^\top W_q^\top W_k \mathbf{x}_j}_{\text{content-content}} + \underbrace{\mathbf{x}_i^\top W_q^\top W_k \mathbf{p}_j}_{\text{content-position}} + \underbrace{\mathbf{p}_i^\top W_q^\top W_k \mathbf{x}_j}_{\text{position-content}} + \underbrace{\mathbf{p}_i^\top W_q^\top W_k \mathbf{p}_j}_{\text{position-position}}

Each term tells us something different about why one token might attend to another:

  • Content-content: "Does this token's meaning relate to that token's meaning?" This is pure semantic matching, independent of where the tokens appear.
  • Content-position: "Given what this token is looking for, does that position matter?" For example, a verb might preferentially attend to its subject, which typically precedes it.
  • Position-content: "Given where this token is, does that token's content matter more?" Early positions might attend differently than late positions.
  • Position-position: "Do these two positions have an inherent affinity?" Adjacent positions might naturally attend to each other.

Why does this decomposition matter? Because it shows us exactly which parts of the attention score use absolute positions. Terms (b), (c), and (d) all depend on pi\mathbf{p}_i or pj\mathbf{p}_j (or both), the absolute position embeddings. These are the terms we need to fix. Term (a) is already position-free and can stay as is.

From Absolute to Relative

The problem with this decomposition is that all position information uses absolute positions pi\mathbf{p}_i and pj\mathbf{p}_j. Transformer-XL's insight is that we can rewrite these terms to use relative position instead. The redesign makes two key changes:

  1. Replace the key's absolute position with relative distance: Instead of pj\mathbf{p}_j (the absolute position of the key), use ri−j\mathbf{r}_{i-j} (the relative distance from query to key).

  2. Replace the query's position with a learned global bias: The query's absolute position pi\mathbf{p}_i becomes learned vectors u\mathbf{u} and v\mathbf{v} that don't depend on position at all.

The first change makes sense because the key's absolute position is what we want to replace. We want to say "this key is 5 positions before this query" rather than "this key is at absolute position 47." The second change is more subtle: the query's absolute position would still be ambiguous across segment boundaries even after fixing the key's encoding, so we replace it with a global learned bias that captures the query's general tendency to prefer nearby or distant keys, independent of where in the document the query appears.

The resulting formula is:

scoreij=xi⊤Wq⊤Wk,Exj⏟(a)+xi⊤Wq⊤Wk,Rri−j⏟(b)+u⊤Wk,Exj⏟(c)+v⊤Wk,Rri−j⏟(d)\text{score}_{ij} = \underbrace{\mathbf{x}_i^\top W_q^\top W_{k,E} \mathbf{x}_j}_{(a)} + \underbrace{\mathbf{x}_i^\top W_q^\top W_{k,R} \mathbf{r}_{i-j}}_{(b)} + \underbrace{\mathbf{u}^\top W_{k,E} \mathbf{x}_j}_{(c)} + \underbrace{\mathbf{v}^\top W_{k,R} \mathbf{r}_{i-j}}_{(d)}

where:

  • xi,xj\mathbf{x}_i, \mathbf{x}_j: token embeddings at positions ii and jj, unchanged from before
  • ri−j\mathbf{r}_{i-j}: a sinusoidal encoding of the relative distance i−ji - j, not the absolute position
  • Wk,EW_{k,E}: key projection matrix for content (the "E" stands for embeddings)
  • Wk,RW_{k,R}: key projection matrix for relative positions (the "R" stands for relative)
  • u\mathbf{u}: a learned global bias for content attention, shared across all query positions
  • v\mathbf{v}: a learned global bias for position attention, also shared across all positions

Why does this formula make sense? Notice that every term that previously depended on pi\mathbf{p}_i or pj\mathbf{p}_j now depends on either a relative distance encoding ri−j\mathbf{r}_{i-j} or a global learned constant. The relative distance i−ji-j is the same for any pair of tokens separated by the same gap, regardless of their absolute positions. This means the formula produces the same attention contribution for identical content pairs at identical relative distances, regardless of where in the document those pairs appear.

Let's unpack each component:

  1. Term (a): Pure content-based attention. This is unchanged from standard attention. The word "cat" attends to "feline" because of semantic similarity, regardless of position.

  2. Term (b): Content-dependent distance preference. The query content determines how much the model cares about distance. A pronoun might strongly prefer nearby tokens, while a discourse marker might look further back.

  3. Term (c): Global content importance. Some tokens are just important regardless of the query's position. The beginning-of-sentence token might receive attention from everywhere.

  4. Term (d): Global distance preference. The model learns a general preference for certain distances. Typically, nearby tokens receive more attention than distant ones.

Terms (b) and (d) now depend on i−ji - j rather than on ii and jj separately. This means the position signal is the same whether we're at the start of the document or the end, whether we're attending within the current segment or reaching back into cached memory.

Building the Relative Encoding

With the theory in place, let's implement relative position encoding step by step. We need two components: (1) a function to generate sinusoidal encodings for each possible relative distance, and (2) a function to compute attention scores using the four-term formula.

The sinusoidal encoding for relative positions works similarly to absolute position encodings, but instead of encoding absolute positions 0, 1, 2, ..., we encode relative distances ..., -2, -1, 0, 1, 2, .... Negative distances mean the key is before the query; positive distances mean the key is after the query (though in causal attention, we only see non-positive distances). The sinusoidal function ensures smooth, gradual variation between nearby distances, so encodings for similar distances are similar vectors.

In[16]:
Code
import numpy as np


def sinusoidal_relative_encoding(max_distance: int, d_model: int) -> np.ndarray:
    """
    Generate sinusoidal encodings for relative positions.

    Unlike absolute encodings, these encode the distance between positions,
    allowing the same encoding for any pair at the same distance.
    """
    positions = np.arange(-max_distance, max_distance + 1)
    encodings = np.zeros((len(positions), d_model))

    for i, pos in enumerate(positions):
        for j in range(0, d_model, 2):
            div_term = 10000 ** (j / d_model)
            encodings[i, j] = np.sin(pos / div_term)
            if j + 1 < d_model:
                encodings[i, j + 1] = np.cos(pos / div_term)

    return positions, encodings


def relative_attention_scores(
    queries: np.ndarray,
    keys: np.ndarray,
    relative_encodings: np.ndarray,
    rel_positions: np.ndarray,
    u: np.ndarray,
    v: np.ndarray,
    W_k_E: np.ndarray,
    W_k_R: np.ndarray,
) -> np.ndarray:
    """
    Compute Transformer-XL style attention scores with relative positions.

    Args:
        queries: Query vectors, shape (L, d)
        keys: Key vectors (from extended context), shape (L+M, d)
        relative_encodings: Sinusoidal encodings indexed by distance
        rel_positions: Array mapping distance to encoding index
        u, v: Global bias vectors, shape (d,)
        W_k_E, W_k_R: Key projections for content and position

    Returns:
        Attention scores, shape (L, L+M)
    """
    L = queries.shape[0]
    K = keys.shape[0]  # L + M (current + cached)
    M = K - L  # Memory/cache length

    # Term (a): content-to-content
    term_a = queries @ W_k_E @ keys.T  # (L, K)

    # Term (b): content-to-relative-position
    # For each query position i and key position j, we need r_{i-j}
    term_b = np.zeros((L, K))
    for i in range(L):
        for j in range(K):
            # Relative distance: query position (0 to L-1) vs key position (-M to L-1)
            # Key positions 0 to M-1 correspond to cached tokens at positions -M to -1
            # Key positions M to M+L-1 correspond to current tokens at positions 0 to L-1
            if j < M:
                key_pos = j - M  # Negative for cached tokens
            else:
                key_pos = j - M  # 0 to L-1 for current tokens
            query_pos = i
            distance = query_pos - key_pos

            # Find encoding for this distance
            enc_idx = np.where(rel_positions == distance)[0]
            if len(enc_idx) > 0:
                r_ij = relative_encodings[enc_idx[0]]
                term_b[i, j] = queries[i] @ W_k_R @ r_ij

    # Term (c): global content bias
    term_c = u @ W_k_E @ keys.T  # (K,) broadcast to (L, K)
    term_c = np.tile(term_c, (L, 1))

    # Term (d): global position bias
    term_d = np.zeros((L, K))
    for i in range(L):
        for j in range(K):
            if j < M:
                key_pos = j - M
            else:
                key_pos = j - M
            query_pos = i
            distance = query_pos - key_pos

            enc_idx = np.where(rel_positions == distance)[0]
            if len(enc_idx) > 0:
                r_ij = relative_encodings[enc_idx[0]]
                term_d[i, j] = v @ W_k_R @ r_ij

    return term_a + term_b + term_c + term_d
In[17]:
Code
# Demonstrate relative position encoding
rng = np.random.default_rng(42)

L = 4  # Current segment length
M = 4  # Cached segment length
d = 8  # Model dimension

# Generate relative position encodings
max_dist = L + M
rel_positions, rel_encodings = sinusoidal_relative_encoding(max_dist, d)

# Random queries and keys (in practice, these come from token embeddings)
queries = rng.standard_normal((L, d)) * 0.5
keys = rng.standard_normal((L + M, d)) * 0.5

# Learned parameters
u = rng.standard_normal(d) * 0.1
v = rng.standard_normal(d) * 0.1
W_k_E = np.eye(d) + rng.standard_normal((d, d)) * 0.1
W_k_R = np.eye(d) + rng.standard_normal((d, d)) * 0.1

scores = relative_attention_scores(
    queries, keys, rel_encodings, rel_positions, u, v, W_k_E, W_k_R
)
Out[18]:
Console
Attention score matrix shape: (4, 8)

Scores for query position 0 (first token of current segment):
  To cached positions (distances 4-7): [ 0.8931239  -0.1445392   1.11760926 -0.10404486]
  To current positions (distances 0-3): [-1.67922564 -0.77753701  0.42740559  0.55614617]

Scores for query position 3 (last token of current segment):
  To cached positions (distances 7-10): [-0.25324454  0.19148481  0.02378098 -0.03696258]
  To current positions (distances 0-3): [-0.56254582  0.31803436 -0.76636984 -0.96564079]

The scores vary based on both content similarity and relative distance. Notice that the scores to cached positions (which are further away) differ from scores to current positions (which are closer). The relative position encoding ensures consistent treatment of distance regardless of absolute segment boundaries. A token attending to something 3 positions back receives the same position signal whether that's within the current segment or reaching into the cached memory.

To better understand how these four terms contribute to the final attention score, let's decompose them for a single query-key pair across different relative distances:

Out[19]:
Visualization
Bar chart showing individual contributions of the four attention terms (content-to-content, content-to-position, bias-to-content, bias-to-position) at each relative distance.
Individual contributions of the four attention terms at each relative distance. Content-to-content (a) shows token-specific variation based on semantic similarity. The position-related terms (b, c, d) show systematic patterns tied to distance.
Stacked bar chart showing combined attention scores broken down by the four term contributions at each relative distance.
Combined attention score with stacked breakdown showing how content and position terms interact across relative distances. The peak at distance 0 reflects both high self-similarity and maximum position bias.

The visualization shows how content similarity (term a) provides the base signal, while position terms (b, c, d) modulate attention based on distance. The strong preference for position 0 (self-attention) reflects both high content similarity and the global position bias favoring nearby tokens.

Out[20]:
Visualization
Heatmap showing encoding similarity between relative positions from -8 to 8, with high similarity along the diagonal.
Relative position encoding similarity matrix showing how the sinusoidal encoding captures distance relationships. Nearby relative positions have highly similar encodings (warm colors along the diagonal), with similarity decaying smoothly as the distance between positions increases. This smooth decay allows the attention mechanism to naturally prefer nearby tokens while still accessing distant context when semantically relevant.

The similarity matrix shows that encodings for nearby relative positions are similar, gradually diverging for larger distances. This smooth decay allows the attention mechanism to naturally prefer nearby tokens while still accessing distant context when needed. The negative similarities at large distances reflect the oscillatory nature of sinusoidal functions, which ensures that encodings for distant positions become distinct vectors rather than faded versions of nearby ones.

Worked Example: Tracing Information Through Segments

To make the segment-level recurrence concrete, let's trace how information about a specific token propagates across three segments of a simple model. Imagine a 2-layer model with segment length 4 processing the sentence: "The professor lectured. Students took notes. They asked questions."

We have three segments: ["The", "professor", "lectured", "."], ["Students", "took", "notes", "."], and ["They", "asked", "questions", "."]. The token "They" in segment 3 is a pronoun that refers to "Students" in segment 2. Let's trace how the model can resolve this reference.

Segment 1 processing. No cached memory exists. Layer 1 processes ["The", "professor", "lectured", "."] with standard causal attention. The hidden states at layer 1 encode things like "lecturer introducing a topic." Layer 2 receives these layer-1 states as keys and values and produces layer-2 hidden states. These layer-2 states are cached as h12\mathbf{h}_1^2 (segment 1, layer 2 output) and h11\mathbf{h}_1^1 (segment 1, layer 1 output).

Segment 2 processing. Layer 1 of segment 2 receives the layer-1 cached states from segment 1: h10\mathbf{h}_1^0 (the embedding layer outputs). The cache at layer nn uses the previous segment's layer n−1n-1 output. Thus, layer 1 of segment 2 receives the embedding-level outputs of segment 1 as its extended key-value context. Layer 1 computes attention over ["Students" attended to "The", "professor", "lectured", "."] plus itself. The hidden state for "Students" at layer 1 now encodes "a plural subject, appearing after a lecturing context." Layer 2 of segment 2 receives layer-1 outputs of segment 1 as its extended context. The hidden state for "Students" at layer 2 now carries information about both the current segment and the previous one.

Segment 3 processing. Layer 1 of segment 3 receives the layer-1 outputs of segment 2 as its extended context. "They" at layer 1 can attend to the layer-1 representation of "Students" from segment 2. The layer-1 hidden state for "They" begins to encode a connection to the plural subject from the previous segment. Layer 2 of segment 3 receives the layer-2 outputs of segment 2 as its extended context. "They" at layer 2 now attends to the richer layer-2 representation of "Students," which already encoded the connection between "Students" and the lecturing context from segment 1. The final hidden state of "They" captures, through this chain of attention, a connection back to "Students" and even further to the "professor" from segment 1.

This trace illustrates several important properties. First, the coreference link between "They" and "Students" is resolved by a single attention operation at layer 1 of segment 3, which can directly attend to "Students" in the cached memory. Second, deeper connections (like the link between the students and the professor) are accessible through the layer-2 representations, which already aggregated information from segment 1 when processing segment 2. Third, the information path follows the recurrence chain: segment 3 layer 1 attends to segment 2 layer 0 cached states, and segment 3 layer 2 attends to segment 2 layer 1 cached states, which themselves incorporated segment 1 information.

The numerical details would require running an actual trained model, but the logical structure of the information flow is clear: information about tokens from prior segments is compressed into the hidden states and carried forward one segment at a time, with each layer having access to the corresponding layer's cached representation from the previous segment.

Implementing Transformer-XL

We've now covered the two core innovations of Transformer-XL: segment-level recurrence for extending context across segment boundaries, and relative position encoding for handling positions consistently across segments. Let's bring these pieces together into a complete implementation.

A Transformer-XL layer follows the same structure as a standard transformer layer: attention followed by feed-forward, with residual connections and layer normalization. The key differences are in the attention computation, where we must handle cached memory and compute relative position biases. Transformer-XL's original implementation uses "pre-norm" layering, applying layer normalization before the attention and feed-forward sublayers rather than after. This improves training stability for deep networks.

The implementation below simplifies some aspects for clarity, particularly multi-head attention (which runs this same computation in parallel across multiple heads with smaller dimensions) and the exact handling of the pre-norm versus post-norm choice. The core logic of concatenating cached states, computing extended keys and values, and applying the four-term attention score is faithfully represented.

In[21]:
Code
import numpy as np


class TransformerXLLayer:
    """
    A single Transformer-XL layer with segment-level recurrence
    and relative position encoding.
    """

    def __init__(
        self, d_model: int, n_heads: int, d_ff: int, max_rel_dist: int = 512
    ):
        self.d_model = d_model
        self.n_heads = n_heads
        self.d_head = d_model // n_heads
        self.max_rel_dist = max_rel_dist

        # Initialize parameters (simplified: using random initialization)
        rng = np.random.default_rng(42)
        scale = 0.02

        # Attention projections
        self.W_q = rng.standard_normal((d_model, d_model)) * scale
        self.W_k_E = rng.standard_normal((d_model, d_model)) * scale  # Content
        self.W_k_R = rng.standard_normal((d_model, d_model)) * scale  # Position
        self.W_v = rng.standard_normal((d_model, d_model)) * scale
        self.W_o = rng.standard_normal((d_model, d_model)) * scale

        # Global biases for relative attention
        self.u = rng.standard_normal(d_model) * scale
        self.v = rng.standard_normal(d_model) * scale

        # Feed-forward network
        self.W_ff1 = rng.standard_normal((d_model, d_ff)) * scale
        self.b_ff1 = np.zeros(d_ff)
        self.W_ff2 = rng.standard_normal((d_ff, d_model)) * scale
        self.b_ff2 = np.zeros(d_model)

        # Layer norms (simplified: just storing means and variances)
        self.ln1_gamma = np.ones(d_model)
        self.ln1_beta = np.zeros(d_model)
        self.ln2_gamma = np.ones(d_model)
        self.ln2_beta = np.zeros(d_model)

        # Precompute relative position encodings
        self.rel_positions, self.rel_encodings = sinusoidal_relative_encoding(
            max_rel_dist, d_model
        )

    def layer_norm(
        self, x: np.ndarray, gamma: np.ndarray, beta: np.ndarray
    ) -> np.ndarray:
        """Apply layer normalization."""
        mean = x.mean(axis=-1, keepdims=True)
        std = x.std(axis=-1, keepdims=True) + 1e-6
        return gamma * (x - mean) / std + beta

    def relative_attention(
        self,
        hidden: np.ndarray,
        memory: np.ndarray,
    ) -> np.ndarray:
        """
        Compute relative multi-head attention with cached memory.

        Args:
            hidden: Current segment hidden states, shape (L, d_model)
            memory: Cached previous segment states, shape (M, d_model)

        Returns:
            Attention output, shape (L, d_model)
        """
        L = hidden.shape[0]
        M = memory.shape[0] if memory is not None else 0

        # Concatenate memory and hidden for keys/values
        if memory is not None and M > 0:
            extended = np.concatenate([memory, hidden], axis=0)
        else:
            extended = hidden
        K = extended.shape[0]

        # Compute queries, keys, values
        queries = hidden @ self.W_q  # (L, d)
        keys_E = extended @ self.W_k_E  # (K, d)
        values = extended @ self.W_v  # (K, d)

        # Compute attention scores with relative positions
        # Term (a): content-to-content
        scores = queries @ keys_E.T  # (L, K)

        # Terms (b), (c), (d): relative position terms
        # Simplified: we add positional bias based on distance
        for i in range(L):
            for j in range(K):
                # Compute relative distance
                if j < M:
                    key_pos = j - M  # Negative for memory
                else:
                    key_pos = j - M  # 0 to L-1 for current
                query_pos = i
                distance = query_pos - key_pos

                # Find relative encoding
                idx = np.where(
                    self.rel_positions
                    == np.clip(distance, -self.max_rel_dist, self.max_rel_dist)
                )[0]
                if len(idx) > 0:
                    r_ij = self.rel_encodings[idx[0]]
                    # Add position bias terms
                    scores[i, j] += (queries[i] + self.u) @ self.W_k_R @ r_ij
                    scores[i, j] += self.v @ self.W_k_R @ r_ij

        # Scale
        scores = scores / np.sqrt(self.d_head)

        # Causal mask (only for current segment attending to past)
        mask = np.ones_like(scores) * float("-inf")
        for i in range(L):
            # Can attend to all memory + tokens up to and including position i
            mask[i, : M + i + 1] = 0
        scores = scores + mask

        # Softmax
        exp_scores = np.exp(scores - np.max(scores, axis=-1, keepdims=True))
        attention_weights = exp_scores / (
            exp_scores.sum(axis=-1, keepdims=True) + 1e-8
        )

        # Output
        output = attention_weights @ values  # (L, d)
        output = output @ self.W_o

        return output, attention_weights

    def feed_forward(self, x: np.ndarray) -> np.ndarray:
        """Apply position-wise feed-forward network."""
        hidden = np.maximum(0, x @ self.W_ff1 + self.b_ff1)  # ReLU
        return hidden @ self.W_ff2 + self.b_ff2

    def forward(
        self, hidden: np.ndarray, memory: np.ndarray
    ) -> tuple[np.ndarray, np.ndarray]:
        """
        Forward pass for one layer.

        Args:
            hidden: Current segment hidden states
            memory: Cached memory from previous segment

        Returns:
            output: New hidden states
            new_memory: States to cache for next segment
        """
        # Self-attention with memory
        attn_out, attn_weights = self.relative_attention(hidden, memory)
        hidden = self.layer_norm(
            hidden + attn_out, self.ln1_gamma, self.ln1_beta
        )

        # Feed-forward
        ff_out = self.feed_forward(hidden)
        output = self.layer_norm(hidden + ff_out, self.ln2_gamma, self.ln2_beta)

        return output, attn_weights
In[22]:
Code
# Demonstrate the layer
layer = TransformerXLLayer(d_model=64, n_heads=4, d_ff=256, max_rel_dist=32)

# Simulate processing two segments
segment_length = 8
hidden_dim = 64

# Segment 1: no memory yet
rng = np.random.default_rng(123)
segment1_input = rng.standard_normal((segment_length, hidden_dim)) * 0.5
segment1_output, attn1 = layer.forward(segment1_input, memory=None)

# Segment 2: use segment 1's output as memory
segment2_input = rng.standard_normal((segment_length, hidden_dim)) * 0.5
segment2_output, attn2 = layer.forward(segment2_input, memory=segment1_output)
Out[23]:
Console
Transformer-XL Layer Processing

Segment 1 (no memory):
  Input shape: (8, 64)
  Output shape: (8, 64)
  Attention shape: (8, 8)

Segment 2 (with cached memory from segment 1):
  Input shape: (8, 64)
  Memory shape: (8, 64)
  Output shape: (8, 64)
  Attention shape: (8, 16)
  Attention to memory: 5.324
  Attention to current: 2.676

The attention shape for segment 2 is (8, 16). This reflects 8 query positions attending to 16 key positions (8 cached + 8 current). The total attention sums show how much of the model's attention budget goes to memory versus the current segment. A roughly balanced split indicates the model is actively using both sources of context.

Out[24]:
Visualization
Heatmap of attention weights for segment 2 queries attending to segment 1 cached memory and current segment positions, showing a causal mask over the current segment.
Attention pattern when processing segment 2 with segment 1's cached states as memory. Current segment tokens (rows) can attend to all 8 cached memory positions (left half) plus their own causal context (right half, triangular pattern). Earlier tokens in the current segment devote more attention to memory since they have fewer local predecessors.
Out[25]:
Visualization
Stacked bar chart showing the proportion of attention allocated to cached memory versus the current segment for each query position, with early positions showing higher memory attention.
Attention distribution per position showing how much attention each query position allocates to memory versus current segment. Early positions rely more on memory due to limited local context, while later positions progressively shift attention toward current-segment tokens as more local predecessors become available.

The visualization reveals how attention is distributed between memory and the current segment. Early positions in the current segment allocate substantial attention to memory because they have limited local context. Later positions can attend more to the growing local context while still accessing memory for longer-range dependencies. This natural transition from memory-heavy to local-heavy attention across positions reflects the model's learned understanding that nearby context is usually more relevant but that distant context is invaluable when local cues are scarce.

Evaluation: Comparing Context Approaches

How does Transformer-XL's approach compare to other long-context methods? We can evaluate on a synthetic task that explicitly requires long-range dependencies: the copying task. The model must reproduce a sequence of tokens after a long delay filled with noise. This task is a clean test because there is no way to guess the target tokens; the only path to success is maintaining an accurate memory of the original sequence across the delay.

The copying task is deliberately extreme: real language modeling rarely requires verbatim recall of distant tokens. But it is a diagnostic probe that reveals the hard boundary of a model's effective memory. Any model that cannot solve the copying task at a given delay length cannot reliably use information from that far back in any task.

In[26]:
Code
def create_copying_task(
    seq_length: int, copy_length: int, delay_length: int, vocab_size: int = 10
) -> tuple[np.ndarray, np.ndarray]:
    """
    Create a copying task instance.

    The input consists of:
    1. A sequence of tokens to remember (copy_length tokens)
    2. A delay period filled with blanks (delay_length tokens)
    3. A signal token indicating the model should start reproducing

    The target is the original sequence of tokens.
    """
    rng = np.random.default_rng(42)

    # Tokens to copy (1 to vocab_size-2, reserve 0 for blank, vocab_size-1 for signal)
    to_copy = rng.integers(1, vocab_size - 1, size=copy_length)

    # Build input sequence
    blank_token = 0
    signal_token = vocab_size - 1

    input_seq = np.concatenate(
        [
            to_copy,
            np.full(delay_length, blank_token),
            [signal_token],
            np.full(
                copy_length - 1, blank_token
            ),  # Positions where model outputs
        ]
    )

    # Target: just the copied tokens at the end
    target = np.concatenate(
        [
            np.full(copy_length + delay_length + 1, -1),  # -1 = ignore
            to_copy[:-1],  # Predict each token in the copy
        ]
    )

    return input_seq, target, to_copy


# Create examples with different delay lengths
delays = [50, 100, 200, 400]
copy_len = 10
examples = [create_copying_task(500, copy_len, d) for d in delays]
Out[27]:
Console
Copying Task Examples

Task: Remember the first 10 tokens, output them after a delay

Delay length: 50
  Tokens to copy: [1 7 6 4 4 7 1 6 2 1]
  Input length: 70
  Required context: 60 positions

Delay length: 100
  Tokens to copy: [1 7 6 4 4 7 1 6 2 1]
  Input length: 120
  Required context: 110 positions

Delay length: 200
  Tokens to copy: [1 7 6 4 4 7 1 6 2 1]
  Input length: 220
  Required context: 210 positions

Delay length: 400
  Tokens to copy: [1 7 6 4 4 7 1 6 2 1]
  Input length: 420
  Required context: 410 positions

This copying task becomes impossible for models without sufficient context. If the segment length is 100 and the delay is 200, a standard transformer with context fragmentation can never succeed because the tokens to copy fall outside any segment that needs to reproduce them. Transformer-XL can succeed as long as the required context falls within its effective reach of N×LN \times L tokens.

Out[28]:
Visualization
Line plot comparing required context length against segment length threshold, showing failure region for standard transformers.
Required context length versus segment length threshold for the copying task. The solid line shows the minimum context needed to solve tasks with different delay lengths. Standard transformers fail when required context exceeds the segment length (shaded regions above each dashed horizontal line). Transformer-XL extends the failure threshold proportional to the number of layers, but information decay still imposes practical limits at very long delays.

The figure illustrates the fundamental limitation of fixed context windows. As delay length increases, the required context eventually exceeds any fixed segment length. Standard transformers fail in the shaded regions. Transformer-XL extends the failure threshold by a factor proportional to the number of layers, but the memory cache size and information decay still impose practical limits.

Limitations of Recurrent Memory

While Transformer-XL's segment-level recurrence significantly extends effective context, it comes with important limitations that shape when and how the technique should be applied. These limitations determine when recurrent memory is the right tool and when other approaches such as sparse attention or full long-context models would serve better.

Information decay over segments. Hidden states are finite-dimensional vectors. As information propagates through multiple segments, it inevitably compresses and degrades. A fact stated in segment 1 may be perfectly preserved in segment 2's hidden states, partially preserved in segment 3, and largely lost by segment 10. Unlike attention over the full sequence, recurrence cannot perfectly preserve arbitrary information over arbitrary distances.

The rate of decay depends on how distinctive the information is. A highly unusual word or a rare syntactic structure creates a strong signal in the hidden states that resists compression. Common words and typical sentence structures blend into the background representation more quickly. This means the model's effective memory is not uniform: it tends to remember remarkable events for longer and forget mundane details sooner. For language modeling, this selective decay can be useful because remarkable events are also more likely to be referenced later. For tasks requiring precise recall of any specific token regardless of distinctiveness, this selective decay can be problematic.

Out[29]:
Visualization
Line plot showing exponential decay of information signal strength across segments for three different initial signal levels.
Simulated information decay across segments in recurrent memory. Information about a specific token (introduced in segment 0) degrades as it propagates through hidden states. The rate of decay depends on the token''s distinctiveness: distinctive tokens (high initial signal) persist longer than common tokens (low initial signal). After about 8 segments, even distinctive information drops below the recovery threshold (dashed red line).

This decay means Transformer-XL works best for gradual, statistical dependencies rather than precise long-range retrieval. Language modeling benefits because most predictions depend on local context with only soft influence from distant text. Tasks requiring exact recall of distant tokens may still fail even with recurrence.

Unidirectional information flow. The recurrence mechanism flows strictly backward in time. Segment 5 can access information from segments 1-4, but segment 2 cannot access information from segment 5. This asymmetry limits bidirectional tasks. For language understanding tasks like question answering where the question appears after the context, the question cannot inform how the context is processed.

Some architectures address this with bidirectional memory or multiple passes, but these increase complexity and computation. XLNet, a successor to Transformer-XL, uses a permutation-based training objective to recover some bidirectional context while maintaining an autoregressive generation model. But the fundamental tradeoff between sequential efficiency and bidirectional context remains: you can have one cheaply but getting both requires additional machinery.

Memory cache management. Storing hidden states for memory consumes GPU memory proportional to a simple product:

Memory=M×d×Nlayers×B\text{Memory} = M \times d \times N_{\text{layers}} \times B

where:

  • MM: the memory/cache length (number of tokens cached from previous segment)
  • dd: the hidden dimension of the model
  • NlayersN_{\text{layers}}: the number of transformer layers (each layer maintains its own cache)
  • BB: the batch size

For large models, this becomes substantial. A 24-layer model with hidden dimension 1024 and memory length 512 requires storing over 12 million floats per sample. Batch processing multiplies this further.

In[30]:
Code
def compute_memory_requirements(
    memory_length: int,
    hidden_dim: int,
    num_layers: int,
    batch_size: int,
    bytes_per_float: int = 4,  # FP32
) -> dict:
    """
    Compute memory requirements for Transformer-XL cache.
    """
    floats_per_layer = memory_length * hidden_dim
    total_floats = floats_per_layer * num_layers * batch_size
    total_bytes = total_floats * bytes_per_float
    total_mb = total_bytes / (1024 * 1024)
    total_gb = total_mb / 1024

    return {
        "memory_length": memory_length,
        "hidden_dim": hidden_dim,
        "num_layers": num_layers,
        "batch_size": batch_size,
        "total_floats": total_floats,
        "memory_mb": total_mb,
        "memory_gb": total_gb,
    }


# Compare different configurations
configs = [
    (512, 768, 12, 8),  # BERT-base scale
    (512, 1024, 24, 8),  # BERT-large scale
    (1024, 1024, 24, 8),  # Larger memory
    (2048, 2048, 36, 4),  # GPT-2 XL scale
]

memory_stats = [compute_memory_requirements(*c) for c in configs]
Out[31]:
Console
Memory Cache Requirements for Transformer-XL

Config                             Floats     Memory (MB)  Memory (GB)
----------------------------------------------------------------------
512m × 768d × 12L × 8b         37,748,736           144.0         0.14
512m × 1024d × 24L × 8b       100,663,296           384.0         0.38
1024m × 1024d × 24L × 8b      201,326,592           768.0         0.75
2048m × 2048d × 36L × 4b      603,979,776          2304.0         2.25

The memory requirements grow quickly with model size. A GPT-2 XL scale configuration with 2048-token memory already requires over 2 GB just for the cache, not counting the model weights or activations. This overhead becomes a significant factor when deploying recurrent memory models on resource-constrained hardware. The memory requirement is not avoidable in the way that attention computation cost can sometimes be reduced through approximations; these hidden states must be stored in their full precision to serve as accurate context for the next segment.

Out[32]:
Visualization
Line plot showing linear scaling of memory cache size in GB as memory length increases from 0 to 1000 tokens, with fixed model parameters.
Memory cache requirements scale linearly with memory length for a fixed model configuration (d=1024, 24 layers, batch size 8). Doubling the memory length exactly doubles the cache size, as the formula is simply a product. At a 4096-token memory length, the cache alone consumes nearly 3 GB, before accounting for model weights or activations.
Out[33]:
Visualization
Bar chart comparing memory cache size in GB across Small, Base, Large, and XL model configurations, showing rapid growth with model scale.
Memory cache requirements across model configurations from Small to XL. As models scale up in hidden dimension, number of layers, and memory length simultaneously, the cache overhead grows rapidly. The XL configuration requires nearly 3 GB of cache memory per batch, illustrating why large recurrent memory models are demanding to serve even on high-end hardware.

Training and inference divergence. During training, the memory cache is populated from training data, maintaining realistic statistics. During inference on new text, the cache may be empty at the start, creating a "cold start" problem where initial segments lack the memory context the model learned to expect. Some implementations address this by processing a warmup prefix that isn't used for generation, but this adds latency. The cold start problem is especially pronounced for tasks where the very beginning of a document is critical, such as generating a response to a prompt that constitutes the first segment.

Gradient computation complexity. Although the cached hidden states don't receive gradients (the StopGrad operation), backpropagation still needs to flow through the attention computation over the extended sequence. This slightly increases training complexity compared to pure segment-independent processing, though it's far less expensive than full sequence backpropagation. In practice, the StopGrad decision also means the model cannot learn to produce hidden states that are optimally informative as memory: the gradient signal for improving the quality of what gets cached is indirect, coming only through the representations that the next segment learns to extract from the memory.

Limited to one segment of direct memory. In the standard Transformer-XL formulation, the model only caches one previous segment's hidden states at each layer. Information from earlier segments is only accessible through the layered recurrence chain, and even then only up to NN segments back for an NN-layer model. If you need to access information from 100 segments ago in a 12-layer model, you simply cannot, regardless of how much memory you have available. Some variants extend this by caching multiple previous segments, but this multiplies memory requirements and attention cost proportionally.

When to Use Recurrent Memory

Transformer-XL's approach shines in specific scenarios where its properties are well-matched to the task. Understanding these scenarios helps you decide when recurrent memory is the right choice and when other long-context approaches would serve better.

Recurrent memory works best in situations where information flows naturally in one direction and nearby context matters more than distant context. Language modeling on long documents is the paradigmatic use case: you are always generating the next token based on everything that came before, and the most recent context is almost always the most relevant. Streaming applications where text arrives incrementally have a similar structure: you process each new piece as it arrives, carrying forward a memory of what came before. Document summarization, long-form text generation, and code completion in large files all fit this profile.

Memory-constrained settings are another natural fit for recurrent memory. When the alternative would be either truncating context or paying the quadratic cost of full attention over long sequences, caching a fixed-size memory provides a middle ground. The memory cost is linear in sequence length rather than quadratic, and the cached states compress a segment's information into a fixed-size buffer rather than growing without bound.

Transformer-XL's approach is less suitable when the task requires looking at context that may be far in the past, or when the relevant information is unpredictably distributed across the document. Question answering over long documents is a prime example: the answer may appear anywhere in the document, and the question text, which appears at the end, cannot inform how the earlier document was processed. Bidirectional tasks like natural language inference, which require comparing two pieces of text that may not be adjacent, also do not map well to unidirectional recurrence. For these tasks, either full bidirectional attention or retrieval-based approaches tend to work better.

Modern alternatives like FlashAttention and sparse attention patterns have reduced the cost of longer context windows, somewhat diminishing the need for recurrent approaches. However, when truly long sequences must be processed incrementally, segment-level recurrence remains a powerful tool. The recurrent approach also has the appealing property of being constant-cost per token at inference time: no matter how long the document is, processing each new segment costs the same amount. This makes it particularly well-suited for streaming inference where latency per token is a constraint.

The key tradeoff to keep in mind when choosing recurrent memory is specificity versus scale. Recurrent memory can handle very long documents with modest computational cost, but it trades precision for scale. If your task requires precise access to information from many segments ago, recurrent memory will likely disappoint. If your task involves gradual accumulation of context where the recent past is most important and the distant past provides soft background influence, recurrent memory is an excellent choice.

Summary

This chapter explored Transformer-XL's recurrent memory mechanism, which extends effective context beyond the fixed segment length by caching and reusing hidden states from previous segments.

Context fragmentation, where fixed-length segment processing breaks cross-segment dependencies, fundamentally limits standard transformers. Transformer-XL addresses this by caching the hidden states from the previous segment and concatenating them with the current segment's keys and values. Queries still come only from the current segment, creating an asymmetric attention pattern that enables information flow from past to present while keeping computational cost constant per segment.

The recurrence mechanism requires relative position encodings because absolute positions would be ambiguous across segments. Transformer-XL redesigns the attention score computation to depend on the relative distance between query and key positions rather than their absolute locations. This involves decomposing the standard attention score into four terms and replacing the absolute-position-dependent terms with relative distance encodings and global learned biases. The result is an attention mechanism that behaves consistently regardless of where in the document a pair of tokens appears.

The effective context length grows with network depth. Information propagates one segment further back at each layer, so an NN-layer model has a theoretical reach of NN segments beyond the directly attended memory. In practice, information decay limits this reach, but the extension is still substantial, and the same depth that provides expressiveness also provides temporal reach.

Key implementation considerations include:

  • Memory cache storage scales with memory length, hidden dimension, number of layers, and batch size
  • The StopGrad operation prevents backpropagation through cached states, limiting training signal for long-range learning
  • Cold start at inference time may require warmup prefixes to populate the cache
  • Unidirectional information flow limits applicability to bidirectional tasks
  • Information decay across segments means recurrent memory works better for statistical dependencies than precise recall

Recurrent memory represents a principled approach to the long-context problem: accept that we cannot attend to everything at once, but ensure that information can flow across the boundaries we impose. While modern advances in attention efficiency have expanded what "at once" can mean, the core insight, that hidden states can carry forward context without explicit attention, remains valuable for streaming and memory-efficient processing of long sequences.

Key Parameters

When implementing Transformer-XL or similar recurrent memory mechanisms, the following parameters have the greatest impact on model behavior:

  • segment_length: The number of tokens processed in each forward pass. Larger segments capture more local context but increase memory usage quadratically (due to attention). Typical values range from 128 to 512 tokens.

  • memory_length: The number of tokens cached from the previous segment. Usually set equal to segment_length, but can be larger to extend context reach at the cost of increased memory and computation.

  • num_layers: Deeper networks extend effective context linearly. A 24-layer model can theoretically access 24x more context than a single layer, though information decay limits practical gains.

  • d_model: The hidden dimension affects both model capacity and memory requirements. Cache memory scales linearly with this parameter. Common values are 768 (BERT-base) to 1024 (GPT-2).

  • max_rel_dist: The maximum relative distance for position encodings. Should be at least segment_length + memory_length to cover all possible query-key distances. Setting this too small causes position information to saturate for distant tokens.

Quiz

Ready to test your understanding? Take this quick quiz to reinforce what you've learned about Transformer-XL and recurrent memory mechanisms.

Recurrent Memory Quiz

Question 1 of 100 of 10 completed
What is 'context fragmentation' in standard transformers?

Comments

No comments yet. Be the first to share your thoughts!

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2025recurrentmemory, author = {Michael Brenndoerfer}, title = {Recurrent Memory: Extending Transformer Context}, year = {2025}, url = {https://mbrenndoerfer.com/writing/recurrent-memory-transformer-xl-segment-recurrence}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-09-30} }
APAAcademic
Michael Brenndoerfer (2025). Recurrent Memory: Extending Transformer Context. Retrieved from https://mbrenndoerfer.com/writing/recurrent-memory-transformer-xl-segment-recurrence
MLAAcademic
Michael Brenndoerfer. "Recurrent Memory: Extending Transformer Context." 2026. Web. September 30, 2026. <https://mbrenndoerfer.com/writing/recurrent-memory-transformer-xl-segment-recurrence>.
CHICAGOAcademic
Michael Brenndoerfer. "Recurrent Memory: Extending Transformer Context." Accessed September 30, 2026. https://mbrenndoerfer.com/writing/recurrent-memory-transformer-xl-segment-recurrence.
HARVARDAcademic
Michael Brenndoerfer (2025) 'Recurrent Memory: Extending Transformer Context'. Available at: https://mbrenndoerfer.com/writing/recurrent-memory-transformer-xl-segment-recurrence (Accessed: September 30, 2026).
SimpleBasic
Michael Brenndoerfer (2025). Recurrent Memory: Extending Transformer Context. https://mbrenndoerfer.com/writing/recurrent-memory-transformer-xl-segment-recurrence

About the author

Continue with the full handbook

This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.

Explore Language AI Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.