Embeddings for categorical variables
Be able to replace one-hot with learnt embeddings in tabular data.
Prerequisites
- DEmbeddings — words as vectorsrequired
- DFeature engineeringrequired
Intuition
One-hot gives every category its own column with a one in it. That works — until the categories become many.
| The number of categories | One-hot | An embedding (dim 8) |
|---|---|---|
| 5 | 5 columns | 40 numbers — a worse deal |
| 100 | 100 columns | 800 numbers |
| 50 000 postcodes | 50 000 columns | 400 000 numbers, but dense |
The real problem with one-hot is not the number of columns but that all the categories are equally different from one another. The postcodes 11122 and 11123 are neighbours in reality but as different as 11122 and 98765 in the one-hot space.
An embedding is a learnt table: every category gets a vector of a few numbers, and those numbers are adjusted during the training. Categories that behave alike end up near each other — without anybody having said that they belong together.
It is the same idea as word embeddings, applied to tabular data.
Formal
The dimension. A rule of thumb from fast.ai that works surprisingly well:
| Categories | The dimension |
|---|---|
| 10 | 6 |
| 100 | 22 |
| 1 000 | 79 |
| 50 000 | 600 |
The exact formula matters less than the order of magnitude — and that the dimension is much smaller than the number of categories.
An embedding table is a lookup, not a matrix multiplication. Mathematically it is one-hot times a weight matrix, but the implementation looks the row up directly — which is what makes 50 000 categories practically possible.
Three things that have to be handled:
| The problem | The solution |
|---|---|
| An unknown category in production | reserve index 0 for <unknown> |
| Rare categories | merge everything below a threshold into <rare> |
| Leakage | build the category ordering on the training data, not on everything |
When does it pay off?
| The situation | Choose |
|---|---|
| Few categories, a linear model | one-hot |
| Trees or boosting | one-hot or ordinal — trees handle it anyway |
| Many categories + a neural network | an embedding |
| Many categories, no deep learning | target encoding with out-of-fold |
The unexpected return is that the embeddings become interpretable. If you train on shop sales and project the shop embeddings down to two dimensions, shops of the same type or in the same region often group together — even though you never fed that information in. It is a good check that the model has learnt something real.
Code
import torch, torch.nn as nn
CATEGORIES = {"shop": 1100, "product": 4200, "weekday": 7}
def dim_for(n):
return min(600, int(round(1.6 * n ** 0.56)))
for name, n in CATEGORIES.items():
print(f"{name:<9} {n:>5} categories → dimension {dim_for(n)}")
# shop 1100 categories → dimension 79
# product 4200 categories → dimension 168
# weekday 7 categories → dimension 5
class TabularNet(nn.Module):
def __init__(self, categories: dict[str, int], n_numeric: int):
super().__init__()
# +1 for index 0 = an unknown category
self.emb = nn.ModuleDict({k: nn.Embedding(n + 1, dim_for(n))
for k, n in categories.items()})
width = sum(e.embedding_dim for e in self.emb.values()) + n_numeric
self.head = nn.Sequential(nn.Linear(width, 128), nn.ReLU(),
nn.Dropout(0.2), nn.Linear(128, 1))
def forward(self, cat: dict[str, torch.Tensor], num: torch.Tensor):
parts = [self.emb[k](v) for k, v in cat.items()] + [num]
return self.head(torch.cat(parts, dim=1)).squeeze(-1)
m = TabularNet(CATEGORIES, n_numeric=4)
print(sum(p.numel() for p in m.parameters()), "parameters")
# After training: do shops with the same profile resemble each other?
E = m.emb["shop"].weight.detach()
E = E / E.norm(dim=1, keepdim=True)
similarity = E @ E.t()
similarity.fill_diagonal_(-1)
print("shop 5 most resembles:", int(similarity[5].argmax()))
# An unknown category in production is mapped to 0 — it does not crash
index = {"Malmö": 1, "Lund": 2}
print(index.get("Kiruna", 0)) # 0 = <unknown>
Mastery means
- Explains when embeddings beat one-hot
- Chooses the dimension for an embedding
- Interprets a learnt embedding
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Entity Embeddings of Categorical Variables — arXiv (open access; licence per article)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0