Embeddings and Vector Search · Retrieval · lesson 6 of 8
BM25 from scratch
about 22 minutes · free · runs in your browser
Step 1 of 2
Rare words carry the signal
Before embeddings there was keyword search, and it is still the better half of a good retrieval system. The idea underneath BM25 is one sentence: a word that appears in every document tells you nothing; a word that appears in three tells you a great deal.
That is inverse document frequency:
import math
idf = math.log(1 + (total_docs - matching_docs + 0.5) / (matching_docs + 0.5))
A term in 1 of 100 documents scores highly. A term in all 100 scores near zero — which is why "the" contributes nothing without anyone maintaining a list of stop words.
Your turn: write idf(term, documents) where each document is a list of words.
Rarer terms must score higher, and a term appearing in every document must score close to
zero without being negative.
You start from this, and edit it in the browser:
import math
def idf(term, documents):
"""Inverse document frequency of a term across the documents."""
return 0.0
Step 2 of 2
The rest of BM25
Two more ideas complete it.
Saturation. A document mentioning "yeast" twenty times is not twenty times more about
yeast than one mentioning it once. BM25 damps repeated terms with k1, so the second
mention adds a lot and the twentieth adds almost nothing.
Length normalisation. A long document contains more of every word by accident, so it
would win everything. b scales the penalty by how much longer than average it is.
tf · (k1 + 1)
score = idf · ---------------------------------
tf + k1 · (1 - b + b · dl / avgdl)
k1 = 1.5 and b = 0.75 are the defaults everywhere, and they are defaults because
they work; tuning them is a late optimisation, not a starting point.
Your turn: write bm25(query_terms, document, documents, k1=1.5, b=0.75) scoring one
document against a query.
You start from this, and edit it in the browser:
import math
def idf(term, documents):
total = len(documents)
matching = sum(1 for d in documents if term in d)
return math.log(1 + (total - matching + 0.5) / (matching + 0.5))
def ():