complete restructure
This commit is contained in:
@@ -0,0 +1,133 @@
|
||||
"""RAG retrieval benchmark: p95 latency of get_conversation_context.
|
||||
|
||||
Standalone dev tool (do not import into the package). Seeds temporary
|
||||
SQLite databases with 1k and 5k user/response rows using deterministic
|
||||
fake embeddings (the embedding HTTP call is monkeypatched, so no network
|
||||
is needed), then measures p95 of ChatDatabase.get_conversation_context
|
||||
over repeated queries.
|
||||
|
||||
Run from the repo root:
|
||||
|
||||
uv run python scripts/bench_rag.py
|
||||
|
||||
Prints p95@1k and p95@5k in milliseconds.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
import vibe_bot.db.messages as db_messages
|
||||
from vibe_bot import llm_client
|
||||
from vibe_bot.database import ChatDatabase
|
||||
|
||||
EMBEDDING_DIM = 256
|
||||
NUM_CLUSTERS = 16
|
||||
TOPICS = [
|
||||
"sailing",
|
||||
"baking",
|
||||
"gardening",
|
||||
"astronomy",
|
||||
"chess",
|
||||
"pottery",
|
||||
"mountaineering",
|
||||
"photography",
|
||||
"brewing",
|
||||
"carpentry",
|
||||
"weaving",
|
||||
"falconry",
|
||||
"cartography",
|
||||
"metallurgy",
|
||||
"botany",
|
||||
"masonry",
|
||||
]
|
||||
SIZES = (1000, 5000)
|
||||
WARMUP_QUERIES = 20
|
||||
MEASURED_QUERIES = 200
|
||||
USER_ID = "bench-user"
|
||||
|
||||
|
||||
def fake_embedding(text: str, *, model: str, url: str, api_key: str) -> list[float]:
|
||||
"""Deterministic cluster-structured fake embedding (no network)."""
|
||||
topic = text.split(maxsplit=1)[0].lower()
|
||||
seed = int.from_bytes(hashlib.sha256(topic.encode("utf-8")).digest()[:8], "big")
|
||||
rng = np.random.default_rng(seed)
|
||||
center = np.zeros(EMBEDDING_DIM, dtype=np.float32)
|
||||
span = EMBEDDING_DIM // NUM_CLUSTERS
|
||||
cluster = seed % NUM_CLUSTERS
|
||||
center[cluster * span : (cluster + 1) * span] = 1.0
|
||||
noise = rng.standard_normal(EMBEDDING_DIM, dtype=np.float32)
|
||||
noise /= np.linalg.norm(noise)
|
||||
vector = center + 0.3 * noise
|
||||
return [float(x) for x in vector]
|
||||
|
||||
|
||||
def seed_db(db_path: str, n_rows: int) -> list[str]:
|
||||
"""Seed n_rows user/response pairs; return the user contents as queries."""
|
||||
db = ChatDatabase(db_path=db_path)
|
||||
queries: list[str] = []
|
||||
for i in range(n_rows):
|
||||
topic = TOPICS[i % NUM_CLUSTERS]
|
||||
content = f"{topic} question number {i}"
|
||||
queries.append(content)
|
||||
db.add_message(
|
||||
message_id=f"bench-{i}",
|
||||
user_id=USER_ID,
|
||||
username="bench",
|
||||
content=content,
|
||||
)
|
||||
db.add_message(
|
||||
message_id=f"bench-{i}_response",
|
||||
user_id="bench-bot",
|
||||
username="bench-bot",
|
||||
content=f"response {i}",
|
||||
role="assistant",
|
||||
embed=False,
|
||||
)
|
||||
return queries
|
||||
|
||||
|
||||
def measure(db_path: str, queries: list[str]) -> float:
|
||||
"""Return p95 in ms of get_conversation_context over repeated queries."""
|
||||
db = ChatDatabase(db_path=db_path)
|
||||
for query in queries[:WARMUP_QUERIES]:
|
||||
db.get_conversation_context(USER_ID, query)
|
||||
|
||||
latencies_ms: list[float] = []
|
||||
for query in queries[:MEASURED_QUERIES]:
|
||||
start = time.perf_counter()
|
||||
db.get_conversation_context(USER_ID, query)
|
||||
latencies_ms.append((time.perf_counter() - start) * 1000)
|
||||
return float(np.percentile(latencies_ms, 95))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
llm_client.embedding = fake_embedding # type: ignore[assignment]
|
||||
db_messages.MAX_HISTORY_MESSAGES = 10**9
|
||||
|
||||
results: dict[int, float] = {}
|
||||
for size in SIZES:
|
||||
with tempfile.TemporaryDirectory(prefix="bench_rag_") as tmp:
|
||||
db_path = str(Path(tmp) / f"bench_{size}.db")
|
||||
queries = seed_db(db_path, size)
|
||||
results[size] = measure(db_path, queries)
|
||||
|
||||
p95_1k = results[1000]
|
||||
p95_5k = results[5000]
|
||||
passed = p95_1k <= 10.0 and p95_5k <= 75.0
|
||||
print(f"p95@1k = {p95_1k:.2f} ms")
|
||||
print(f"p95@5k = {p95_5k:.2f} ms")
|
||||
print(f"ratio = {p95_5k / p95_1k:.2f}x (informational)")
|
||||
print(f"target: p95@1k <= 10ms, p95@5k <= 75ms -> {'PASS' if passed else 'FAIL'}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user