Language Models2020intermediate10 min read
REALM: Retrieval-Augmented Language Model Pre-Training
REALM: التدريب المسبق لنموذج لغوي مُعزَّز بالاسترجاع
Guu, K. · Lee, K. · Tung, Z. · Pasupat, P. · Chang, M.-W. — ICML
The problem
By 2020, large language models like BERT and T5 stored world knowledge implicitly in their parameters. To know more facts, you had to build bigger networks — expensive and opaque. No one could tell what the model knew or where inside the weights that knowledge lived. And when facts changed, you had to retrain the entire model. The field needed a way to give language models access to external knowledge that was modular, interpretable, and updatable.
The contribution
REALM: a framework that augments language model pre- with a learned neural knowledge retriever. During pre-training, the model retrieves relevant Wikipedia passages to help predict masked tokens — and learns which passages are useful via through the retrieval step. The retriever uses (MIPS) over a cached document index, refreshed asynchronously. Fine-tuned on Open-domain QA, REALM outperforms all prior systems on NaturalQuestions, WebQuestions, and CuratedTrec by 4–16% absolute accuracy, including models 30× its size.
The impact
REALM established the paradigm of learned retrieval during pre-training — the idea that a language model should be trained end-to-end with a retriever from the start, not just at . It directly inspired RAG (Lewis et al., 2020), RETRO (Borgeaud et al., 2022), and the entire ecosystem that powers modern AI assistants. Its core insight — that retrieval is a you can optimize — remains the foundation of knowledge-grounded language models.
Imagine a student taking an open-book exam. A student who memorized everything (BERT) can answer from memory but struggles with rare facts and can't update what they know. A student with an encyclopedia but no idea where to look ( + reader) wastes time flipping pages.
REALM is the student who learned how to use the index: before each question, they flip to exactly the right page, read the relevant paragraph, and answer. The key insight? This student practiced using the index during study sessions, not just during the exam — so by test time, retrieval is second nature.
The problem: knowledge trapped in parameters
Language models like BERT capture world knowledge by training on massive text corpora. Ask BERT to fill in "The ___ is the currency of the United Kingdom" and it correctly predicts "pound." But this knowledge is stored implicitly inside millions of neural network parameters — and that creates three problems.
First, capacity is limited: to store more knowledge, you need a bigger model, which means more compute and more memory. Second, knowledge is opaque: you can't inspect which parameters encode which facts. Third, knowledge is frozen: when facts change (a new president, an updated record), you have to retrain the entire model.
Before REALM, retrieval-based QA systems existed but used fixed retrievers like BM25 — keyword matching that can't learn from context. The retriever and the reader were separate systems, never trained together. This meant the retriever couldn't adapt to what the reader actually needed.
The idea: retrieve, then predict — end-to-end
REALM's insight is elegant: model retrieval as a latent variable in the language model. Before predicting masked tokens, the model first retrieves a document from a knowledge corpus (like Wikipedia). The document is treated as a hidden variable — the model doesn't know in advance which document is useful, but it learns to find the right ones because helpful retrievals improve the prediction.
Think of it as a two-step generative process. Step 1: given a masked sentence, use a knowledge retriever to pick a document. Step 2: use a knowledge-augmented to read both the sentence and the document, then predict the masked word. The total probability of the answer is the sum over all possible documents, weighted by how likely each is to be retrieved.
The key breakthrough is that both the retriever and the encoder are trained jointly — the gradient flows from the prediction loss all the way back through the retrieval decision. A retrieval that helps predict the masked word is rewarded; an unhelpful retrieval is penalized. This is what makes REALM fundamentally different from prior systems that bolted a fixed retriever onto a reader.
The marginal likelihood: summing over all documents
REALM models the probability of an answer by treating the retrieved document as a latent variable and marginalizing over all possible documents. The model first selects a document from the corpus, then uses it alongside the query to produce an answer. The total probability of the output is the sum of contributions from every possible document.
The retriever models the probability of selecting each document using a softmax over relevance scores. These scores are computed as the inner product between a query and a document embedding, both produced by BERT encoders with linear projections.
Architecture: two BERTs and a search index
REALM has two main components, each built on a BERT encoder.
The Knowledge Retriever encodes the input query and each document separately. The query passes through a BERT encoder, and its is projected to a lower-dimensional embedding. Each document (title + body) is encoded the same way. Retrieval is the inner product of these two embeddings, and MIPS finds the top-k documents efficiently.
The Knowledge-Augmented Encoder takes the input and a retrieved document, concatenates them into one sequence [CLS] query [SEP] document [SEP], and feeds this into a second BERT. This allows rich cross- between the query and the evidence before making a prediction. For , the output at each [MASK] position is used to predict the original . For Open-QA, the model predicts start and end positions of the answer span within the document.
The engineering trick: asynchronous index refresh
Here's a challenge: if the retriever's parameters change every training step, the cached document embeddings go stale. Re-embedding millions of documents on every step is impractical. REALM solves this with an elegant two-job system.
A trainer job runs gradient updates on the model parameters as normal. In parallel, an index builder job takes a snapshot of the current retriever parameters, re-embeds all documents, rebuilds the MIPS index, and ships the new index back to the trainer. This cycle repeats every few hundred steps.
Between refreshes, the index is slightly stale — but the authors show empirically that the staleness is small enough that training remains stable. The key requirement is that refreshes happen frequently enough that the index doesn't drift too far from the current parameters.
Salient span masking: focusing on world knowledge
Standard BERT masks random tokens — but many masked tokens only require local syntax to predict (e.g., "the" or "of"). These don't teach the retriever anything about world knowledge. REALM's solution is salient span masking: instead of masking random tokens, it masks named entities and dates — spans that genuinely require world knowledge to predict.
A BERT-based NER tagger trained on CoNLL-2003 identifies named entities, and a regex catches dates. REALM selects and masks one of these salient spans per sentence. This focuses the learning signal on cases where retrieval can actually help, dramatically improving pre-training quality.
The idea in code
Simplified to show the idea — not the real implementation.
import numpy as np
def realm_forward(query, corpus_embeddings, encoder, retriever, top_k=5):
"""REALM's retrieve-then-predict in one forward pass."""
# Step 1: Embed the query
q_emb = retriever.embed_input(query) # (d,)
# Step 2: Score all documents via inner product (MIPS in practice)
scores = corpus_embeddings @ q_emb # (num_docs,)
# Step 3: Get top-k documents
top_idx = np.argsort(scores)[-top_k:]
top_scores = scores[top_idx]
# Step 4: Softmax over retrieved documents
retrieval_probs = softmax(top_scores) # (top_k,)
# Step 5: For each retrieved doc, predict masked token
# Then marginalize: p(y|x) = sum_z p(y|z,x) * p(z|x)
total_prob = 0
for i, doc_idx in enumerate(top_idx):
doc = corpus[doc_idx]
prediction_prob = encoder.predict(query, doc) # p(y|z,x)
total_prob += prediction_prob * retrieval_probs[i]
return total_prob
# Key insight: gradients flow through retrieval_probs back to the
# retriever — rewarding retrievals that improve prediction accuracy.
# The MIPS index is refreshed asynchronously every ~500 steps.Results: smaller model, bigger knowledge
REALM was evaluated on three benchmarks, and the results were striking.
On NaturalQuestions-Open, REALM achieved 40.4% exact match accuracy — compared to 33.3% for the previous best retrieval-based system (ORQA) and 36.6% for the massive T5-11B model (which has 30× more parameters). On WebQuestions, REALM scored 40.7% (vs 36.4% ORQA). On CuratedTrec, REALM hit 46.8% (vs 30.1% ORQA) — a massive 16.7-point improvement.
These results demonstrated that explicit retrieval with a learned retriever during pre-training can outperform brute-force memorization in much larger models. The model is also more interpretable: you can inspect which documents it retrieved, and more modular: you can update the knowledge corpus without retraining.
What REALM unlocked
2020
REALM
First system to pre-train a language model jointly with a learned neural retriever. Set state-of-the-art on three Open-QA benchmarks by 4–16%.
2020
RAG — Retrieval-Augmented Generation
Extended the retrieve-then-generate paradigm with a seq2seq generator (BART) instead of an extractive reader. Used DPR for retrieval and became the namesake of the entire RAG ecosystem.
2022
RETRO — Retrieval-Enhanced Transformer
Scaled retrieval-augmented pre-training to 7B parameters with chunked cross-attention. Retrieved from a 2-trillion token database. Showed that retrieval can substitute for 10× parameter scaling.
REALM's core contribution is not just a model — it's a principle: retrieval is a latent variable that can be optimized end-to-end alongside the language model. Before REALM, retrieval was a fixed preprocessing step. After REALM, it became a differentiable, learnable component. This paradigm shift underlies every modern retrieval-augmented system, from AI chatbots that cite their sources to enterprise search engines that understand context.
CitationGuu, Lee, Tung, Pasupat, Chang. REALM: Retrieval-Augmented Language Model Pre-Training. ICML, 2020.
Terms in this paper
- Retrieval-Augmented Generationالتوليد المعزَّز بالاسترجاع
- Masked Language Modeling (MLM)نمذجة اللغة المُقنَّعة (MLM)
- Maximum Inner Product Searchبحث الحاصل الداخلي الأقصى
- Dense Retrievalالاسترجاع الكثيف
- Knowledge-Intensiveكثيف المعرفة
- Latent Variableالمتغير الكامن
- Open-Domain Question Answeringالإجابة المفتوحة عن الأسئلة
- End-to-End Learningالتعلُّم من طرف إلى طرف
- Embedding Spaceفضاء التضمين
- Fine-Tuningالضبط الدقيق
- Pretrainingالتدريب المسبق
- Encoderالمُرمِّز
- BM25BM25