# How kNN-LM Combines Retrieval and Language Modeling: Implementation Guide

> Learn how kNN-LM combines retrieval and language modeling by interpolating next-token and k-nearest neighbor distributions. Implement this powerful technique with our guide.

- Repository: [labml.ai/annotated_deep_learning_paper_implementations](https://github.com/labmlai/annotated_deep_learning_paper_implementations)
- Tags: how-to-guide
- Published: 2026-03-04

---

**kNN-LM augments a standard transformer language model by interpolating its next-token distribution with a k-nearest neighbor distribution derived from stored context embeddings, enabling the model to memorize rare patterns while maintaining generalization capabilities.**

The kNN-LM architecture enhances autoregressive language modeling by integrating a retrieval component that queries stored context embeddings at inference time. This guide examines how kNN-LM combines retrieval and language modeling within the labmlai/annotated_deep_learning_paper_implementations repository, demonstrating the mechanism that balances parametric knowledge with non-parametric memory to improve perplexity and domain adaptation.

## The Three-Stage Retrieval Pipeline

The implementation follows a strict separation between training data collection, index construction, and inference-time retrieval. According to the source code in `labml_nn/transformers/knn/`, the system extracts hidden representations during training, organizes them into a fast-searchable structure, and retrieves nearest neighbors during inference to combine with the base model's predictions.

### Collecting Context Embeddings

During training, the feed-forward input of the last transformer layer, denoted as `f(c_i)`, is extracted and stored alongside the corresponding target token `w_i`. The `gather_keys()` function in [`labml_nn/transformers/knn/build_index.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/knn/build_index.py) (lines 85-90) writes these embeddings into memory-mapped NumPy arrays named `keys.npy` and `vals.npy`. This process creates a persistent datastore mapping contextual representations to their subsequent tokens, effectively capturing the model's internal states for later retrieval.

### Building the FAISS Index

The stored embeddings are organized into a scalable search structure using the `build_index()` function in [`build_index.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/build_index.py) (lines 13-26). This function initializes a `faiss.IndexIVFPQ` instance—an inverted file index with product quantization—that stores the vectors `f(c_i)` with their integer IDs. This compression technique enables efficient approximate nearest neighbor search across millions of stored contexts while maintaining retrieval accuracy.

### Retrieving Neighbors and Interpolating Distributions

At inference time, the current context embedding `f(c_t)` is retrieved from `conf.model.ff_input`. The `knn()` function in [`labml_nn/transformers/knn/eval_knn.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/knn/eval_knn.py) implements the retrieval and aggregation logic:

1.  **Neighbor Search**: `index.search()` returns the indices of the 10 nearest keys (`idx`) and their L2 distances (lines 36-39)
2.  **Logit Construction**: Retrieved token IDs (`vals_found`) are converted to logits by weighting each token with the cosine similarity (`dot_prod`) between the query and neighbor keys, scattering these scores into a zero-initialized tensor `logits_token` (lines 45-61)
3.  **Distribution Mixing**: The `validation_loss()` function (lines 104-107) combines the transformer logits (`res`) with k-NN logits (`res_knn`) using a scalar interpolation weight `knn_weight` (denoted as `c` in the paper), implementing the formula **k-NN-LM = (Transformer pₜ) ⨁ (k-NN pₙ)** where `⨁` represents convex combination.

## Practical Implementation Example

The following code demonstrates the complete workflow, from loading a trained model to generating predictions with retrieval augmentation:

```python

# 1️⃣ Load a previously trained transformer experiment

from labml_nn.transformers.knn.build_index import load_experiment
conf = load_experiment('4984b85c20bf11eb877a69c1a03717cd')
conf.model.eval()

```

```python

# 2️⃣ Build the FAISS index (run once after training)

from labml_nn.transformers.knn.build_index import gather_keys, build_index
gather_keys(conf)               # creates keys.npy / vals.npy

build_index(conf)                # trains and saves faiss.index

```

```python

# 3️⃣ Run inference with retrieval + interpolation

from labml_nn.transformers.knn.eval_knn import load_index, knn
import faiss, numpy as np, torch

# Load the pre‑built index

index, keys_store, vals_store = load_index(conf)

# Example query tensor (batch × seq × d_model)

queries = conf.model.ff_input            # shape = (B, T, d_model)

# Retrieve k‑NN logits (here n_tokens = vocab size)

knn_logits = knn(queries, index, keys_store, vals_store, conf.n_tokens)

# Combine with transformer logits (e.g., weight = 0.4 for k‑NN)

weight = 0.4
final_logits = weight * knn_logits + (1 - weight) * conf.model(queries)

```

## Key Source Files

The implementation spans four files in the `labml_nn/transformers/knn/` directory:

-   **[`build_index.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/build_index.py)**: Implements `gather_keys()` to extract `f(c_i)` embeddings and `build_index()` to construct the FAISS IVF-PQ index
-   **[`eval_knn.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/eval_knn.py)**: Contains `knn()` for neighbor retrieval and `validation_loss()` for the interpolation logic that merges distributions
-   **[`train_model.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/train_model.py)**: Defines the training configuration (`Configs`) and model architecture used during the key collection phase
-   **[`__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/__init__.py)**: Provides high-level documentation and mathematical formulation of the kNN-LM approach

## Summary

-   kNN-LM **combines retrieval and language modeling** by maintaining a datastore of context embeddings `f(c_i)` gathered from the transformer's last layer feed-forward inputs during training
-   The system uses **FAISS IndexIVFPQ** for efficient approximate nearest neighbor search, enabling scalable retrieval from millions of stored contexts
-   At inference, the model retrieves the **k-nearest neighbors** (default k=10) and computes a probability distribution based on their target tokens weighted by cosine similarity to the query
-   Final predictions use **convex interpolation** between the parametric transformer distribution and the non-parametric retrieval distribution, controlled by the `knn_weight` hyperparameter
-   This hybrid architecture achieves lower perplexity by memorizing rare patterns through retrieval while preserving the transformer's generalization capabilities

## Frequently Asked Questions

### What is the computational overhead of adding retrieval to language modeling?

The retrieval step requires querying the FAISS index, which introduces latency proportional to the number of inverted lists probed during search. However, the implementation uses product quantization to compress vectors and limits the search to a subset of clusters, making the overhead manageable compared to the transformer forward pass for moderately sized datastores up to several hundred million tokens.

### How does kNN-LM handle domain adaptation without fine-tuning?

kNN-LM excels at domain adaptation because the retrieval component can access training examples from the target domain stored in the non-parametric datastore. While the parametric transformer generalizes poorly to distribution shifts, the k-NN component retrieves relevant in-domain contexts, effectively adapting the model's output distribution without updating any base model parameters.

### What determines the optimal balance between transformer and retrieval predictions?

The scalar `knn_weight` (λ) controls the interpolation between the two distributions and is treated as a hyperparameter tuned on validation perplexity. Higher values emphasize the retrieval distribution for rare or memorizable n-grams, while lower values rely on the transformer's parametric knowledge for common linguistic patterns.

### Can the datastore be updated after the initial training phase?

Yes, the datastore consisting of `keys.npy` and `vals.npy` can be appended with new context-token pairs, and the FAISS index can be rebuilt or incrementally updated. This allows the model to incorporate new factual knowledge or domain-specific text without retraining the base transformer, though the current [`build_index.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/build_index.py) implementation focuses on static datastore construction from a fixed training corpus.