1. What is it?#
An embedding is a learned vector that represents a discrete thing — a token, a word, a user, a
product. The embedding table is a matrix of shape (V, d): one row per vocabulary item.
token id 1547 → row 1547 of the table → [0.03, -0.11, ..., 0.42] (d numbers)That’s it. Embedding lookup is an array index, not a matrix multiply — though it is mathematically equivalent to multiplying a one-hot vector by the table.
2. Why does it exist?#
Because models do arithmetic, and “the token cat” is not a number you can do arithmetic with.
You need a representation that (a) is continuous, (b) has enough dimensions to encode meaning,
and (c) can be learned.
One-hot encoding fails (a) and is absurdly wasteful: a 128,000-dimensional vector with one 1 in it. Embeddings compress that into ~4,000 dense dimensions where geometric relationships encode semantic ones.
3. Simple analogy#
A library’s coordinate system. Instead of a shelf number (arbitrary), each book gets coordinates in a “meaning space”: one axis roughly fiction↔nonfiction, another technical↔popular, and thousands more that resist naming. Books about similar things end up near each other, and the model learns those coordinates itself from data.
The famous consequence: king - man + woman ≈ queen. Directions in the space carry meaning.
(This works better for classic word embeddings than for modern contextual ones, but the
intuition holds.)
4. Tiny example#
package main
import (
"fmt"
"math/rand"
)
func main() {
const V, d = 10, 4
table := make([][]float32, V) // the embedding "matrix": V rows of d numbers
for i := range table {
table[i] = make([]float32, d)
for j := range table[i] {
table[i][j] = float32(rand.NormFloat64())
}
}
tokenIDs := []int{3, 7, 1, 3} // a sequence
embeddings := make([][]float32, len(tokenIDs))
for i, id := range tokenIDs {
embeddings[i] = table[id] // a lookup: no arithmetic at all
}
fmt.Println(len(embeddings), "x", len(embeddings[0])) // 4 x 4
fmt.Println(embeddings[0], embeddings[3]) // identical: both are token 3
}
That single indexing operation is the entire embedding layer. Equivalent but wasteful:
// The same result the slow way: a one-hot vector times the table.
oneHot := make([]float32, V)
oneHot[tokenIDs[0]] = 1
slow := make([]float32, d)
for v := 0; v < V; v++ { // V*d multiplications, all but d of them by zero
for j := 0; j < d; j++ {
slow[j] += oneHot[v] * table[v][j]
}
}
fmt.Println(slow) // same as embeddings[0]
The matmul version costs 2·4·10·4 = 320 FLOPs; the indexing version costs zero FLOPs and
4 memory reads. For a real model (V=128256, d=4096), the matmul version would cost
2·S·128256·4096 FLOPs to accomplish S array lookups. Always index.
5. Technical explanation#
Size#
Embedding table: V × d parameters
Llama 3 8B: 128,256 × 4,096 = 525 M params = 1.05 GB in FP16
Llama 3 70B: 128,256 × 8,192 = 1.05 B params = 2.1 GB in FP16
Gemma 2B: 256,000 × 2,048 = 524 M params ← 25% of a 2B model!For small models with large vocabularies, the embedding table is a huge fraction of the parameters. This matters for quantization (embeddings are often kept at higher precision) and for on-device deployment.
Weight tying#
Many models share the input embedding and the output projection (lm_head):
logits = hidden @ embedding_table.T # tiedSaves V × d parameters (1 GB for Llama-3-8B). Llama 3 8B ties; Llama 3 70B does not. Check
config.tie_word_embeddings — getting it wrong changes your memory estimate by gigabytes.
Positional information#
The embedding table alone has no notion of order — "dog bites man" and "man bites dog" would
give identical bags of vectors. Position must be injected:
| Scheme | How | Used by |
|---|---|---|
| Learned absolute | add a learned vector per position | BERT, GPT-2 |
| Sinusoidal | add fixed sin/cos patterns | original transformer |
| RoPE (rotary) | rotate q and k by a position-dependent angle | Llama, Mistral, Qwen, most modern |
| ALiBi | add a distance-based bias to attention scores | Bloom, MPT |
RoPE dominates modern LLMs and has direct inference consequences:
# RoPE is applied to q and k, NOT to the value vectors, and NOT to the residual stream
q_rot = rotate(q, position)
k_rot = rotate(k, position)Why it matters for you:
- The KV cache stores rotated keys. If you change position (e.g. when reusing a prefix at a different offset), the cached keys are wrong. This constrains prefix caching (Section V.11).
- RoPE has a
thetabase parameter that controls the frequency spectrum. Long-context extension methods (NTK scaling, YaRN) change it — see Section XIII.04. - RoPE is applied every step during decode, so it must be a fast, fused kernel.
The output side: the LM head#
hidden (B, S, d) @ W_out (d, V) → logits (B, S, V)For V=128256, d=4096, B=1, S=1: 2 · 4096 · 128256 = 1.05 GFLOPs — for a single token! That is
7% of the entire 8B model’s per-token FLOPs, in one matrix.
And the output is (B, S, V): at B=32, S=1, FP32 that’s 16 MB per step; during prefill with
S=2048 it would be 33 GB. This is why you only compute logits for the last position during
prefill. A common performance bug in custom code is computing them for all positions.
6. Under the hood#
Embedding lookup is a gather: random access into a large table.
Table: 1 GB in HBM
Lookup: read d × 2 bytes = 8 KB per token, from an unpredictable locationProperties:
- No spatial locality across tokens (adjacent tokens have unrelated ids).
- Cache-hostile; the prefetcher cannot help.
- But: cheap in absolute terms (8 KB per token), so it rarely matters.
- Exception: recommendation systems with terabyte-scale embedding tables, where the gather is the bottleneck and specialized systems exist.
For LLMs, the embedding lookup is a rounding error. The output projection is not.
7. Performance implications#
- Input embedding: negligible cost. A gather of a few KB per token.
- Output projection: significant. ~5-10% of decode FLOPs, and it produces a large tensor.
- Vocabulary size is a real design parameter. Doubling V doubles the LM head cost and the embedding memory.
- Logits in FP32 vs FP16: sampling in FP32 is more numerically stable but doubles the logits tensor. Most engines compute logits in the model dtype and upcast only for the softmax.
- Vocabulary parallelism: in tensor-parallel setups, the vocabulary is sharded across GPUs so each computes a slice of the logits, followed by an AllGather or a distributed argmax.
8. Production implications#
- Check
tie_word_embeddingswhen estimating memory. - Compute logits only where needed. Last position during prefill; all positions only if you need per-token logprobs.
- Keep sampling on the GPU (Section II.09) — moving a
(B, V)logits tensor to the host every step costs milliseconds. - Embeddings are often excluded from aggressive quantization, because the gather is cheap and the quality cost of quantizing them can be disproportionate.
- Logprob APIs are expensive. Returning top-k logprobs per token requires keeping and sorting the full logits. Price and rate-limit accordingly.
9. Common mistakes#
Implementing embedding as a one-hot matmul. Astronomically wasteful.
Computing logits for all prefill positions. Wastes FLOPs and can OOM.
Forgetting weight tying in memory estimates.
Applying RoPE to values. It goes on q and k only.
Reusing cached KV at a different position offset with RoPE without recomputation or a position-aware scheme. Produces subtly wrong outputs.
Ignoring vocabulary size when comparing models. A “3B” model with a 256k vocabulary has very different serving characteristics from a 3B model with 32k.
10. Hands-on exercise#
A. Measure the tables. For three open models of different sizes, compute the embedding table’s share of total parameters. Which model is most embedding-heavy?
B. Gather cost. Benchmark table[ids] on GPU for a (128256, 4096) FP16 table with batches
of 1, 32, 2048 random ids. What bandwidth do you achieve? Compare to a sequential read of the
same number of bytes.
C. LM head cost. Time the final projection for B ∈ {1, 8, 64} and compare to the time for one transformer layer. What fraction of a decode step is the LM head?
D. RoPE. Implement RoPE from scratch for one head. Verify that the dot product between a
query at position i and a key at position j depends only on i-j (the relative position
property). This property is why RoPE works.
E. The prefill logits trap. Write code that computes logits for all prefill positions and measure memory. Then fix it to compute only the last. Report both.
11. Interview questions#
- What is an embedding table and how large is it for a typical LLM?
- Why is embedding lookup a gather rather than a matmul?
- What is weight tying and how much memory does it save?
- Explain RoPE and two of its consequences for inference.
- Why do you only compute logits for the last position during prefill?
- How does vocabulary size affect serving cost?
- What is vocabulary parallelism and why is it used?
12. Further reading#
- [ESTABLISHED] Su et al., “RoFormer: Enhanced Transformer with Rotary Position Embedding” (2021)
- [ESTABLISHED] Press et al., “Train Short, Test Long: Attention with Linear Biases” (ALiBi, 2021)
- [FUNDAMENTAL] Mikolov et al., word2vec papers — for embedding intuition
- Next: 07 — Softmax and numerical stability