diff --git a/reme/components/keyword_index/bm25_index.py b/reme/components/keyword_index/bm25_index.py index 4bd4c14d..4860d0ad 100644 --- a/reme/components/keyword_index/bm25_index.py +++ b/reme/components/keyword_index/bm25_index.py @@ -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 diff --git a/scripts/benchmark_bm25_length.py b/scripts/benchmark_bm25_length.py new file mode 100644 index 00000000..8754cec3 --- /dev/null +++ b/scripts/benchmark_bm25_length.py @@ -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()) diff --git a/tests/unit/test_bm25_live_length.py b/tests/unit/test_bm25_live_length.py new file mode 100644 index 00000000..5ba674d3 --- /dev/null +++ b/tests/unit/test_bm25_live_length.py @@ -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"])