ScaNN (Scalable Nearest Neighbors) is Google Research's library for maximum inner product search (MIPS): given a query q and millions of vectors x, return the k largest values of 〈q, x〉. Its best-known idea is anisotropic vector quantization. When you compress vectors into codes, you should not minimise plain reconstruction error. You should weight the error component parallel to each datapoint more heavily, because that component is what shifts the scores of the queries that would actually retrieve it.
The derivation, with a two-dimensional intuition, is on the companion page ScaNN, the math. This page is the engineering side. You will train an anisotropic codebook yourself in NumPy, work out what the threshold parameter does to the weights, configure the real library with parameters checked against its docs, estimate per-query cost by hand, measure recall properly, and operate the index in production.
The loss in one paragraph
Write the quantization residual r = x − x̃ as r∥, its projection onto x, plus r⊥, the rest. Guo et al. (ICML 2020) show that if you weight queries by how much they care about x (for instance only queries with 〈q, x〉 ≥ T), the expected score error becomes h∥·‖r∥‖² + h⊥·‖r⊥‖², with h∥ ≥ h⊥. Only the ratio η = h∥/h⊥ matters for training. For the threshold weight, the paper's Theorem 3.4 gives the large-dimension limit η/(d−1) → (T/‖x‖)² / (1 − (T/‖x‖)²).
For unit-norm data, that means η ≈ (d−1)·T²/(1−T²). Here is what that does for a few values of T:
| d | T | η (large-d approximation) | Meaning |
|---|---|---|---|
| 128 | 0.1 | 1.28 | almost isotropic: close to plain k-means |
| 128 | 0.2 | 5.29 | parallel error costs about 5x orthogonal |
| 128 | 0.3 | 12.56 | strongly anisotropic |
| 768 | 0.2 | 31.96 | same T, much stronger weighting in high d |
The last row matters most for practice. A random unit query has inner products around 1/√d with any given point, so a fixed T is a much stricter filter in 768 dimensions than in 128. Do not copy a threshold across embedding models of different widths without re-tuning it.
Three stages of a search
Partitioning clusters the dataset into num_leaves leaves. A query scores every centroid and descends into the best num_leaves_to_search. Asymmetric hashing (AH) splits each vector into blocks of dimensions_per_block dimensions and replaces each block with the index of its nearest of 16 learned centers, which is 4 bits. At query time, ScaNN builds a lookup table of q·center per block, and a point's approximate score is the sum of its table entries. Those 16-entry tables fit in SIMD registers, which is where ScaNN's speed comes from. Reordering rescores the best few hundred candidates with exact dot products, which repairs most of the ranking damage that quantization caused.
Training an anisotropic codebook yourself
Training alternates the same two steps as k-means. The assignment step uses the anisotropic loss instead of squared distance. The update step has a closed form (Theorem 4.2 in the paper). Each center solves a small linear system that pulls it along the directions of its members, instead of taking their plain mean. This version fixes h⊥ = 1, which is exact for unit-norm data:
import numpy as np
def eta_for(d, T, norms):
r = (T / norms) ** 2
assert np.all(r < 1), "T must be below every datapoint norm"
return (d - 1) * r / (1 - r) # large-d limit of h_par / h_perp
def train_avq(X, k, T, iters=10, seed=0):
n, d = X.shape
rng = np.random.default_rng(seed)
C = X[rng.choice(n, k, replace=False)].copy()
norms = np.linalg.norm(X, axis=1)
eta = eta_for(d, T, norms)
U = X / norms[:, None] # unit directions of the datapoints
for _ in range(iters):
R = X[:, None, :] - C[None, :, :] # n x k x d (demo sizes only)
par = np.einsum("nkd,nd->nk", R, U) # parallel residual length
loss = (eta[:, None] - 1) * par**2 + np.einsum("nkd,nkd->nk", R, R)
assign = loss.argmin(axis=1)
for j in range(k): # closed-form center update
m = assign == j
if not m.any():
continue
Uj, ej = U[m], eta[m]
A = m.sum() * np.eye(d) + (Uj * (ej - 1)[:, None]).T @ Uj
C[j] = np.linalg.solve(A, X[m].sum(axis=0))
return C, assignThe loss line is the whole idea: η·r∥² + r⊥² rewritten as (η−1)·r∥² + ‖r‖². With η = 1 the update reduces to the cluster mean, which is ordinary k-means, a useful sanity test for your implementation. To see the effect, train two codebooks on the same normalised sample, one with T = 0.05 (η near 1) and one with T = 0.2. Score held-out queries using only the quantized vectors, and compare top-10 recall against brute force. Expect the gap to be largest in the high-recall regime the paper targets.
The real library applies this per block, as product quantization, and that couples the blocks. The parallel component depends on the whole vector, so changing one block's code changes the best choice in the others. ScaNN handles this by refining the assignments iteratively rather than picking each block independently. Keep that in mind if you port the toy above to PQ.
Using the library
Install with pip install scann. Supported targets are Linux with Python 3.9–3.13. Since version 1.4.0, the TensorFlow integration needs the scann[tf] extra. The builder below follows the project's own example notebook:
import numpy as np, scann
X = X / np.linalg.norm(X, axis=1, keepdims=True) # cosine via dot_product
searcher = (
scann.scann_ops_pybind.builder(X, 10, "dot_product")
.tree(num_leaves=1000, num_leaves_to_search=100, training_sample_size=250000)
.score_ah(2, anisotropic_quantization_threshold=0.2)
.reorder(100)
.build()
)
neighbors, scores = searcher.search_batched(Q) # defaults from the builder
neighbors, scores = searcher.search_batched(Q, leaves_to_search=150,
pre_reorder_num_neighbors=250)
nbr, sc = searcher.search(Q[0], final_num_neighbors=5) # single query
searcher.serialize("/srv/index/v42") # directory of artifacts
searcher = scann.scann_ops_pybind.load_searcher("/srv/index/v42")| Knob | Rule of thumb from the docs | Effect of raising it |
|---|---|---|
| Brute force vs tree | under ~20k points, use score_brute_force | — |
num_leaves | roughly √n | smaller leaves, more centroid work |
num_leaves_to_search | tune to your recall target | higher recall, lower QPS (the main dial) |
dimensions_per_block | 2 for AH | smaller codes, more error |
reorder(n) | n greater than k | better ranking, more exact dot products |
The search-time overrides leaves_to_search, pre_reorder_num_neighbors and final_num_neighbors let you sweep recall against latency without rebuilding. search_batched_parallel splits large query batches across threads.
Worked example: cost of one query
You can estimate cost before you benchmark. Take n = 1,000,000 normalised float32 embeddings with d = 128 and the configuration above.
- Memory. Raw vectors take 1M × 128 × 4 B = 512 MB. With 2-dimension blocks there are 64 blocks of 4 bits each, so 32 B per point and 32 MB of codes. Reordering needs the original vectors as well, so budget for both, or reorder against a lower-precision copy and accept slightly worse final ranking.
- Partition stage. 1,000 centroids × 128 dims is 128K multiply-adds.
- LUT build. 64 blocks × 16 centers × 2 dims is about 2K multiply-adds. Negligible.
- AH scan. 100 of 1,000 leaves is about 10% of the data, roughly 100K points × 64 lookups = 6.4M table lookups. This dominates.
- Reorder. 100 × 128 = 12.8K multiply-adds.
The scan is about fifty times the next-largest term, so query cost is roughly linear in the fraction of leaves searched. That is why num_leaves_to_search is the dial you tune, and why better partitions, which find the true neighbours in fewer leaves, pay off more than anything else. Treat these as operation counts, not latencies, and confirm them with a benchmark on your own hardware.
The same arithmetic tells you how to scale. Doubling n at a fixed leaf fraction doubles the scan, so either accept the extra cost or re-tune: with num_leaves at √n, each leaf grows only by √2, and you can often search a slightly smaller fraction for the same recall. Batching queries does not shrink the per-query lookup count, but it improves cache reuse of leaf data and amortises Python call overhead, which is why search_batched and search_batched_parallel are the right entry points for offline jobs and high-QPS services alike.
Measuring recall honestly
Always measure recall against exact ground truth computed on the same normalised vectors. Sweep the search-time knobs and keep the cheapest setting that meets your target:
import time
def compute_recall(neighbors, true_neighbors): # as in the ScaNN example notebook
total = 0
for gt_row, row in zip(true_neighbors, neighbors):
total += np.intersect1d(gt_row, row).shape[0]
return total / true_neighbors.size
gt = np.argsort(-(Q @ X.T), axis=1)[:, :10] # exact top-10 (chunk Q for large sets)
for leaves in (20, 50, 100, 200):
for pre in (100, 250):
t0 = time.perf_counter()
nb, _ = searcher.search_batched(Q, leaves_to_search=leaves, pre_reorder_num_neighbors=pre)
qps = len(Q) / (time.perf_counter() - t0)
print(leaves, pre, round(compute_recall(nb, gt), 4), int(qps))Use queries drawn from real traffic, not a slice of the indexed data. Self-queries find themselves at score 1 and inflate recall. Report recall at your actual k, and look at the tail as well: a mean of 0.95 can hide a cluster of queries at 0.6.
ScaNN, HNSW or IVF-PQ
| ScaNN (tree + AH) | HNSW | IVF-PQ | |
|---|---|---|---|
| Strength | throughput at high recall for MIPS on CPU | low latency, strong recall, easy inserts | very compact, scales to billions |
| Memory | codes plus raw vectors for reorder | vectors plus graph links (largest) | codes only, unless refining |
| Updates | plan for periodic rebuilds | incremental inserts are natural | add to lists, retrain centroids eventually |
| Score type | inner product native; cosine via normalising | any metric | L2 or inner product |
If your index changes every minute and latency matters most, HNSW is usually simpler. If memory is the hard limit, IVF-PQ (possibly with OPQ rotations) wins. ScaNN is strongest on large, mostly static, inner-product workloads where you want high recall at high QPS on commodity CPUs.
Running it in production
- Build blue/green. Serialize each build to a versioned directory, load it in a fresh process, run a recall smoke test, then switch traffic. Keep the previous directory for rollback.
- Rebuild on drift. Partitions and codebooks are trained on a sample. When new data comes from a different distribution, such as a new embedding model, new languages or a new product line, leaves become unbalanced and recall decays quietly. Track leaf-size skew and a nightly recall probe against brute force on a sample.
- Never mix embedding models. Vectors from two model versions share no geometry. Re-embed everything and rebuild.
- Training sample size. Keep
training_sample_sizewell abovenum_leaves, and make it representative, stratified if your corpus is skewed. - Managed options. Some databases ship ScaNN-based indexes. AlloyDB is one (see AlloyDB). That can be the better choice than operating the library yourself.
Failure modes
- Unnormalised data with dot_product. If you wanted cosine similarity, long vectors win every query. Normalise both the dataset and the queries.
- reorder not above k. Final results are then pure quantized scores, and the ranking quality drops sharply.
- Threshold copied across dimensions. As the η table shows, T = 0.2 means very different things at d = 128 and d = 768. Re-tune T per model.
- Tiny datasets. Under about 20k points, a tree adds error and buys little. Use brute force, which is exact or nearly so.
- Optimistic benchmarks. Self-queries, cached warm runs, or measuring single queries when production sends batches (or the reverse) all mislead. Benchmark the shape of your real traffic.
What to do next
- Normalise your embeddings and compute exact top-k ground truth for a few thousand real queries.
- Run the NumPy trainer at T = 0.05 and T = 0.2 on a sample to see the anisotropic effect on your own data.
- Build ScaNN with num_leaves near √n, dimensions_per_block 2, and reorder above k.
- Sweep leaves_to_search and pre_reorder_num_neighbors, and keep the cheapest pair that meets your recall target.
- Re-tune anisotropic_quantization_threshold per embedding model and dimension.
- Ship it blue/green with a nightly recall probe and a rebuild trigger on drift.