An annotated companion · AI Primer

Matryoshka Representation Learning, annotated

About this page. This is a companion, not a copy. It follows the paper section by section, quotes only a sentence or two per section (clearly marked), and explains everything in its own words. The paper is distributed under arXiv's standard licence, so its results appear here as selected numbers written for this page, with attribution, and its figures are redrawn from scratch. The live experiment in §3 runs on synthetic data in your browser. Read the original alongside: every section links to it.

How to read this page

  • Any dotted word explains itself when you hover it, tab to it, or tap it.
  • Every symbol inside an equation does the same.
  • The truncation experiment in §3 and the retrieval-cost calculator in §4.3 are live.

Each idea climbs the same ladder: everyday picture, tiny example, diagram, the math, why it matters. The dimensions and compression lesson reproduces a Matryoshka-style truncation experiment in NumPy.

Abstract · original

“This leads us to ask: can we design a flexible representation that can adapt to multiple downstream tasks with varying computational resources?”Kusupati et al. (2022), Abstract

Everyday picture

A set of Russian nesting dolls: open the big doll and a smaller, complete doll is inside, and a smaller one inside that. A Matryoshka embedding is a list of numbers whose first 8 numbers are already a usable (coarse) embedding, the first 16 a better one, and so on up to the full length. You pick how much to keep, per task, without retraining or re-encoding anything.

What the paper claims

  • Each prefix of the embedding is at least as accurate as a separate model trained at that small size.
  • For ImageNet classification, the same accuracy with an embedding up to 14× smaller.
  • For large-scale retrieval, up to 14× faster in wall-clock time at the same accuracy.
  • Up to 2% better accuracy on rare classes, and as robust as ordinary embeddings.
  • No extra cost at inference, and it works for images, text and image-text models.

Why it matters today

Many production text-embedding models are trained this way, so you can store shorter vectors for cheaper search. When a provider lets you choose the embedding size, this is usually the technique behind it.

1 Introduction · original

Everyday picture

Computing an embedding is paid once per item. Using it is paid every time you search or classify, and that cost grows with the number of dimensions × the number of items. At web scale the using dominates. Yet ordinary training spreads information across all dimensions, so you can't just drop some: every number carries a bit of everything.

Tiny example

Store 10 million embeddings of 2,048 numbers at 4 bytes each: 10,000,000 × 2,048 × 4 = 81.9 GB. Keep only the first 256 numbers and it is 10.2 GB, 8× less, with every search 8× cheaper too. The question is whether those 256 numbers are any good. Normally they aren't. Matryoshka training makes them good.

Why it matters

The usual alternatives each cost something: train several models of different sizes (and store every item several times), compress afterwards (and lose accuracy), or pick features on the fly (expensive). Matryoshka training gets all sizes from one model and one stored vector. The storage arithmetic is worked through in the compression lesson.

2 Related work · original

The paper sits between three lines of work. Representation learning produces embeddings (supervised, contrastive or masked-language training); Matryoshka training adds to any of them. Efficient search uses approximate nearest-neighbour indexes such as HNSW, which cut the dependence on the number of items but not on the dimension; Matryoshka embeddings cut the dimension and combine with those indexes. And earlier “nested dropout” learned ordered representations by optimizing every length; this paper optimizes only about log₂(d) lengths, which is what makes it practical at scale.

3 Matryoshka Representation Learning · original

Everyday picture

A teacher marks one essay several times: once reading only the first paragraph, once the first two, and so on, and the student is graded on every reading. The student quickly learns to put the most important point first. Matryoshka training grades the embedding at several lengths at once, so the network learns to pack the most useful information into the first numbers.

Hover or tap a doll (a prefix length) to see what it costs and what it is good for.

Figure 1 of the paper (the nested representation), redrawn. Based on Kusupati et al. (2022), Figure 1.

Reading it: the long bar at the bottom is one 2,048-number embedding. Each shorter bar above it marks a prefix: the first 8 numbers, the first 16, 32, and so on. Every one of them is a complete embedding on its own, and each larger one contains the smaller ones, like nesting dolls. Each is twice the length of the one below it, so there are only 9 of them for 2,048 numbers. The lengths are drawn on a doubling (log) scale; the 8-number doll is really 256 times shorter than the full one.

The math

In words: “take the first m numbers of each embedding, for every m in the chosen set; give each prefix its own small classifier; add up the classification losses of all prefixes; and train the network and all the classifiers to make that total small.”

With the numbers: say for one image the cross-entropy losses for m = 8, 16, 32, …, 2,048 are 1.2, 0.8, 0.7, 0.65, 0.6, 0.57, 0.54, 0.52 and 0.5: worst with the first 8 numbers, best with all 2,048. With every cm = 1 the training loss for this image is the sum of all nine, 6.08; the gradient from the 8-number term pushes the most important information into the first 8 numbers, because that is the only place it can help.

In Python:

# the nested sizes m
M = [8, 16, 32, 64, 128, 256, 512, 1024, 2048]
# L for each prefix F(x)_1:m
loss = [1.2, 0.8, 0.7, 0.65, 0.6, 0.57, 0.54, 0.52, 0.5]
# every c_m = 1
c = [1] * len(M)
# one image
N = 1
round(sum(c_m * L_m for c_m, L_m in zip(c, loss)) / N, 2)  # → 6.08

With a ResNet-50 on ImageNet, d = 2,048 and the set of sizes is ℳ = {8, 16, 32, …, 1024, 2048}: 9 sizes, halving each time. The paper sets every cm to 1.

A cheaper variant: MRL-E

Instead of 9 separate classifiers, share one weight matrix and let each size use its first m columns. That roughly halves the classifier memory, which matters when there are millions of classes. The paper calls this MRL-E; it is within about 1% of full MRL from 16 dimensions up.

Try it: truncate an embedding, live

Computing…

Reading it: this is a real experiment computed in your browser on synthetic data: 600 “documents” of 64 numbers, and for each a noisy “query” copy. The task is to find each query's own document by cosine similarity using only the first m numbers (the y-axis is recall@1: how often the top result is the right one). The solid blue curve uses embeddings whose information is packed into the first numbers, the way Matryoshka training arranges it. The dashed red curve holds exactly the same information, rotated so it is spread evenly over all 64 numbers, the way ordinary training leaves it. Both reach 100% with all 64 numbers. Cut them short and the front-loaded one keeps working far longer: that gap is what Matryoshka training buys. The slider reads both curves off at one size and reports what that size would save in storage for 10 million vectors; hover the chart to read any point.

Why it matters

The trick costs almost nothing: a few extra small classifiers during training, nothing at inference. You get one model and one stored vector serving every budget.

4 Applications · original

4.1 Where it was tried · original

Supervised image models (ResNet-50 on ImageNet, a Vision Transformer on 300 million images), a contrastive image-text model (ALIGN, similar to CLIP) and a masked language model (BERT). For 768-number models the sizes were {12, 24, 48, 96, 192, 384, 768}. The same training settings as the ordinary baselines were reused, with no special tuning.

4.2 Classification · original

What they found

  • At every size, the Matryoshka prefix was at least as accurate as a separate ResNet-50 trained at that size alone, and up to 2% better at the smallest sizes when judged by nearest-neighbour lookup.
  • Sizes that were never explicitly trained, between the powers of two, still worked well: the information spreads smoothly.
  • Compressing an ordinary embedding afterwards (with SVD), or keeping random features, lost much more accuracy as size shrank.

Adaptive classification

Everyday picture: a doctor's triage. Easy cases are decided at a glance; only unclear cases get the full examination. The classifier looks at the first 8 numbers; if it is confident enough, it stops. Otherwise it looks at 16, then 32, and so on, with the confidence thresholds learned on held-out data.

With the numbers: on ImageNet this cascade reached 76.30% accuracy, the same as an ordinary 512-number model, while using on average only about 37 numbers per image: about 14× smaller, and just 0.8 points below the full 2,048-number model.

4.3 Retrieval · original

Everyday picture

Finding a book in a huge library: first skim the spine colours to pull 200 likely candidates off the shelves (cheap, rough), then read those 200 blurbs carefully (expensive, precise). Adaptive retrieval does exactly this with one embedding: shortlist with its first few numbers, rerank the shortlist with all of them.

Across all sizes, Matryoshka embeddings retrieved up to 3% better than ordinary embeddings of the same size on ImageNet, measured by mAP@10 (mean average precision in the top 10).

The cost model

In words: “compare the query with every item using only Ds numbers, then compare it with the K shortlisted items using Dr numbers, instead of comparing it with every item using all d numbers.”

With the numbers: ImageNet's database has about 1.3 million images. Exhaustive search with 2,048 numbers costs 1.3M × 2,048 ≈ 2.66 × 10⁹ multiply-adds per query, matching the paper's 2.6 GFLOPs. Adaptive retrieval with Ds = 16, K = 200 and Dr = 2,048 costs 1.3M × 16 + 200 × 2,048 ≈ 2.12 × 10⁷: about 126× less. The paper reports ~128× in theory, and a measured 14× in wall-clock time using HNSW, at the same accuracy.

In Python:

# items in the database
N = 1.3e6
d, D_s, K, D_r = 2048, 16, 200, 2048
exhaustive = N * d
# shortlist cheaply, then rerank K
adaptive = N * D_s + K * D_r
# × 10⁹ multiply-adds
round(exhaustive / 1e9, 2)  # → 2.66
# × 10⁷ multiply-adds
round(adaptive / 1e7, 2)  # → 2.12
round(exhaustive / adaptive)  # → 126

Reading it: the two bars compare multiply-adds per query, on a log scale: exhaustive search at full size, and the two-stage shortlist-then-rerank search. The defaults are the paper's ImageNet-1K setting. Grow the database to ImageNet-4K's 4.2 million images and the paper needed Ds = 64 to keep accuracy, which the calculator puts at about 32× cheaper (the paper: ~32× in theory, ~6× measured). The rerank term K × Dr is tiny, so almost all the cost is the first pass: the shorter the shortlisting prefix, the bigger the saving.

Funnel retrieval

Choosing Ds and Dr is fiddly, so the paper also proposes a funnel: halve the shortlist and double the length at each step, for example shortlists 200 → 100 → 50 → 25 → 10 with lengths 16 → 32 → 64 → 128 → 256 → 2,048. It was as accurate as full-size search at about 128× lower theoretical cost.

Why it matters today

This is the pattern behind “search with short vectors, rescore with long ones” in modern vector databases. Combine it with ANN indexes and quantization and the savings multiply.

5 Further analysis · original

  • Robustness: at least as robust as ordinary embeddings on shifted versions of ImageNet, 0.6 points better on the hard ImageNet-A set, and up to 3% better mAP@10 when queries come from a shifted set.
  • Rare classes: up to 2% more accurate on new, rarely seen classes, without losing accuracy elsewhere. Rare classes seem to need the higher dimensions most.
  • Disagreement between sizes: some images are classified better at small sizes. With perfect routing of each image to its best size, accuracy could rise by up to 4.6%.
  • Coarse before fine: small prefixes lose fine-grained accuracy fast but keep broad categories (“a dog”) well: the nesting mirrors a hierarchy from general to specific.
  • Ablations: existing models can be given Matryoshka structure with cheap partial fine-tuning; the halving spacing beats evenly spaced sizes; and retrieval gains level off beyond a certain shortlist size.

6 Discussion · original

The authors conclude that one Matryoshka embedding matches fixed-size models at every size, giving about 14× smaller classifiers on average and 128× cheaper (14× faster) retrieval at the same accuracy. Future directions they list: learning the loss weights cm, using different losses at different sizes (for example, favouring recall at 8 numbers and robustness at 2,048), and learning search structures on top of the nested embedding.

What changed since

In the paperCommon todayLesson
Image models and BERTMany text-embedding models trained with Matryoshka losses, with selectable output sizescompression
Shortlist with a short prefix, rerank with the full vectorStandard two-stage vector search, often with quantized vectors in the first stageANN
Fixed thresholds for adaptive classificationThe same cascade idea in routing requests to cheap or expensive modelscost

Glossary

Every term with hover guidance on this page, in one place.