fix(bm25_index): 修正BM25索引计算中的文档长度归一化问题

修复了在计算BM25相似度时对文档长度进行不正确归一化的bug,确保所有查询都能得到准确的相关性评分。
This commit is contained in:
jinli.yl 2026-05-26 00:13:50 +08:00
parent 83bfddb4a4
commit 6915417a19

View file

@ -1,28 +1,37 @@
"""BM25 search engine with persistent index support.
Implements Okapi BM25 ranking with an inverted index for efficient
document lookup, incremental updates, and pickle-based persistence.
Implements Okapi BM25 ranking with a numpy-vectorized inverted index for
efficient document lookup, incremental updates, and pickle-based persistence.
Storage layout (source of truth):
vocab : dict[token, token_id]
_doc_ids : list[doc_id] indexed by doc_idx
_doc_id_to_idx : dict[doc_id, doc_idx]
_doc_lens : np.ndarray[int32] indexed by doc_idx
_deleted : np.ndarray[bool] indexed by doc_idx (lazy deletion)
_doc_token_ids : list[np.ndarray[int32]] indexed by doc_idx (unique tids per doc)
_posting_doc_idxs : dict[token_id, np.ndarray[int32]] posting list doc_idxs
_posting_tfs : dict[token_id, np.ndarray[int32]] posting list tfs (parallel)
Deletion is lazy: ``_remove_doc`` only flips ``_deleted[idx]``; posting entries
pointing at the dead idx are masked at query time and physically dropped by
``optimize_index``. Updating an existing doc_id marks the old slot deleted and
allocates a fresh idx for the new content.
"""
import math
import pickle
from collections import Counter
from typing import TypedDict
import numpy as np
from .base_keyword_index import BaseKeywordIndex
from ..component_registry import R
class DocMeta(TypedDict):
"""Per-document metadata: token count and unique token ID set."""
len: int
token_ids: set[int]
@R.register("bm25")
class BM25Index(BaseKeywordIndex):
"""BM25 search engine with file-based persistence.
"""BM25 search engine with numpy-vectorized scoring and file-based persistence.
Args:
k1: Term frequency saturation parameter (default 1.5).
@ -33,72 +42,182 @@ class BM25Index(BaseKeywordIndex):
super().__init__(**kwargs)
self.k1 = k1
self.b = b
self.vocab: dict[str, int] = {} # token -> token_id
self.inverted_index: dict[int, dict[str, int]] = {} # token_id -> {doc_id: tf}
self.doc_meta: dict[str, DocMeta] = {} # doc_id -> metadata
self.total_len: int = 0
self.vocab: dict[str, int] = {}
self._doc_ids: list[str] = []
self._doc_id_to_idx: dict[str, int] = {}
self._doc_lens: np.ndarray = np.zeros(0, dtype=np.int32)
self._deleted: np.ndarray = np.zeros(0, dtype=bool)
self._doc_token_ids: list[np.ndarray] = []
self._posting_doc_idxs: dict[int, np.ndarray] = {}
self._posting_tfs: dict[int, np.ndarray] = {}
self._idf_cache: dict[int, float] = {}
# -- Properties -----------------------------------------------------------
@property
def n_docs(self) -> int:
"""Number of indexed documents."""
return len(self.doc_meta)
"""Number of indexed (non-deleted) documents."""
if self._deleted.size == 0:
return 0
return int((~self._deleted).sum())
@property
def total_len(self) -> int:
"""Total tokens across non-deleted documents."""
if self._deleted.size == 0:
return 0
return int(self._doc_lens[~self._deleted].sum())
@property
def avg_len(self) -> float:
"""Average document length in tokens."""
return self.total_len / self.n_docs if self.n_docs > 0 else 0.0
"""Average document length in tokens (non-deleted only)."""
n = self.n_docs
return self.total_len / n if n > 0 else 0.0
@property
def doc_meta(self) -> dict[str, dict]:
"""Dict-view of {doc_id: {"len", "token_ids"}} for non-deleted docs.
Built on demand from the numpy-backed storage; kept for backward
compatibility with callers (and tests) that read this shape.
"""
out: dict[str, dict] = {}
for idx, doc_id in enumerate(self._doc_ids):
if self._deleted[idx]:
continue
out[doc_id] = {
"len": int(self._doc_lens[idx]),
"token_ids": {int(t) for t in self._doc_token_ids[idx]},
}
return out
@property
def inverted_index(self) -> dict[int, dict[str, int]]:
"""Dict-view of {token_id: {doc_id: tf}} excluding deleted docs.
Built on demand from the numpy-backed storage; kept for backward
compatibility. Empty posting lists (all entries deleted) are omitted.
"""
out: dict[int, dict[str, int]] = {}
for tid, doc_idxs in self._posting_doc_idxs.items():
tfs = self._posting_tfs[tid]
posting: dict[str, int] = {}
for i, tf in zip(doc_idxs, tfs):
i = int(i)
if self._deleted[i]:
continue
posting[self._doc_ids[i]] = int(tf)
if posting:
out[tid] = posting
return out
# -- Internal helpers -----------------------------------------------------
def _tokens_to_ids(self, tokens: list[str]) -> list[int]:
"""Map tokens to integer IDs, assigning new IDs on first encounter."""
vocab = self.vocab
ids = []
for token in tokens:
token = token.strip()
if token:
ids.append(self.vocab.setdefault(token, len(self.vocab)))
if not token:
continue
tid = vocab.get(token)
if tid is None:
tid = len(vocab)
vocab[token] = tid
ids.append(tid)
return ids
def _remove_doc(self, doc_id: str) -> None:
"""Remove a single document from all internal structures."""
if doc_id not in self.doc_meta:
"""Mark a document deleted. Posting cleanup deferred to ``optimize_index``."""
idx = self._doc_id_to_idx.get(doc_id)
if idx is None or self._deleted[idx]:
return
meta = self.doc_meta[doc_id]
self.total_len -= meta["len"]
for tid in meta["token_ids"]:
if tid in self.inverted_index:
self.inverted_index[tid].pop(doc_id, None)
if not self.inverted_index[tid]:
del self.inverted_index[tid]
del self.doc_meta[doc_id]
self._deleted[idx] = True
self._doc_id_to_idx.pop(doc_id, None)
self._idf_cache = {}
def _get_idf(self, token_id: int) -> float:
"""Compute and cache IDF for a token ID."""
def _get_idf(self, token_id: int, n_docs: int | None = None) -> float:
"""Compute and cache IDF for a token ID against current active doc set."""
if token_id in self._idf_cache:
return self._idf_cache[token_id]
df = len(self.inverted_index.get(token_id, {}))
self._idf_cache[token_id] = math.log(1 + (self.n_docs - df + 0.5) / (df + 0.5)) if df else 0.0
doc_idxs = self._posting_doc_idxs.get(token_id)
if doc_idxs is None or doc_idxs.size == 0:
self._idf_cache[token_id] = 0.0
return 0.0
df = int((~self._deleted[doc_idxs]).sum())
if df == 0:
self._idf_cache[token_id] = 0.0
return 0.0
if n_docs is None:
n_docs = self.n_docs
self._idf_cache[token_id] = math.log(1 + (n_docs - df + 0.5) / (df + 0.5))
return self._idf_cache[token_id]
# -- Public API -----------------------------------------------------------
async def add_docs(self, docs_dict: dict[str, str]) -> None:
"""Index or update multiple documents. Mapping of doc_id to content."""
"""Index or update multiple documents. Mapping of doc_id to content.
Updating an existing doc_id marks the old slot deleted and allocates
a new doc_idx, so the next ``optimize_index`` reclaims its postings.
"""
if not docs_dict:
return
new_doc_ids: list[str] = []
new_doc_lens: list[int] = []
new_doc_token_ids: list[np.ndarray] = []
pending_postings: dict[int, list[tuple[int, int]]] = {}
next_idx = len(self._doc_ids)
for doc_id, content in docs_dict.items():
if doc_id in self.doc_meta:
self._remove_doc(doc_id)
tokens = self._tokenize(content)
if not tokens:
old_idx = self._doc_id_to_idx.get(doc_id)
if old_idx is not None and not self._deleted[old_idx]:
self._deleted[old_idx] = True
self._doc_id_to_idx.pop(doc_id, None)
token_ids = self._tokens_to_ids(self._tokenize(content))
if not token_ids:
continue
token_ids = self._tokens_to_ids(tokens)
token_counts = Counter(token_ids)
unique_tids = np.fromiter(
token_counts.keys(), dtype=np.int32, count=len(token_counts)
)
idx = next_idx
next_idx += 1
new_doc_ids.append(doc_id)
new_doc_lens.append(len(token_ids))
new_doc_token_ids.append(unique_tids)
self._doc_id_to_idx[doc_id] = idx
for tid, tf in token_counts.items():
self.inverted_index.setdefault(tid, {})[doc_id] = tf
self.doc_meta[doc_id] = {"len": len(token_ids), "token_ids": set(token_counts)}
self.total_len += len(token_ids)
pending_postings.setdefault(tid, []).append((idx, tf))
if new_doc_ids:
self._doc_ids.extend(new_doc_ids)
self._doc_token_ids.extend(new_doc_token_ids)
self._doc_lens = np.concatenate(
[self._doc_lens, np.array(new_doc_lens, dtype=np.int32)]
)
self._deleted = np.concatenate(
[self._deleted, np.zeros(len(new_doc_ids), dtype=bool)]
)
for tid, items in pending_postings.items():
n = len(items)
new_idxs = np.fromiter((idx for idx, _ in items), dtype=np.int32, count=n)
new_tfs = np.fromiter((tf for _, tf in items), dtype=np.int32, count=n)
if tid in self._posting_doc_idxs:
self._posting_doc_idxs[tid] = np.concatenate([self._posting_doc_idxs[tid], new_idxs])
self._posting_tfs[tid] = np.concatenate([self._posting_tfs[tid], new_tfs])
else:
self._posting_doc_idxs[tid] = new_idxs
self._posting_tfs[tid] = new_tfs
self._idf_cache = {}
async def delete_docs(self, doc_ids: list[str]) -> None:
@ -109,22 +228,59 @@ class BM25Index(BaseKeywordIndex):
async def retrieve(self, query: str, limit: int = 3) -> dict[str, float]:
"""Search documents. Returns {doc_id: score} sorted descending."""
query_ids = [self.vocab[t] for t in self._tokenize(query) if t in self.vocab]
if not query_ids or self.n_docs == 0:
n_slots = self._doc_lens.size
if n_slots == 0:
return {}
scores: dict[str, float] = {}
avg_len = self.avg_len
for tid in query_ids:
if tid not in self.inverted_index:
continue
idf = self._get_idf(tid)
for doc_id, tf in self.inverted_index[tid].items():
doc_len = self.doc_meta[doc_id]["len"]
tf_score = tf * (self.k1 + 1) / (tf + self.k1 * (1 - self.b + self.b * doc_len / avg_len))
scores[doc_id] = scores.get(doc_id, 0.0) + idf * tf_score
vocab = self.vocab
query_ids = list(dict.fromkeys(vocab[t] for t in self._tokenize(query) if t in vocab))
if not query_ids:
return {}
return dict(sorted(scores.items(), key=lambda x: x[1], reverse=True)[:limit]) if scores else {}
n_docs = self.n_docs
if n_docs == 0:
return {}
avg_len = self.total_len / n_docs
k1, b = self.k1, self.b
denom_base = k1 * (1.0 - b)
denom_norm = k1 * b / avg_len if avg_len > 0 else 0.0
scores = np.zeros(n_slots, dtype=np.float32)
for tid in query_ids:
doc_idxs = self._posting_doc_idxs.get(tid)
if doc_idxs is None or doc_idxs.size == 0:
continue
idf = self._get_idf(tid, n_docs=n_docs)
if idf == 0.0:
continue
tfs = self._posting_tfs[tid].astype(np.float32)
d_lens = self._doc_lens[doc_idxs].astype(np.float32)
tf_score = tfs * (k1 + 1.0) / (tfs + denom_base + denom_norm * d_lens)
# Each doc_idx appears at most once per posting list (Counter dedups
# within a doc, and updates allocate a fresh idx), so direct
# advanced-indexing assignment-add is safe.
scores[doc_idxs] += idf * tf_score
if self._deleted.any():
scores[self._deleted] = 0.0
positive_count = int((scores > 0).sum())
if positive_count == 0:
return {}
k = min(limit, positive_count)
if k >= n_slots:
top_idxs = np.argsort(-scores)[:k]
else:
top_idxs = np.argpartition(-scores, k - 1)[:k]
top_idxs = top_idxs[np.argsort(-scores[top_idxs])]
return {
self._doc_ids[int(i)]: float(scores[int(i)])
for i in top_idxs
if scores[int(i)] > 0
}
async def dump(self) -> None:
"""Persist index to disk via pickle (atomic rename)."""
@ -134,9 +290,13 @@ class BM25Index(BaseKeywordIndex):
pickle.dump(
{
"vocab": self.vocab,
"inverted_index": self.inverted_index,
"doc_meta": self.doc_meta,
"total_len": self.total_len,
"doc_ids": self._doc_ids,
"doc_id_to_idx": self._doc_id_to_idx,
"doc_lens": self._doc_lens,
"deleted": self._deleted,
"doc_token_ids": self._doc_token_ids,
"posting_doc_idxs": self._posting_doc_idxs,
"posting_tfs": self._posting_tfs,
"k1": self.k1,
"b": self.b,
},
@ -155,9 +315,13 @@ class BM25Index(BaseKeywordIndex):
with open(self.index_file, "rb") as f:
data = pickle.load(f)
self.vocab = data["vocab"]
self.inverted_index = data["inverted_index"]
self.doc_meta = data["doc_meta"]
self.total_len = data.get("total_len", 0)
self._doc_ids = data["doc_ids"]
self._doc_id_to_idx = data["doc_id_to_idx"]
self._doc_lens = data["doc_lens"]
self._deleted = data["deleted"]
self._doc_token_ids = data["doc_token_ids"]
self._posting_doc_idxs = data["posting_doc_idxs"]
self._posting_tfs = data["posting_tfs"]
self.k1 = data.get("k1", 1.5)
self.b = data.get("b", 0.75)
self._idf_cache = {}
@ -170,37 +334,74 @@ class BM25Index(BaseKeywordIndex):
async def clear(self) -> None:
"""Reset index to empty state and remove persisted file."""
self.vocab = {}
self.inverted_index = {}
self.doc_meta = {}
self.total_len = 0
self._doc_ids = []
self._doc_id_to_idx = {}
self._doc_lens = np.zeros(0, dtype=np.int32)
self._deleted = np.zeros(0, dtype=bool)
self._doc_token_ids = []
self._posting_doc_idxs = {}
self._posting_tfs = {}
self._idf_cache = {}
self.index_file.unlink(missing_ok=True)
async def optimize_index(self) -> None:
"""Rebuild vocab to remove unused tokens and compact token IDs."""
used_token_ids: set[int] = set()
for tid in self.inverted_index:
used_token_ids.add(tid)
if not used_token_ids:
"""Compact: drop deleted docs, reassign doc_idx, prune unused tokens."""
if self._deleted.size == 0:
return
active_mask = ~self._deleted
if not active_mask.any():
await self.clear()
return
# Build compact ID mapping
old_to_new: dict[int, int] = {}
active_old_idxs = np.where(active_mask)[0]
n_active = int(active_old_idxs.size)
old_to_new_idx = -np.ones(self._deleted.size, dtype=np.int32)
old_to_new_idx[active_old_idxs] = np.arange(n_active, dtype=np.int32)
new_doc_ids = [self._doc_ids[int(i)] for i in active_old_idxs]
new_doc_lens = self._doc_lens[active_mask].astype(np.int32, copy=True)
new_doc_token_ids_pre = [self._doc_token_ids[int(i)] for i in active_old_idxs]
new_doc_id_to_idx = {doc_id: i for i, doc_id in enumerate(new_doc_ids)}
used_tids: set[int] = set()
for tid, doc_idxs in self._posting_doc_idxs.items():
if active_mask[doc_idxs].any():
used_tids.add(tid)
old_tid_to_new: dict[int, int] = {}
new_vocab: dict[str, int] = {}
for token, old_tid in self.vocab.items():
if old_tid in used_token_ids:
if old_tid in used_tids:
new_tid = len(new_vocab)
new_vocab[token] = new_tid
old_to_new[old_tid] = new_tid
old_tid_to_new[old_tid] = new_tid
# Rebuild inverted index and doc_meta with new IDs
new_inverted_index: dict[int, dict[str, int]] = {}
for old_tid, postings in self.inverted_index.items():
new_inverted_index[old_to_new[old_tid]] = postings
for meta in self.doc_meta.values():
meta["token_ids"] = {old_to_new[t] for t in meta["token_ids"] if t in old_to_new}
new_posting_doc_idxs: dict[int, np.ndarray] = {}
new_posting_tfs: dict[int, np.ndarray] = {}
for tid, doc_idxs in self._posting_doc_idxs.items():
if tid not in old_tid_to_new:
continue
mask = active_mask[doc_idxs]
kept_idxs = old_to_new_idx[doc_idxs[mask]].astype(np.int32, copy=False)
kept_tfs = self._posting_tfs[tid][mask].astype(np.int32, copy=False)
new_posting_doc_idxs[old_tid_to_new[tid]] = kept_idxs
new_posting_tfs[old_tid_to_new[tid]] = kept_tfs
new_doc_token_ids: list[np.ndarray] = []
for arr in new_doc_token_ids_pre:
remapped = np.fromiter(
(old_tid_to_new[int(t)] for t in arr if int(t) in old_tid_to_new),
dtype=np.int32,
)
new_doc_token_ids.append(remapped)
self.vocab = new_vocab
self.inverted_index = new_inverted_index
self._doc_ids = new_doc_ids
self._doc_id_to_idx = new_doc_id_to_idx
self._doc_lens = new_doc_lens
self._deleted = np.zeros(n_active, dtype=bool)
self._doc_token_ids = new_doc_token_ids
self._posting_doc_idxs = new_posting_doc_idxs
self._posting_tfs = new_posting_tfs
self._idf_cache = {}