mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-05 02:41:43 +00:00
perf(bm25): maintain live document length incrementally (#590)
Co-authored-by: machaoxin0407 <221922045+machaoxin0407@users.noreply.github.com>
This commit is contained in:
parent
d529ec5256
commit
49bbbc93ff
3 changed files with 249 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
106
scripts/benchmark_bm25_length.py
Normal file
106
scripts/benchmark_bm25_length.py
Normal 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())
|
||||
135
tests/unit/test_bm25_live_length.py
Normal file
135
tests/unit/test_bm25_live_length.py
Normal 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"])
|
||||
Loading…
Add table
Reference in a new issue