Language Models2022advanced15 min read

RETRO: Improving Language Models by Retrieving from Trillions of Tokens

RETRO: تحسين النماذج اللغوية بالاسترجاع من تريليونات الرموز

Borgeaud, S. · Mensch, A. · Hoffmann, J. · Cai, T. · Rutherford, E. · Millican, K. · Rae, J. W. · Elsen, E. · Sifre, L. — ICML

The problem

By 2021, language models had reached hundreds of billions of parameters, yet scaling kept coupling two things together: more computation and more memorization. Every fact the model knew had to be baked into its weights, making training enormously expensive and updates nearly impossible. If new knowledge emerged or errors needed fixing, you had to retrain from scratch. Worse, there was no way to inspect what the model "remembered" or to remove problematic data after training.

The contribution

RETRO (Retrieval-Enhanced Transformer): a semi-parametric language model that augments a standard Transformer decoder with a frozen BERT retriever and a chunked mechanism. The input is split into 64-token chunks; for each chunk, the model retrieves the k-nearest neighbours from a 2-trillion-token database and integrates them via a dedicated encoder and cross-attention layers inserted every third block. RETRO 7.5B matches GPT-3 (175B) and Jurassic-1 (178B) on the Pile benchmark, using 25× fewer parameters. Pre-trained baselines can also be rapidly RETROfitted with retrieval in just 3% of the original training tokens.

The impact

RETRO demonstrated that explicit memory can substitute for raw parameter scaling, opening a new axis for language model improvement. It showed that retrieval gains remain constant as model size grows, meaning retrieval and parameters are complementary, not redundant. RETRO directly influenced ATLAS, which combined retrieval with few-shot learning, and inspired the broader wave of (RAG) systems now standard in production LLM deployments. The idea that you can update a model's knowledge by swapping its database — without retraining — became a foundational principle in modern AI infrastructure.

Think of a traditional language model as a brilliant but forgetful chef who has memorized thousands of recipes by heart. To learn a new dish, the chef must go back to cooking school for months. Now imagine a different chef who keeps a massive recipe library in the kitchen. Before cooking each course, this chef glances at the most relevant pages — not the whole book, just the right paragraphs. This chef doesn't need to memorize everything; the library is always there.

RETRO is that second chef. Instead of cramming 175 billion parameters' worth of knowledge into its brain, it keeps a 2-trillion-token external database and retrieves the right chunks just in time.

The problem: scaling forces memorization into weights

By 2021, the dominant recipe for improving language models was simple: make them bigger. GPT-3 used 175 billion parameters, Jurassic-1 used 178 billion. Each parameter increase brought better performance, but also higher training costs, longer training times, and a model whose knowledge was frozen the moment training stopped.

This approach couples two fundamentally different capabilities into a single mechanism. The model must simultaneously learn how to reason about language (syntax, semantics, pragmatics) and what facts exist in the world (dates, names, relationships, code patterns). Both are encoded in the same weight matrices — and there is no way to update one without touching the other.

Prior retrieval-augmented approaches like REALM and RAG showed that external retrieval could help, but they were limited to small models (under 300M parameters) and small databases (up to a few billion tokens). No one had demonstrated that retrieval could scale alongside the largest parametric models — or that it could remain effective when both the model and the database grow.

Open in Lab
Drag the model-size slider to see how RETRO's retrieval gain stays constant while parameters grow — equivalent to a ~10× parameter boost.
The demo wakes as you arrive…

The idea: retrieve, then read

RETRO's core insight is deceptively simple: split the input into chunks, retrieve similar text for each chunk, and let the model attend to both its own tokens and the retrieved passages. Instead of a model that stores everything internally, RETRO creates a model that looks things up — a semi-parametric approach where some knowledge lives in the weights and some lives in a searchable database.

Here is the pipeline step by step:

  1. Chunk the input. Split every sequence into fixed-size chunks of 64 tokens each. A 2048-token sequence becomes 32 chunks.
  2. Retrieve neighbours. For each chunk, use a frozen BERT model to compute an , then find the k-nearest neighbours in a precomputed database using approximate nearest-neighbour search (SCaNN). Each retrieved neighbour includes the matching chunk plus its continuation (the next 64 tokens), giving 128 tokens of context per neighbour.
  3. Encode the neighbours. Feed the retrieved chunks through a small bidirectional Transformer encoder that is conditioned on the current chunk's activations — so the representation of the retrieved text is modulated by what the model is currently processing.
  4. Integrate via chunked cross-attention. At specific decoder layers (every 3rd layer starting from layer 6), insert a chunked cross-attention (CCA) layer where the model's hidden states attend to the encoded neighbours. Crucially, the attending window is shifted: the last token of chunk u and all tokens of chunk u+1 attend to the neighbours retrieved for chunk u. This preserves autoregressive causality.
  5. Predict. The final token probabilities depend on both the model's own parameters and the retrieved context.
Open in Lab
Follow a sequence through RETRO's five-stage pipeline: chunking → retrieval → encoding → chunked cross-attention → prediction.
The demo wakes as you arrive…

Chunked cross-attention: the key mechanism

The most novel architectural component in RETRO is chunked cross-attention (CCA). Standard cross-attention would have every token attend to all retrieved neighbours across all chunks — quadratic cost that destroys scalability. CCA restricts each token to attending only the neighbours of its own chunk, keeping the cost linear in the number of retrieved tokens.

But there is a subtlety involving causality. When RETRO retrieves neighbours for chunk C_u, those neighbours are found using C_u itself. If tokens within C_u could attend to those neighbours, the model would be peeking at information derived from future tokens in the same chunk — breaking autoregressive causality. The solution: shift the attending window. Neighbours retrieved for chunk C_u are only accessible starting from the last token of C_u and extending through all of chunk C(u+1). This one-token-overlap trick maintains strict left-to-right generation while still allowing retrieved context to influence the next chunk's predictions.

The RETRO decoder interleaves two block types. A standard LM block applies self-attention and a feed-forward layer: LM(H) = FFW(ATTN(H)). A RETRO block inserts CCA between the self-attention and feed-forward layers: RETRO(H, E) = FFW(CCA(ATTN(H), E)), where E is the encoded neighbour set.

Open in Lab
See how the attending window shifts: click any token to see which retrieved neighbours it can access, and why the boundary preserves causality.
The demo wakes as you arrive…
RETRO(H,E)=FFW ⁣(CCA ⁣(ATTN(H), E))\text{RETRO}(H, E) = \text{FFW}\!\bigl(\text{CCA}\!\bigl(\text{ATTN}(H),\, E\bigr)\bigr)
RETRO block — self-attention, then chunked cross-attention, then feed-forward — H is the sequence of hidden states; E is the set of encoded retrieval neighbours. ATTN is causal self-attention, CCA is chunked cross-attention (attending only to the neighbours of each chunk), and FFW is the standard feed-forward network. A standard LM block omits CCA: LM(H) = FFW(ATTN(H)).
CCA(H,E)um+i−1=CA(hum+i−1,  Eu)for i∈[1,m]\text{CCA}(H, E)_{um+i-1} = \text{CA}(h_{um+i-1},\; E_u) \quad \text{for } i \in [1, m]
Chunked cross-attention — each attending token cross-attends to one chunk's neighbours — For each chunk u and each position i within the attending window, the cross-attention operates over the encoded neighbours E_u of chunk C_u. The attending window H⁺_u is shifted by one token: it contains the last token of chunk C_u and the first m-1 tokens of chunk C(u+1). CA is standard multi-head cross-attention with relative positional encodings.

The database: a trillion-token key-value store

RETRO's retrieval database is built from MassiveText, a multilingual dataset totalling over 5 trillion tokens drawn from the web, books, news, Wikipedia, and GitHub. The database is structured as a key-value memory: each key is the average BERT embedding of a 64-token chunk, and each value is the chunk itself concatenated with its 64-token continuation — giving 128 tokens per entry.

At training time, RETRO retrieves from a 600-billion-token subset. At evaluation, the full 1.75-trillion-token database is used. Retrieval uses SCaNN (Scalable Nearest Neighbour) for approximate k-nearest-neighbour search in L₂ distance over BERT embeddings, with query time of ~10ms per chunk. To prevent trivial cheating, neighbours from the same document as the training sequence are filtered out.

A critical comparison with token-level retrieval: kNN-LM stores one 1024-float vector per token, requiring 15 terabytes just for Wikipedia's 4 billion tokens. RETRO stores one embedding per 64-token chunk, needing only 215 gigabytes for the same Wikipedia data — a 70× compression. This chunk-level approach is what makes trillion-token retrieval feasible.

Open in Lab
Compare token-level retrieval (kNN-LM) vs chunk-level retrieval (RETRO). Toggle the database size to see storage requirements.
The demo wakes as you arrive…

Architecture: encoder, decoder, and the RETRO block

RETRO's architecture has two main components: a retrieval encoder and a decoder that interleaves standard Transformer blocks with RETRO blocks.

The retrieval encoder is a small bidirectional Transformer (2 layers, d=896) shared across all neighbours and all chunks. It takes each retrieved neighbour as input and conditions it on the decoder's intermediate activations for the corresponding chunk via cross-attention. This conditioning is critical: it allows the encoder's representation of retrieved text to be modulated by what the model currently needs, rather than being a static embedding.

The decoder follows a standard autoregressive Transformer design with RMSNorm (instead of LayerNorm) and relative positional encodings. RETRO blocks are inserted every 3 layers starting from layer 6. For the 7.5B model, this adds roughly 8% more parameters compared to the baseline (7.0B → 7.5B). All four model sizes share the same encoder dimensions, so the relative overhead shrinks as the decoder grows.

Open in Lab
Click on any component to see its role and parameter count.
The demo wakes as you arrive…

The idea in code

Chunked cross-attention — simplified implementationpython

Simplified to show the idea — not the real implementation.

import numpy as np

def chunked_cross_attention(H, E, m=64):
    """
    H: (n, d) decoder hidden states for a full sequence
    E: (l, k, r, d) encoded neighbours — l chunks, k neighbours, r tokens each
    m: chunk size (64 tokens)
    Returns: (n, d) output with retrieval information integrated
    """
    n, d = H.shape
    l = n // m
    output = H.copy()

    for u in range(l):
        # Attending window: last token of chunk u + first (m-1) tokens of chunk u+1
        start = u * m + (m - 1)        # last token of chunk C_u
        end = min(start + m, n)         # through chunk C(u+1)

        # Flatten all k neighbours' r tokens into one key-value set
        E_u = E[u].reshape(-1, d)       # (k*r, d)

        for i in range(start, end):
            q = H[i]                    # (d,) query = current hidden state
            scores = E_u @ q            # (k*r,) dot-product attention
            weights = softmax(scores)
            output[i] = H[i] + weights @ E_u  # residual connection

    return output

# Key insight: chunk C_u's neighbours only affect
# the last token of C_u and all of C(u+1).
# This preserves autoregressive causality while
# allowing retrieved context to shape the next chunk.

Results: matching giants with a fraction of the parameters

RETRO 7.5B achieved remarkable results across multiple benchmarks. On the Pile, it outperformed the 178B Jurassic-1 model on a majority of test subsets despite being 25× smaller. On Wikitext103, RETRO achieved a of 3.92 when retrieving from the full MassiveText database. On LAMBADA, accuracy scaled from 52% (172M) to 73% (7.5B) with retrieval enabled.

Three scaling dimensions were explored. First, model scaling: the retrieval gain remained constant as models grew from 172M to 7.5B parameters — retrieval adds a fixed ~10× effective parameter multiplier regardless of model size. Second, database scaling: expanding the retrieval database from Wikipedia (4B tokens) to MassiveText (1.75T tokens) yielded dramatic improvements. Third, neighbour scaling: performance improved with up to 10 neighbours for small models and up to 40 for the 7.5B model, after which quality degradation offset the benefit.

On (Natural Questions), RETRO 7.5B achieved 45.5 exact match — competitive with REALM (40.4) and RAG (44.5), though below FiD (51.4), likely because the T5 in FiD relies more heavily on the retrieved passages than RETRO's decoder-only architecture.

Open in Lab
Compare RETRO 7.5B against GPT-3 (175B), Jurassic-1 (178B), and Gopher (280B) across Pile subsets. Toggle subsets to see where retrieval helps most.
The demo wakes as you arrive…

RETROfitting: adding retrieval to any pre-trained model

One of RETRO's most practical contributions is RETROfitting — converting an existing pre-trained Transformer into a retrieval-enhanced model without retraining the original weights. The procedure freezes all pre-trained parameters and trains only the new CCA layers and retrieval encoder weights (less than 10% of total parameters for the 7B model).

RETROfitting requires only about 3% of the original training tokens (roughly 6 million sequences). Performance quickly surpasses the baseline and approaches that of a RETRO model trained from scratch. Critically, because the original weights are frozen, disabling retrieval at inference time recovers exactly the original model's performance — there is zero degradation.

This has important practical implications: any organization with a pre-trained Transformer can add retrieval capabilities cheaply, and can fall back to the original model at any time. The retrieval database itself can be swapped, updated, or filtered without touching model weights.

Dataset leakage: how much of the gain is real?

A natural concern: does RETRO simply copy-paste from its database? If evaluation data is present in the training set, a retrieval model could exploit that overlap more easily than a standard model — the retrieved neighbours might contain the exact answer.

The authors addressed this rigorously. They defined a filtered bits-per-byte metric: for each evaluation chunk, they measured its overlap with the nearest training chunk (longest common substring), then computed loss only on chunks below a given overlap threshold α. At α = 12.5% (less than 8 contiguous tokens shared), RETRO still outperformed baseline models across all sizes. The slope of the performance curve was steeper for RETRO — meaning it does exploit leakage more — but even on completely novel chunks, retrieval provided a meaningful improvement.

Additionally, the authors created a "future" Wikipedia dataset from articles written in September 2021 — months after the training data was collected. On this guaranteed-unleaked data, RETRO still showed consistent gains across all model sizes, confirming that the benefit comes from genuine generalization, not mere memorization.

Privacy, safety, and updatability

RETRO's semi-parametric design opens both risks and opportunities for AI safety. On the risk side, having direct access to training data at inference time exacerbates privacy concerns — the model can potentially surface memorized personal information from the database. On the opportunity side, retrieval offers a path to mitigation that parametric models lack.

First, data removal: offensive, biased, or private content can be removed from the retrieval database retroactively, without retraining. This is impossible with standard models where such content is entangled in billions of weight values. Second, updatability: to keep the model current, you update the database — orders of magnitude cheaper than retraining from scratch. Third, interpretability: you can inspect the neighbours that influenced any given prediction, making the model's reasoning more transparent than a purely parametric system. Fourth, differential privacy: because retrieval is separate from model weights, differential privacy techniques could guarantee that no private information is stored in the weights, while allowing controlled access through the database at inference.

The road to retrieval-augmented language models

  1. 2017

    Continuous Cache

    Added probability mass to tokens with similar previous activations, extending an LSTM's context via a token-level retrieval cache.

  2. 2020

    kNN-LM

    Extended continuous cache to Transformers with a Wikipedia-scale database. Token-level retrieval interpolated with model probabilities — no retraining needed, but storage costs explode at scale.

  3. 2020

    REALM

    First end-to-end trained retrieval for language model pre-training. Updated retriever and database during training — powerful but expensive and limited to small models.

  4. 2020

    RAG

    Retrieval-Augmented Generation: prepended retrieved passages to the prompt and trained with DPR retriever. Set state of the art on knowledge-intensive QA, but retrieval was one-shot per query — no chunked integration.

  5. 2022

    RETRO

    Scaled retrieval to trillions of tokens with chunk-level retrieval, frozen BERT embeddings, and chunked cross-attention. First to show retrieval gains persist at 7B+ scale. 7.5B matches 175B+ models.

  6. 2022

    ATLAS

    Combined retrieval-augmented pre-training with few-shot learning. Built on RETRO's insights but used end-to-end retriever training at moderate scale, achieving state-of-the-art few-shot results on knowledge-intensive tasks.

Design choices: what matters and what doesn't

The paper includes a thorough ablation study on a 247M model. Key findings: neighbour continuations matter more than the neighbours themselves (56% vs 22% of total retrieval gain), suggesting the model learns to anticipate what comes next rather than just copy. Training with 2 neighbours is computationally optimal — 1 neighbour hurts significantly, while 4 neighbours adds overhead with negligible gain. CCA every 3 layers from layer 6 outperforms applying cross-attention only at a single layer (top, middle, or bottom) and is far cheaper than applying it at every layer. Query conditioning the encoder on decoder activations and using relative positional encodings in CCA both provide pure improvements with no computational penalty. A deeper encoder (6 layers instead of 2) yields only 0.15% lower loss at 20% more training time — not worthwhile.

CitationBorgeaud, Mensch, Hoffmann, Cai, Rutherford, Millican, Rae, Elsen, Sifre et al.. Improving Language Models by Retrieving from Trillions of Tokens. ICML, 2022.

Terms in this paper