perf(bm25): maintain live document length incrementally (#590)

Co-authored-by: machaoxin0407 <221922045+machaoxin0407@users.noreply.github.com>
This commit is contained in:
machaoxin0407 2026-10-03 22:38:11 +08:00 • committed by GitHub
parent d529ec5256
commit 49bbbc93ff
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 249 additions and 1 deletions

View file

@ -47,6 +47,7 @@ class BM25Index(BaseKeywordIndex):
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._total_len = 0
self._deleted: np.ndarray = np.zeros(0, dtype=bool)
self._doc_token_ids: list[np.ndarray] = []
@ -106,7 +107,7 @@ class BM25Index(BaseKeywordIndex):
@property
def total_len(self) -> int:
"""Sum of token counts across live documents."""
return 0 if self._deleted.size == 0 else int(self._doc_lens[~self._deleted].sum())
return self._total_len
@property
def avg_len(self) -> float:
@ -165,6 +166,7 @@ class BM25Index(BaseKeywordIndex):
if idx is None or self._deleted[idx]:
return
self._deleted[idx] = True
self._total_len -= int(self._doc_lens[idx])
self._doc_id_to_idx.pop(doc_id, None)
self._idf_cache = {}
@ -210,6 +212,7 @@ class BM25Index(BaseKeywordIndex):
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)])
self._total_len += int(self._doc_lens[-len(new_doc_ids) :].sum())
def _extend_postings(self, pending: dict[int, list[tuple[int, int]]]) -> None:
"""Append pending (doc_idx, tf) pairs to each token's posting list."""
@ -409,6 +412,8 @@ class BM25Index(BaseKeywordIndex):
self._doc_id_to_idx = data["doc_id_to_idx"]
self._doc_lens = data["doc_lens"]
self._deleted = data["deleted"]
# Derived state: old snapshots remain valid without a format change.
self._total_len = int(self._doc_lens[~self._deleted].sum())
self._doc_token_ids = data["doc_token_ids"]
self._posting_doc_idxs = data["posting_doc_idxs"]
self._posting_tfs = data["posting_tfs"]
@ -467,6 +472,7 @@ class BM25Index(BaseKeywordIndex):
self._doc_ids = []
self._doc_id_to_idx = {}
self._doc_lens = np.zeros(0, dtype=np.int32)
self._total_len = 0
self._deleted = np.zeros(0, dtype=bool)
self._doc_token_ids = []
self._posting_doc_idxs = {}
@ -553,6 +559,7 @@ class BM25Index(BaseKeywordIndex):
self._doc_ids = new_doc_ids
self._doc_id_to_idx = {doc_id: i for i, doc_id in enumerate(new_doc_ids)}
self._doc_lens = self._doc_lens[active_mask].astype(np.int32, copy=True)
self._total_len = int(self._doc_lens.sum())
self._deleted = np.zeros(n_active, dtype=bool)
self._doc_token_ids = new_doc_token_ids
self._posting_doc_idxs = new_posting_idxs

View file

@ -0,0 +1,106 @@
"""Compare live-length lookup and real BM25 queries with the previous scan.
Run from the repository: python -m scripts.benchmark_bm25_length --repeats 15
Setup is excluded; no files, models or network are used. JSON goes to stdout.
"""
# pylint: disable=protected-access
import argparse
import asyncio
import json
import platform
import statistics
import time
from unittest.mock import patch
import numpy as np
from reme.components.keyword_index import BM25Index
from reme.components.tokenizer import RegexTokenizer
def scanning_length(index):
"""Original length calculation, including lazy-deleted slots."""
return int(index._doc_lens[~index._deleted].sum())
async def measure(operation, repeats):
"""Warm up and report milliseconds per operation."""
for _ in range(3):
await operation()
samples = []
for _ in range(repeats):
start = time.perf_counter()
await operation()
samples.append((time.perf_counter() - start) * 1000)
return {"median_ms": statistics.median(samples), "p95_ms": float(np.percentile(samples, 95))}
async def benchmark(size, repeats):
"""Exercise a rare-term query globally and within 100 selected documents."""
index = BM25Index()
index.tokenizer = RegexTokenizer(filter_stopwords=False)
for start in range(0, size, 10000):
docs = {}
for i in range(start, min(start + 10000, size)):
docs[str(i)] = "alpha beta rare" if i % 1000 == 0 else "alpha beta"
await index.add_docs(docs)
selected = [str(i) for i in range(min(size, 100))]
async def length():
return index.total_len
async def query():
return await index.retrieve("rare", 10)
async def filtered():
return await index.retrieve_filtered("rare", 10, selected)
results = []
for deleted_fraction in (0.0, 0.5):
if deleted_fraction:
await index.delete_docs([str(i) for i in range(1, size, 2)])
for name, operation in (("length", length), ("retrieve", query), ("retrieve_filtered", filtered)):
expected = await operation()
cached = await measure(operation, repeats)
with patch.object(BM25Index, "total_len", property(scanning_length)):
assert await operation() == expected
baseline = await measure(operation, repeats)
results.append(
{
"slots": size,
"deleted_fraction": deleted_fraction,
"operation": name,
"baseline": baseline,
"incremental": cached,
},
)
return results
async def main():
"""Print reproducible environment and benchmark results."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--sizes", nargs="+", type=int, default=[10000, 100000, 1000000])
parser.add_argument("--repeats", type=int, default=15)
args = parser.parse_args()
rows = []
for size in args.sizes:
rows.extend(await benchmark(size, args.repeats))
print(
json.dumps(
{
"python": platform.python_version(),
"platform": platform.platform(),
"numpy": np.__version__,
"repeats": args.repeats,
"results": rows,
},
indent=2,
),
)
if __name__ == "__main__":
asyncio.run(main())

View file

@ -0,0 +1,135 @@
"""Derived BM25 length stays consistent across mutations and old snapshots."""
# pylint: disable=protected-access,missing-function-docstring
import random
from unittest.mock import patch
import numpy as np
import pytest
from reme.components.keyword_index import BM25Index
from reme.components.tokenizer import RegexTokenizer
class ScanningIndex(BM25Index):
"""Original implementation used as a scoring oracle."""
@property
def total_len(self):
return int(self._doc_lens[~self._deleted].sum())
def assert_length(index):
assert index.total_len == int(index._doc_lens[~index._deleted].sum())
def make_index(cls=BM25Index):
index = cls()
index.tokenizer = RegexTokenizer(filter_stopwords=False)
return index
@pytest.mark.asyncio
async def test_mutations_and_scores_match_scanning_oracle(tmp_path, monkeypatch):
monkeypatch.chdir(tmp_path)
index, reference = make_index(), make_index(ScanningIndex)
rng = random.Random(417)
for iteration in range(80):
doc_id = f"doc-{rng.randrange(12)}"
if iteration % 4 == 0:
for target in (index, reference):
await target.delete_docs([doc_id, doc_id, "absent"])
else:
text = " ".join(rng.choices(["alpha", "beta", "gamma"], k=rng.randrange(8)))
for target in (index, reference):
await target.add_docs({doc_id: text})
if iteration % 9 == 0:
for target in (index, reference):
await target.optimize_index()
assert_length(index)
assert index.total_len == reference.total_len
assert index.avg_len == reference.avg_len
assert await index.retrieve("alpha beta", 6) == await reference.retrieve("alpha beta", 6)
selected = [f"doc-{i}" for i in range(5)]
assert await index.retrieve_filtered("gamma beta", 3, selected) == await reference.retrieve_filtered(
"gamma beta",
3,
selected,
)
await index.clear()
assert index.total_len == 0
await index.add_docs({"new": "alpha beta"})
await index.delete_docs(["new"])
await index.optimize_index()
assert index.total_len == 0
@pytest.mark.asyncio
async def test_snapshot_roundtrip_rebuilds_length_without_new_fields(tmp_path, monkeypatch):
monkeypatch.chdir(tmp_path)
index = make_index()
await index.add_docs({"old": "alpha beta", "live": "gamma beta alpha"})
await index.delete_docs(["old"])
snapshot = index._snapshot()
assert "total_len" not in snapshot and "_total_len" not in snapshot
restored = make_index()
await restored.add_docs({"unrelated": "word " * 100})
restored._restore(snapshot)
assert restored.total_len == 3
assert await restored.retrieve("beta") == await index.retrieve("beta")
await index.dump()
loaded = make_index()
await loaded.load()
assert loaded.total_len == 3
await loaded.add_docs({"live": "alpha"})
assert loaded.total_len == 1
@pytest.mark.asyncio
async def test_failed_batch_counts_only_published_arrays(tmp_path, monkeypatch):
monkeypatch.chdir(tmp_path)
index = make_index()
await index.add_docs({"old": "alpha beta", "keep": "gamma"})
tokenize = index._tokenize
def fail(text):
if text == "fail":
raise ValueError("tokenizer failed")
return tokenize(text)
with patch.object(index, "_tokenize", side_effect=fail):
with pytest.raises(ValueError, match="tokenizer failed"):
await index.add_docs({"pending": "alpha alpha", "old": "fail"})
# Existing partial-batch behavior retires old but does not append pending.
assert_length(index)
assert index.total_len == 1
@pytest.mark.asyncio
async def test_postings_failure_still_counts_appended_arrays(tmp_path, monkeypatch):
monkeypatch.chdir(tmp_path)
index = make_index()
with patch.object(index, "_extend_postings", side_effect=RuntimeError("posting failed")):
with pytest.raises(RuntimeError, match="posting failed"):
await index.add_docs({"new": "alpha beta"})
assert_length(index)
assert index.total_len == 2
@pytest.mark.asyncio
async def test_query_never_scans_lengths_with_live_mask():
class NoBooleanScan(np.ndarray):
"""Integer posting lookups are allowed, corpus-wide mask scans are not."""
def __getitem__(self, key):
if isinstance(key, np.ndarray) and key.dtype == bool:
raise AssertionError("query scanned all document lengths")
return super().__getitem__(key)
index = make_index()
await index.add_docs({"one": "alpha beta", "two": "beta beta"})
index._doc_lens = index._doc_lens.view(NoBooleanScan)
assert index.total_len == 4
assert await index.retrieve("beta")
assert await index.retrieve_filtered("beta", 1, ["two"])