Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
dc4a96e8d5 | ||
|
|
92219878d9 | ||
|
|
1a3963f306 | ||
|
|
7f656f5e47 | ||
|
|
8e2b678464 |
@@ -17,7 +17,7 @@ once deployed.
|
||||
|
||||
| Tool | Use |
|
||||
|---|---|
|
||||
| `search_docs` | BM25-default search with optional version / platform / bundle filters; cross-encoder reranked when `RERANK_URL` is set |
|
||||
| `search_docs` | Dense-default search with optional version / platform / bundle filters; rerank only when `RERANK_ENABLED=true` |
|
||||
| `get_page` | Full markdown of one page with metadata header + source URL |
|
||||
| `list_versions` | Discover available versions, doc types, and bundle slugs |
|
||||
| `list_cluster` | Cross-version peers of a page (synthesized from same-GUID overlap) |
|
||||
@@ -51,19 +51,19 @@ peer mapping is free (no fuzzy matching needed).
|
||||
|
||||
## Retrieval
|
||||
|
||||
Eval against 22 hand-curated golden queries — see
|
||||
[`eval/results/baseline.md`](eval/results/baseline.md):
|
||||
Eval against 22 hand-curated golden queries. May 2026 baseline is
|
||||
[`eval/results/baseline.md`](eval/results/baseline.md); after nomic
|
||||
prefixes + heading-recursive chunking (2026-09-30, live index):
|
||||
|
||||
| Retriever | MRR | Recall@5 | nDCG@5 | latency |
|
||||
| Retriever | P@1 | MRR | Recall@5 | nDCG@5 |
|
||||
|---|---:|---:|---:|---:|
|
||||
| dense (Ollama nomic-embed-text) | 0.539 | 0.621 | 0.558 | 88 ms |
|
||||
| BM25 (SQLite FTS5) | 0.880 | 0.909 | 0.883 | 3 ms |
|
||||
| hybrid (dense + BM25 + RRF) | 0.692 | 0.818 | 0.713 | 69 ms |
|
||||
| **bm25 + jina-rerank** | **0.920** | **0.939** | **0.927** | 490 ms (CPU) / ~50 ms (GPU) |
|
||||
| **dense** (nomic prefixes) | **0.955** | **0.966** | **0.985** | **0.972** |
|
||||
| hybrid (dense + BM25 + RRF) | 0.955 | 0.961 | 0.955 | 0.955 |
|
||||
| BM25 (SQLite FTS5) | 0.864 | 0.882 | 0.909 | 0.886 |
|
||||
| bm25 + jina-rerank | 0.773 | 0.827 | 0.848 | 0.816 |
|
||||
|
||||
HPE docs use controlled vocabulary, so lexical match dominates; the
|
||||
cross-encoder cleans up the long tail. See PLAN.md Phase 7/8 for the
|
||||
reasoning.
|
||||
`search_docs` defaults to dense. Rerank is opt-in (`RERANK_ENABLED=true`)
|
||||
— it now hurts. Full table: [`eval/results/post-prefix.md`](eval/results/post-prefix.md).
|
||||
|
||||
## Architecture
|
||||
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -32,16 +32,16 @@ services:
|
||||
# add that hostname here. "*" disables the rebind check entirely.
|
||||
MCP_ALLOWED_HOSTS: "hvm-docs-mcp,localhost,127.0.0.1"
|
||||
|
||||
# Phase 6 — reranker sidecar (jina-reranker-v2-base via llama.cpp).
|
||||
# Phase 6 — reranker sidecar is wired but off. Eval 2026-09-30
|
||||
# (post nomic prefixes): dense MRR=0.966 vs bm25+rerank 0.827.
|
||||
# Set RERANK_ENABLED=true to turn the sidecar back on.
|
||||
RERANK_URL: http://hvm-rerank:8080
|
||||
RERANK_POOL: "200"
|
||||
RERANK_TIMEOUT: "30"
|
||||
RERANK_ENABLED: "false"
|
||||
|
||||
# Phase 8 — hybrid retrieval (BM25 + dense + RRF).
|
||||
# Eval on the HVM corpus (eval/results/baseline.md, 2026-05-22) shows
|
||||
# BM25-default + reranker beats hybrid on every metric (MRR 0.920 vs
|
||||
# 0.875). Leaving HYBRID_SEARCH off so search_docs runs BM25-first +
|
||||
# reranker; dense is the fallback when BM25 finds nothing.
|
||||
# Phase 8 — hybrid retrieval (BM25 + dense + RRF). Off: dense is
|
||||
# the default (eval/results/post-prefix.md). BM25 is fallback only.
|
||||
HYBRID_SEARCH: "false"
|
||||
|
||||
# Phase 10 — usage telemetry.
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
"""Markdown formatters for MCP tool output.
|
||||
|
||||
No third-party imports — tests can run without mcp/httpx/chromadb.
|
||||
Citation numbers are per-call (1-based, dense) and stateless.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def format_search_hit(
|
||||
n: int,
|
||||
title: str,
|
||||
url: str,
|
||||
text: str,
|
||||
extra: str = "",
|
||||
) -> str:
|
||||
"""One numbered hit. `n` is 1-based in final reranked/fused order."""
|
||||
head = f"[{n}] **{title}**"
|
||||
if url:
|
||||
head += f" — {url}"
|
||||
parts = [head]
|
||||
if extra:
|
||||
parts.append(extra)
|
||||
body = (text or "").strip()
|
||||
if body:
|
||||
parts.append(body)
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def format_search_hits(
|
||||
hits: list[tuple[str, str, str, str]],
|
||||
) -> str:
|
||||
"""Render hits as `[1] **title** — url` then text.
|
||||
|
||||
`hits` is a list of (title, url, text, extra) in display order.
|
||||
Empty list → empty string (no invented citations).
|
||||
"""
|
||||
if not hits:
|
||||
return ""
|
||||
blocks = [
|
||||
format_search_hit(n, title, url, text, extra)
|
||||
for n, (title, url, text, extra) in enumerate(hits, start=1)
|
||||
]
|
||||
return "\n\n".join(blocks) + "\n"
|
||||
+66
-32
@@ -31,6 +31,7 @@ from mcp.server.mcpserver import MCPServer
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from pydantic import Field
|
||||
|
||||
from .format import format_search_hits
|
||||
from .usage import TimedCall
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -60,6 +61,9 @@ API_LESSONS_MD = Path(__file__).resolve().parent / "api_lessons.md"
|
||||
RERANK_URL = os.environ.get("RERANK_URL", "").rstrip("/") or None
|
||||
RERANK_POOL = int(os.environ.get("RERANK_POOL", "50"))
|
||||
RERANK_TIMEOUT = float(os.environ.get("RERANK_TIMEOUT", "30"))
|
||||
# Opt-in. Watchtower keeps the old container env (RERANK_URL is set in
|
||||
# live compose); default off so a code-only ship actually stops reranking.
|
||||
RERANK_ENABLED = os.environ.get("RERANK_ENABLED", "").lower() in ("true", "1", "yes", "on")
|
||||
|
||||
HYBRID_SEARCH = os.environ.get("HYBRID_SEARCH", "").lower() in ("true", "1", "yes", "on")
|
||||
RRF_K = int(os.environ.get("RRF_K", "60"))
|
||||
@@ -145,6 +149,17 @@ def _collection():
|
||||
return _CHROMA
|
||||
|
||||
|
||||
def _query_dense(col, query: str, n: int, where: dict | None = None, **extra):
|
||||
"""Dense query with nomic prefixes applied at embed time only."""
|
||||
from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts
|
||||
qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0]
|
||||
kwargs: dict[str, Any] = {"query_embeddings": [qvec], "n_results": n}
|
||||
if where:
|
||||
kwargs["where"] = where
|
||||
kwargs.update(extra)
|
||||
return col.query(**kwargs)
|
||||
|
||||
|
||||
def _bm25():
|
||||
"""Lazy BM25Index handle. None if the FTS5 db isn't built."""
|
||||
global _BM25
|
||||
@@ -258,9 +273,10 @@ def search_docs(
|
||||
"""Search the HPE Morpheus VM Essentials (HVM) docs corpus.
|
||||
|
||||
Returns the top-k most relevant chunks (with full source page URLs)
|
||||
given a natural-language query. Optional filters narrow the search
|
||||
to one version, one platform, or one bundle. Use list_versions()
|
||||
first if you need to discover the available facet values.
|
||||
given a natural-language query. Hits are numbered [1]… for citation.
|
||||
Optional filters narrow the search to one version, one platform, or
|
||||
one bundle. Use list_versions() first if you need to discover the
|
||||
available facet values.
|
||||
|
||||
Call this tool whenever the user asks anything that should be
|
||||
answerable from the official product documentation — install,
|
||||
@@ -282,11 +298,10 @@ def search_docs(
|
||||
bm25_where = _where_for_bm25(version, platform, bundle_id)
|
||||
pool = max(k * 5, 50)
|
||||
|
||||
# Retrieval mode selection. Eval on this corpus (2026-05-22, 22 golden
|
||||
# queries) showed BM25 MRR=0.88 vs dense MRR=0.54 vs hybrid MRR=0.69 —
|
||||
# HPE structured docs use controlled vocabulary, so lexical match wins.
|
||||
# Dense is kept as fallback when BM25 has no tokens to chew on (e.g.
|
||||
# purely stopword queries). HYBRID_SEARCH=true forces RRF fusion.
|
||||
# Retrieval mode. Eval 2026-09-30 (22 queries, post nomic prefixes +
|
||||
# heading-recursive chunking): dense MRR=0.966 P@1=0.955 vs
|
||||
# bm25+rerank MRR=0.827 P@1=0.773. Default is dense. HYBRID_SEARCH
|
||||
# still forces RRF. Rerank is opt-in (RERANK_ENABLED) — it now hurts.
|
||||
bm = _bm25()
|
||||
docs: list[str] = []
|
||||
metas: list[dict] = []
|
||||
@@ -296,7 +311,7 @@ def search_docs(
|
||||
|
||||
if HYBRID_SEARCH and bm is not None:
|
||||
try:
|
||||
dense_res = col.query(query_texts=[query], n_results=pool, where=where)
|
||||
dense_res = _query_dense(col, query, pool, where)
|
||||
dense_ids = (dense_res.get("ids") or [[]])[0]
|
||||
bm_hits = bm.query(query, n=pool, where=bm25_where)
|
||||
bm_ids = [cid for cid, _s in bm_hits]
|
||||
@@ -309,7 +324,18 @@ def search_docs(
|
||||
else "dense_only")
|
||||
retrieval_mode = "hybrid"
|
||||
except Exception as e:
|
||||
log.warning("hybrid failed, falling back to BM25→dense: %s", e)
|
||||
log.warning("hybrid failed, falling back to dense: %s", e)
|
||||
|
||||
if not docs:
|
||||
try:
|
||||
res = _query_dense(col, query, k, where)
|
||||
docs = (res.get("documents") or [[]])[0]
|
||||
metas = (res.get("metadatas") or [[]])[0]
|
||||
dists = (res.get("distances") or [[]])[0]
|
||||
retrieval_mode = "dense"
|
||||
top1_source = "dense_only"
|
||||
except Exception as e:
|
||||
log.warning("dense retrieval failed, falling back to BM25: %s", e)
|
||||
|
||||
if not docs and bm is not None:
|
||||
try:
|
||||
@@ -323,16 +349,10 @@ def search_docs(
|
||||
retrieval_mode = "bm25"
|
||||
top1_source = "bm25_only"
|
||||
except Exception as e:
|
||||
log.warning("BM25 retrieval failed, falling back to dense: %s", e)
|
||||
|
||||
if not docs:
|
||||
res = col.query(query_texts=[query], n_results=k, where=where)
|
||||
docs = (res.get("documents") or [[]])[0]
|
||||
metas = (res.get("metadatas") or [[]])[0]
|
||||
dists = (res.get("distances") or [[]])[0]
|
||||
log.warning("BM25 retrieval failed: %s", e)
|
||||
|
||||
reranker_fired = False
|
||||
if RERANK_URL and docs:
|
||||
if RERANK_URL and RERANK_ENABLED and docs:
|
||||
# Pull a deeper pool to give the reranker something to chew on.
|
||||
# We over-fetch up to RERANK_POOL chunks from whichever retriever
|
||||
# already won, then ask the reranker to pick the final top-k.
|
||||
@@ -342,7 +362,7 @@ def search_docs(
|
||||
extra = bm.query(query, n=pool_size, where=bm25_where) if bm else []
|
||||
extra_ids = [cid for cid, _s in extra]
|
||||
else:
|
||||
extra_res = col.query(query_texts=[query], n_results=pool_size, where=where)
|
||||
extra_res = _query_dense(col, query, pool_size, where)
|
||||
extra_ids = (extra_res.get("ids") or [[]])[0]
|
||||
if extra_ids:
|
||||
d2, m2, _ = _enrich_from_chroma(col, extra_ids, None)
|
||||
@@ -371,22 +391,21 @@ def search_docs(
|
||||
if not docs:
|
||||
return f"_No matches for `{query}`._"
|
||||
|
||||
out = [f"# {len(docs)} result(s) for `{query}`", ""]
|
||||
hits: list[tuple[str, str, str, str]] = []
|
||||
for doc, meta, dist in zip(docs, metas, dists):
|
||||
bid = meta.get("bundle_id", "")
|
||||
pid = meta.get("page_id", "")
|
||||
title = meta.get("title") or pid
|
||||
ver = meta.get("version") or ""
|
||||
url = _source_url(bid, pid)
|
||||
header = f"## {title}"
|
||||
extra = f"score={1 - dist:.3f}"
|
||||
if ver:
|
||||
header += f" _(v{ver})_"
|
||||
out.append(header)
|
||||
out.append(f"[{bid}/{pid}]({url}) · score={1 - dist:.3f}")
|
||||
out.append("")
|
||||
out.append(doc.strip())
|
||||
out.append("")
|
||||
return "\n".join(out)
|
||||
extra = f"v{ver} " + extra
|
||||
if pid:
|
||||
extra += f" page_id: `{pid}`"
|
||||
hits.append((title, url, doc, extra))
|
||||
header = f"# {len(docs)} result(s) for `{query}`\n\n"
|
||||
return header + format_search_hits(hits)
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
@@ -810,6 +829,21 @@ def weekly_digest(
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _pub_sort_key(pub: str) -> tuple[int, int]:
|
||||
"""Sortable (year, month) for an HPE "Published" value like "July 2026".
|
||||
|
||||
Plain string comparison ranks "March 2026" above "July 2026".
|
||||
Unparseable values sort last.
|
||||
"""
|
||||
for fmt in ("%B %Y", "%b %Y", "%B %d, %Y", "%Y-%m-%d"):
|
||||
try:
|
||||
d = _dt.datetime.strptime(pub.strip(), fmt)
|
||||
return (d.year, d.month)
|
||||
except ValueError:
|
||||
continue
|
||||
return (0, 0)
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def corpus_status() -> str:
|
||||
"""Freshness + size of the knowledge base.
|
||||
@@ -831,7 +865,7 @@ def corpus_status() -> str:
|
||||
for slug, b in cat.items():
|
||||
pub = (b.get("dates") or {}).get("Published")
|
||||
if pub:
|
||||
if latest_pub is None or pub > latest_pub:
|
||||
if latest_pub is None or _pub_sort_key(pub) > _pub_sort_key(latest_pub):
|
||||
latest_pub = pub
|
||||
per_bundle.append((slug, pub))
|
||||
if latest_pub:
|
||||
@@ -850,7 +884,7 @@ def corpus_status() -> str:
|
||||
"",
|
||||
]
|
||||
if per_bundle:
|
||||
per_bundle.sort(key=lambda kv: kv[1], reverse=True)
|
||||
per_bundle.sort(key=lambda kv: _pub_sort_key(kv[1]), reverse=True)
|
||||
lines.append("## Most-recently-edited bundles (by HPE)")
|
||||
for slug, when in per_bundle[:5]:
|
||||
b = cat.get(slug, {})
|
||||
@@ -1058,8 +1092,8 @@ def find_doc_inconsistencies(
|
||||
return f"Couldn't open Chroma collection: {e}"
|
||||
where = _build_where(version, platform, bundle_id)
|
||||
try:
|
||||
res = col.query(query_texts=[scope_query], n_results=max_pages * 3,
|
||||
where=where, include=["metadatas"])
|
||||
res = _query_dense(col, scope_query, max_pages * 3, where,
|
||||
include=["metadatas"])
|
||||
except Exception as e:
|
||||
_call.set(error=f"query: {e}")
|
||||
return f"Scope query failed: {e}"
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
"""Citation markdown renderer — no Chroma, no mcp."""
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from docs_mcp.format import format_search_hits
|
||||
|
||||
|
||||
class FormatSearchHitsTests(unittest.TestCase):
|
||||
def test_two_hits_numbered_with_urls(self) -> None:
|
||||
out = format_search_hits([
|
||||
("Install on Linux", "https://docs.example.com/install", "chunk one", ""),
|
||||
("HA setup", "https://docs.example.com/ha", "chunk two", ""),
|
||||
])
|
||||
self.assertIn("[1] **Install on Linux** — https://docs.example.com/install", out)
|
||||
self.assertIn("[2] **HA setup** — https://docs.example.com/ha", out)
|
||||
self.assertIn("chunk one", out)
|
||||
self.assertIn("chunk two", out)
|
||||
self.assertNotIn("[3]", out)
|
||||
# [1] before [2]
|
||||
self.assertLess(out.index("[1]"), out.index("[2]"))
|
||||
# each number is followed by its URL on the same logical hit
|
||||
first = out.split("[2]")[0]
|
||||
self.assertIn("https://docs.example.com/install", first)
|
||||
self.assertNotIn("https://docs.example.com/ha", first)
|
||||
|
||||
def test_empty_hits_invents_nothing(self) -> None:
|
||||
self.assertEqual(format_search_hits([]), "")
|
||||
|
||||
def test_missing_url_omits_emdash(self) -> None:
|
||||
out = format_search_hits([("Title only", "", "body", "")])
|
||||
self.assertIn("[1] **Title only**", out)
|
||||
self.assertNotIn("—", out)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+156
@@ -0,0 +1,156 @@
|
||||
"""Paired permutation test between two eval JSONL sidecars.
|
||||
|
||||
Compares per-query scores from two `eval.run_eval` sidecar files so
|
||||
"P@1 went 0.88 → 0.91 on 25 queries" is not treated as a win.
|
||||
|
||||
python -m eval.pvalue \\
|
||||
--a eval/results/baseline.jsonl \\
|
||||
--b eval/results/new.jsonl \\
|
||||
--metric rr
|
||||
|
||||
Exit 0 even when the difference is not significant — this is a report,
|
||||
not a gate. No third-party deps; `random.Random(seed)` is enough.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import random
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def load_sidecar(path: Path) -> list[dict]:
|
||||
rows: list[dict] = []
|
||||
with open(path) as fh:
|
||||
for line in fh:
|
||||
line = line.strip()
|
||||
if line:
|
||||
rows.append(json.loads(line))
|
||||
return rows
|
||||
|
||||
|
||||
def paired_permutation(
|
||||
a: list[float],
|
||||
b: list[float],
|
||||
n_resamples: int = 10000,
|
||||
seed: int = 0,
|
||||
) -> dict:
|
||||
"""Two-sided paired permutation test on per-query scores.
|
||||
|
||||
Null: each pair is exchangeable (randomly flipping the sign of
|
||||
A_i - B_i). p_value is the fraction of permutations whose
|
||||
|mean diff| is at least as large as the observed |mean(A-B)|.
|
||||
"""
|
||||
if len(a) != len(b):
|
||||
raise ValueError(f"paired lengths differ: {len(a)} vs {len(b)}")
|
||||
if not a:
|
||||
raise ValueError("no paired queries to compare")
|
||||
diffs = [x - y for x, y in zip(a, b)]
|
||||
n = len(diffs)
|
||||
observed = sum(diffs) / n
|
||||
abs_obs = abs(observed)
|
||||
rng = random.Random(seed)
|
||||
extreme = 0
|
||||
for _ in range(n_resamples):
|
||||
total = 0.0
|
||||
for d in diffs:
|
||||
total += d if rng.random() < 0.5 else -d
|
||||
if abs(total / n) >= abs_obs - 1e-15:
|
||||
extreme += 1
|
||||
p_value = extreme / n_resamples
|
||||
return {
|
||||
"A_mean": sum(a) / n,
|
||||
"B_mean": sum(b) / n,
|
||||
"Diff(A-B)": observed,
|
||||
"p_value": p_value,
|
||||
"significant": p_value < 0.05,
|
||||
"n": n,
|
||||
"n_resamples": n_resamples,
|
||||
}
|
||||
|
||||
|
||||
def _index(rows: list[dict], metric: str) -> dict[tuple[str, str], float]:
|
||||
"""Map (retriever, query) -> score."""
|
||||
out: dict[tuple[str, str], float] = {}
|
||||
for row in rows:
|
||||
retriever = str(row.get("retriever") or "")
|
||||
query = str(row.get("query") or "")
|
||||
if metric == "p_at_1":
|
||||
score = float(row.get("p_at_1") or 0)
|
||||
else:
|
||||
score = float(row.get("rr") or 0)
|
||||
out[(retriever, query)] = score
|
||||
return out
|
||||
|
||||
|
||||
def compare(
|
||||
rows_a: list[dict],
|
||||
rows_b: list[dict],
|
||||
metric: str = "rr",
|
||||
retriever: str | None = None,
|
||||
n_resamples: int = 10000,
|
||||
seed: int = 0,
|
||||
) -> list[dict]:
|
||||
"""Join on (retriever, query). One result dict per shared retriever."""
|
||||
ia, ib = _index(rows_a, metric), _index(rows_b, metric)
|
||||
retrievers = sorted({r for r, _ in ia} & {r for r, _ in ib})
|
||||
if retriever:
|
||||
retrievers = [r for r in retrievers if r == retriever]
|
||||
if not retrievers:
|
||||
raise ValueError(f"retriever {retriever!r} not in both sidecars")
|
||||
reports = []
|
||||
for name in retrievers:
|
||||
queries = sorted({q for r, q in ia if r == name} & {q for r, q in ib if r == name})
|
||||
if not queries:
|
||||
continue
|
||||
a_scores = [ia[(name, q)] for q in queries]
|
||||
b_scores = [ib[(name, q)] for q in queries]
|
||||
report = paired_permutation(a_scores, b_scores, n_resamples=n_resamples, seed=seed)
|
||||
report["retriever"] = name
|
||||
report["metric"] = metric
|
||||
reports.append(report)
|
||||
if not reports:
|
||||
raise ValueError("no overlapping (retriever, query) pairs")
|
||||
return reports
|
||||
|
||||
|
||||
def render(reports: list[dict]) -> str:
|
||||
lines = ["# Permutation test", ""]
|
||||
for r in reports:
|
||||
sig = "yes" if r["significant"] else "no"
|
||||
lines += [
|
||||
f"## `{r['retriever']}` ({r['metric']}, n={r['n']})",
|
||||
"",
|
||||
f"- A_mean: `{r['A_mean']:.4f}`",
|
||||
f"- B_mean: `{r['B_mean']:.4f}`",
|
||||
f"- Diff(A-B): `{r['Diff(A-B)']:.4f}`",
|
||||
f"- p_value: `{r['p_value']:.4f}` ({r['n_resamples']} resamples)",
|
||||
f"- significant (p < 0.05): **{sig}**",
|
||||
"",
|
||||
]
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
p = argparse.ArgumentParser(description="Paired permutation test on two eval JSONL sidecars.")
|
||||
p.add_argument("--a", type=Path, required=True, help="sidecar JSONL (system A)")
|
||||
p.add_argument("--b", type=Path, required=True, help="sidecar JSONL (system B)")
|
||||
p.add_argument("--metric", choices=("rr", "p_at_1"), default="rr")
|
||||
p.add_argument("--retriever", default=None, help="restrict to one retriever name")
|
||||
p.add_argument("--n-resamples", type=int, default=10000)
|
||||
p.add_argument("--seed", type=int, default=0)
|
||||
args = p.parse_args()
|
||||
reports = compare(
|
||||
load_sidecar(args.a),
|
||||
load_sidecar(args.b),
|
||||
metric=args.metric,
|
||||
retriever=args.retriever,
|
||||
n_resamples=args.n_resamples,
|
||||
seed=args.seed,
|
||||
)
|
||||
print(render(reports), end="")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,21 @@
|
||||
# Retrieval eval — k=5
|
||||
|
||||
_22 hand-curated queries, generated 2026-09-30 02:40:07_
|
||||
|
||||
Live container after nomic prefixes + heading-recursive chunking
|
||||
(`hvm-docs` image `b200c5f76c56`, Watchtower 2026-09-30T02:36Z).
|
||||
Compared with [`baseline.md`](baseline.md) (2026-05-22, pre-prefix).
|
||||
|
||||
| Retriever | P@1 | MRR | Recall@5 | nDCG@5 | avg latency |
|
||||
| --- | ---: | ---: | ---: | ---: | ---: |
|
||||
| `dense` | 0.955 | 0.966 | 0.985 | 0.972 | 183ms |
|
||||
| `bm25` | 0.864 | 0.882 | 0.909 | 0.886 | 8ms |
|
||||
| `hybrid_rrf` | 0.955 | 0.961 | 0.955 | 0.955 | 113ms |
|
||||
| `bm25+rerank` | 0.773 | 0.827 | 0.848 | 0.816 | 162ms |
|
||||
| `hybrid_rrf+rerank` | 0.727 | 0.813 | 0.894 | 0.821 | 254ms |
|
||||
|
||||
Dense flipped from worst (May MRR 0.539) to best. Rerank now hurts.
|
||||
`search_docs` therefore defaults to dense; rerank is opt-in via
|
||||
`RERANK_ENABLED`.
|
||||
|
||||
Dense P@1 miss: `create a user account` (MRR 0.250).
|
||||
+3
-1
@@ -51,7 +51,9 @@ class DenseRetriever:
|
||||
self.pool = pool
|
||||
|
||||
def retrieve(self, query: str, k: int = 10) -> list[tuple[str, str]]:
|
||||
res = self.col.query(query_texts=[query], n_results=self.pool)
|
||||
from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts
|
||||
qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0]
|
||||
res = self.col.query(query_embeddings=[qvec], n_results=self.pool)
|
||||
ids = (res.get("ids") or [[]])[0]
|
||||
return _collapse_to_pages(ids, k)
|
||||
|
||||
|
||||
+37
-7
@@ -34,6 +34,12 @@ def load_queries(path: Path) -> list[dict]:
|
||||
return [json.loads(line) for line in fh if line.strip()]
|
||||
|
||||
|
||||
def p_at_1(retrieved: list[tuple[str, str]], expected: list[tuple[str, str]]) -> float:
|
||||
if not retrieved or not expected:
|
||||
return 0.0
|
||||
return 1.0 if retrieved[0] in set(expected) else 0.0
|
||||
|
||||
|
||||
def reciprocal_rank(retrieved: list[tuple[str, str]], expected: list[tuple[str, str]]) -> float:
|
||||
expected_set = set(expected)
|
||||
for i, page in enumerate(retrieved, start=1):
|
||||
@@ -65,8 +71,12 @@ def main() -> int:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--queries", type=Path, default=Path("eval/queries.jsonl"))
|
||||
p.add_argument("--k", type=int, default=5)
|
||||
p.add_argument("--ks", default="1,5,10,20", help="comma-separated k-curve")
|
||||
p.add_argument("--output", type=Path, default=Path("eval/results/baseline.md"))
|
||||
p.add_argument("--compare", type=Path, default=None)
|
||||
args = p.parse_args()
|
||||
ks = sorted({int(x) for x in args.ks.split(",") if x.strip()}) or [args.k]
|
||||
max_k = max(ks + [args.k])
|
||||
|
||||
if not args.queries.exists():
|
||||
print(f"queries file not found: {args.queries}")
|
||||
@@ -109,34 +119,39 @@ def main() -> int:
|
||||
rows: dict[str, dict[str, float]] = {}
|
||||
per_query: list[dict] = []
|
||||
for r in retrievers:
|
||||
mrr_sum = recall_sum = ndcg_sum = 0.0
|
||||
mrr_sum = recall_sum = ndcg_sum = p1_sum = 0.0
|
||||
elapsed_sum = 0.0
|
||||
for q in queries:
|
||||
expected = [(e["bundle_id"], e["page_id"]) for e in q["expected"]]
|
||||
t0 = time.time()
|
||||
retrieved = r.retrieve(q["query"], k=max(args.k, 10))
|
||||
retrieved = r.retrieve(q["query"], k=max(max_k, 10))
|
||||
elapsed = time.time() - t0
|
||||
mrr = reciprocal_rank(retrieved, expected)
|
||||
p1 = p_at_1(retrieved, expected)
|
||||
recall = recall_at_k(retrieved, expected, args.k)
|
||||
ndcg = ndcg_at_k(retrieved, expected, args.k)
|
||||
mrr_sum += mrr
|
||||
p1_sum += p1
|
||||
recall_sum += recall
|
||||
ndcg_sum += ndcg
|
||||
elapsed_sum += elapsed
|
||||
per_query.append({
|
||||
"retriever": r.name, "query": q["query"],
|
||||
"mrr": mrr, "recall@k": recall, "ndcg@k": ndcg,
|
||||
"mrr": mrr, "p_at_1": int(p1), "recall@k": recall, "ndcg@k": ndcg,
|
||||
"top1": list(retrieved[0]) if retrieved else None,
|
||||
"ranked": [list(p) for p in retrieved],
|
||||
"elapsed_s": round(elapsed, 3),
|
||||
})
|
||||
n = len(queries)
|
||||
rows[r.name] = {
|
||||
"P@1": p1_sum / n,
|
||||
"MRR": mrr_sum / n,
|
||||
f"Recall@{args.k}": recall_sum / n,
|
||||
f"nDCG@{args.k}": ndcg_sum / n,
|
||||
"avg_latency_s": elapsed_sum / n,
|
||||
}
|
||||
print(f" {r.name}: MRR={rows[r.name]['MRR']:.3f} "
|
||||
print(f" {r.name}: P@1={rows[r.name]['P@1']:.3f} "
|
||||
f"MRR={rows[r.name]['MRR']:.3f} "
|
||||
f"Recall@{args.k}={rows[r.name][f'Recall@{args.k}']:.3f} "
|
||||
f"nDCG@{args.k}={rows[r.name][f'nDCG@{args.k}']:.3f} "
|
||||
f"avg={rows[r.name]['avg_latency_s']*1000:.0f}ms")
|
||||
@@ -144,10 +159,10 @@ def main() -> int:
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
md = [f"# Retrieval eval — k={args.k}", "",
|
||||
f"_{len(queries)} hand-curated queries, generated {time.strftime('%Y-%m-%d %H:%M:%S')}_", "",
|
||||
"| Retriever | MRR | Recall@{k} | nDCG@{k} | avg latency |".replace("{k}", str(args.k)),
|
||||
"| --- | ---: | ---: | ---: | ---: |"]
|
||||
"| Retriever | P@1 | MRR | Recall@{k} | nDCG@{k} | avg latency |".replace("{k}", str(args.k)),
|
||||
"| --- | ---: | ---: | ---: | ---: | ---: |"]
|
||||
for name, m in rows.items():
|
||||
md.append(f"| `{name}` | {m['MRR']:.3f} | {m[f'Recall@{args.k}']:.3f} "
|
||||
md.append(f"| `{name}` | {m['P@1']:.3f} | {m['MRR']:.3f} | {m[f'Recall@{args.k}']:.3f} "
|
||||
f"| {m[f'nDCG@{args.k}']:.3f} | {m['avg_latency_s']*1000:.0f}ms |")
|
||||
md += ["", "## Per-query results", "",
|
||||
"| Retriever | Query | MRR | top-1 |", "| --- | --- | ---: | --- |"]
|
||||
@@ -155,7 +170,22 @@ def main() -> int:
|
||||
top1 = f"`{r['top1'][0]}/{r['top1'][1][:24]}...`" if r["top1"] else "—"
|
||||
md.append(f"| `{r['retriever']}` | {r['query'][:60]} | {r['mrr']:.3f} | {top1} |")
|
||||
args.output.write_text("\n".join(md) + "\n")
|
||||
sidecar = args.output.with_suffix(".jsonl")
|
||||
with open(sidecar, "w") as fh:
|
||||
for r in per_query:
|
||||
fh.write(json.dumps({
|
||||
"query": r["query"],
|
||||
"retriever": r["retriever"],
|
||||
"ranked": r.get("ranked") or [],
|
||||
"rr": r["mrr"],
|
||||
"p_at_1": r["p_at_1"],
|
||||
}) + "\n")
|
||||
print(f"wrote {args.output}")
|
||||
print(f"wrote {sidecar}")
|
||||
if args.compare:
|
||||
from eval.pvalue import compare, load_sidecar, render
|
||||
print()
|
||||
print(render(compare(load_sidecar(sidecar), load_sidecar(args.compare))), end="")
|
||||
return 0
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
"""Stdlib tests for eval metrics + the permutation test.
|
||||
|
||||
Must not open Chroma — the template has no corpus. Run with:
|
||||
|
||||
python -m unittest eval.test_metrics
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import unittest
|
||||
|
||||
from eval.pvalue import compare, paired_permutation
|
||||
from eval.run_eval import ndcg_at_k, p_at_1, recall_at_k, reciprocal_rank
|
||||
|
||||
|
||||
A, B, C, X, Y = ("b", "a"), ("b", "b"), ("b", "c"), ("b", "x"), ("b", "y")
|
||||
|
||||
|
||||
class MetricTests(unittest.TestCase):
|
||||
def test_reciprocal_rank(self) -> None:
|
||||
# Q1: expected at rank 1
|
||||
self.assertEqual(reciprocal_rank([A, B, C], [A]), 1.0)
|
||||
# Q2: expected at rank 2
|
||||
self.assertEqual(reciprocal_rank([B, A], [A]), 0.5)
|
||||
# Q3: miss
|
||||
self.assertEqual(reciprocal_rank([X, Y], [A]), 0.0)
|
||||
|
||||
def test_p_at_1(self) -> None:
|
||||
self.assertEqual(p_at_1([A, B], [A]), 1.0)
|
||||
self.assertEqual(p_at_1([B, A], [A]), 0.0)
|
||||
self.assertEqual(p_at_1([], [A]), 0.0)
|
||||
self.assertEqual(p_at_1([A], []), 0.0)
|
||||
|
||||
def test_recall_at_k(self) -> None:
|
||||
self.assertEqual(recall_at_k([A, B, C], [A], 1), 1.0)
|
||||
self.assertEqual(recall_at_k([B, A], [A], 1), 0.0)
|
||||
self.assertEqual(recall_at_k([B, A], [A], 2), 1.0)
|
||||
self.assertEqual(recall_at_k([X, Y], [A], 5), 0.0)
|
||||
self.assertEqual(recall_at_k([A, B], [A, C], 1), 0.5)
|
||||
|
||||
def test_ndcg_at_k(self) -> None:
|
||||
self.assertEqual(ndcg_at_k([A], [A], 1), 1.0)
|
||||
# expected at rank 2: dcg = 1/log2(3), idcg = 1
|
||||
self.assertAlmostEqual(
|
||||
ndcg_at_k([B, A], [A], 2),
|
||||
(1.0 / math.log2(3)) / 1.0,
|
||||
)
|
||||
self.assertEqual(ndcg_at_k([X, Y], [A], 5), 0.0)
|
||||
|
||||
|
||||
class PermutationTests(unittest.TestCase):
|
||||
def test_identical_lists_not_significant(self) -> None:
|
||||
scores = [0.5, 1.0, 0.0, 1.0, 0.5]
|
||||
report = paired_permutation(scores, list(scores), n_resamples=2000, seed=0)
|
||||
self.assertEqual(report["Diff(A-B)"], 0.0)
|
||||
self.assertEqual(report["p_value"], 1.0)
|
||||
self.assertFalse(report["significant"])
|
||||
|
||||
def test_large_paired_difference_is_significant(self) -> None:
|
||||
a = [1.0] * 20
|
||||
b = [0.0] * 20
|
||||
report = paired_permutation(a, b, n_resamples=5000, seed=0)
|
||||
self.assertGreater(report["Diff(A-B)"], 0.9)
|
||||
self.assertLess(report["p_value"], 0.05)
|
||||
self.assertTrue(report["significant"])
|
||||
|
||||
def test_compare_joins_on_retriever_and_query(self) -> None:
|
||||
rows_a = (
|
||||
[{"retriever": "dense", "query": f"q{i}", "rr": 1.0, "p_at_1": 1} for i in range(20)]
|
||||
+ [{"retriever": "bm25", "query": "q0", "rr": 0.0, "p_at_1": 0}]
|
||||
)
|
||||
rows_b = (
|
||||
[{"retriever": "dense", "query": f"q{i}", "rr": 0.0, "p_at_1": 0} for i in range(20)]
|
||||
+ [{"retriever": "bm25", "query": "q0", "rr": 0.0, "p_at_1": 0}]
|
||||
)
|
||||
reports = compare(rows_a, rows_b, metric="rr", n_resamples=2000, seed=0)
|
||||
by_name = {r["retriever"]: r for r in reports}
|
||||
self.assertIn("dense", by_name)
|
||||
self.assertTrue(by_name["dense"]["significant"])
|
||||
self.assertFalse(by_name["bm25"]["significant"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+114
@@ -0,0 +1,114 @@
|
||||
"""Page-level miss dump using this clone's Dense/BM25 retrievers.
|
||||
|
||||
python -m eval.trace --queries eval/queries.jsonl
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from eval.run_eval import load_queries, p_at_1
|
||||
|
||||
|
||||
def classify_top1_source(top1, dense_pages, bm25_pages) -> str:
|
||||
if top1 is None:
|
||||
return "neither"
|
||||
in_d, in_b = top1 in dense_pages, top1 in bm25_pages
|
||||
if in_d and in_b:
|
||||
return "both"
|
||||
if in_d:
|
||||
return "dense_only"
|
||||
if in_b:
|
||||
return "bm25_only"
|
||||
return "neither"
|
||||
|
||||
|
||||
def first_ranks(pages: list[tuple[str, str]]) -> dict[str, int]:
|
||||
out: dict[str, int] = {}
|
||||
for i, (bid, pid) in enumerate(pages, start=1):
|
||||
key = f"{bid}/{pid}"
|
||||
if key not in out:
|
||||
out[key] = i
|
||||
return out
|
||||
|
||||
|
||||
def render_misses(rows: list[dict]) -> str:
|
||||
misses = [r for r in rows if not r.get("hit")]
|
||||
if not misses:
|
||||
return "# Eval misses\n\n_(none)_\n"
|
||||
lines = [f"# Eval misses ({len(misses)})", ""]
|
||||
for row in misses:
|
||||
lines += [
|
||||
f"## {row['query']}",
|
||||
"",
|
||||
f"- expected: `{row['expected']}`",
|
||||
f"- top-5: `{row['ranked_pages'][:5]}`",
|
||||
f"- top1_source: `{row['top1_source']}`",
|
||||
"",
|
||||
]
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--queries", type=Path, default=Path("eval/queries.jsonl"))
|
||||
p.add_argument("--trace-out", type=Path, default=Path("eval/results/trace.jsonl"))
|
||||
p.add_argument("--misses-out", type=Path, default=Path("eval/results/misses.md"))
|
||||
args = p.parse_args()
|
||||
if not args.queries.exists():
|
||||
print(f"queries file not found: {args.queries}")
|
||||
return 1
|
||||
try:
|
||||
import os
|
||||
import chromadb
|
||||
from chromadb.config import Settings
|
||||
from rag.embeddings import embedding_function
|
||||
from rag.bm25 import BM25Index
|
||||
from eval.retrievers import BM25Retriever, DenseRetriever
|
||||
product = os.environ.get("PRODUCT_NAME", "hvm")
|
||||
root = Path(__file__).resolve().parent.parent
|
||||
col = chromadb.PersistentClient(
|
||||
path=str(root / "chroma"),
|
||||
settings=Settings(anonymized_telemetry=False),
|
||||
).get_collection(f"{product}_docs", embedding_function=embedding_function())
|
||||
bm = BM25Index(str(root / "bm25" / f"{product}_docs.db"))
|
||||
dense_r, bm25_r = DenseRetriever(col), BM25Retriever(bm)
|
||||
except Exception as e:
|
||||
args.trace_out.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.trace_out.write_text("")
|
||||
args.misses_out.write_text("# Eval misses\n\nno index\n")
|
||||
print(f"no index ({e}); wrote empty trace")
|
||||
return 0
|
||||
|
||||
rows = []
|
||||
for q in load_queries(args.queries):
|
||||
expected = [(e["bundle_id"], e["page_id"]) for e in q["expected"]]
|
||||
dense_pages = dense_r.retrieve(q["query"], k=50)
|
||||
bm25_pages = bm25_r.retrieve(q["query"], k=50)
|
||||
ranked = dense_pages or bm25_pages # HVM default retrieval is dense-first
|
||||
top1 = ranked[0] if ranked else None
|
||||
p1 = p_at_1(ranked, expected)
|
||||
rows.append({
|
||||
"query": q["query"],
|
||||
"expected": [list(p) for p in expected],
|
||||
"hit": bool(p1),
|
||||
"p_at_1": int(p1),
|
||||
"dense_rank": first_ranks(dense_pages),
|
||||
"bm25_rank": first_ranks(bm25_pages),
|
||||
"top1": list(top1) if top1 else None,
|
||||
"top1_source": classify_top1_source(top1, set(dense_pages), set(bm25_pages)),
|
||||
"ranked_pages": [list(p) for p in ranked],
|
||||
})
|
||||
args.trace_out.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(args.trace_out, "w") as fh:
|
||||
for row in rows:
|
||||
fh.write(json.dumps(row) + "\n")
|
||||
args.misses_out.write_text(render_misses(rows))
|
||||
print(f"wrote {args.trace_out} ({len(rows)} queries)")
|
||||
print(f"wrote {args.misses_out}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
+168
-47
@@ -1,24 +1,12 @@
|
||||
"""Markdown chunker — paragraph-aware, ~400-600 token target.
|
||||
"""Markdown chunker — heading-recursive, ~400-600 token target.
|
||||
|
||||
Adjust the chunking strategy per product if your page format differs
|
||||
significantly from prose. The output shape (id, text, metadata) is
|
||||
fixed by the downstream Chroma + BM25 indexing in rag/index.py — don't
|
||||
change that.
|
||||
Chunk by semantic section (ATX headings), not raw page/length. A
|
||||
synthetic chunk 0 (title + first paragraph + optional keyword bag) is
|
||||
always emitted first — dense retrieval lands on it. Do not drop it.
|
||||
|
||||
The key knob you'll tune per product is chunk-0. Dense retrieval lands
|
||||
on chunk 0 first for most queries. Make it a synthetic chunk built
|
||||
from:
|
||||
|
||||
- the page title (as natural-language H1)
|
||||
- a 1-sentence task description (you'll have to generate this — for
|
||||
pages that already have a "## Overview" or "## Introduction" the
|
||||
first sentence usually works)
|
||||
- a keyword bag of important terms (filenames, API names, error
|
||||
codes — the rare technical tokens that BM25 lights up on)
|
||||
|
||||
Without a rich chunk 0, dense retrieval gets dominated by the much
|
||||
larger prose body, and short pages (script examples, reference cards)
|
||||
get buried.
|
||||
No chonkie / langchain dependency. The output shape (id, text, metadata)
|
||||
is fixed by rag/index.py — don't change that. `heading_path` is optional
|
||||
metadata (e.g. "Install > Linux") and is not required for indexing.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -26,17 +14,12 @@ import re
|
||||
from typing import Iterator
|
||||
|
||||
|
||||
# Approximate token estimate from char count. Tunable — set per
|
||||
# embedder if the default 4 chars/token is wrong.
|
||||
CHARS_PER_TOKEN = 4
|
||||
TARGET_TOKENS = 500
|
||||
TARGET_CHARS = TARGET_TOKENS * CHARS_PER_TOKEN
|
||||
# Hard cap: nomic-embed-text's context is 2048 tokens. Anything larger
|
||||
# 400s the entire embed batch. 6000 chars works for prose but markdown
|
||||
# tables with lots of `|` separators tokenize ~1.4× denser; a 5839-char
|
||||
# table chunk from the HVM qualification matrix tokenized past 2048 and
|
||||
# crashed the rebuild. 4000 chars stays under 2048 tokens even for
|
||||
# dense table content while leaving headroom for the query side.
|
||||
# nomic-embed-text context is 2048 tokens. Markdown tables with lots of
|
||||
# `|` tokenize ~1.4× denser than prose; 4000 chars stays under 2048 even
|
||||
# for qualification-matrix chunks (a 5839-char table crashed a rebuild).
|
||||
MAX_CHARS = 4000
|
||||
|
||||
|
||||
@@ -97,6 +80,158 @@ def split_paragraphs(md: str) -> list[str]:
|
||||
return [b for b in blocks if b]
|
||||
|
||||
|
||||
def _heading_level(block: str) -> int:
|
||||
first = block.lstrip().split("\n", 1)[0].strip()
|
||||
m = re.match(r"^(#{1,6})\s+\S", first)
|
||||
return len(m.group(1)) if m else 0
|
||||
|
||||
|
||||
def _heading_title(block: str) -> str:
|
||||
first = block.split("\n", 1)[0].strip()
|
||||
return re.sub(r"^#{1,6}\s+", "", first).strip()
|
||||
|
||||
|
||||
def _is_fence(block: str) -> bool:
|
||||
return block.lstrip().startswith("```")
|
||||
|
||||
|
||||
def _joined_len(blocks: list[str]) -> int:
|
||||
if not blocks:
|
||||
return 0
|
||||
return sum(len(b) for b in blocks) + 2 * (len(blocks) - 1)
|
||||
|
||||
|
||||
def _split_len(text: str) -> list[str]:
|
||||
size = TARGET_CHARS
|
||||
return [text[i:i + size] for i in range(0, len(text), size)] or [text]
|
||||
|
||||
|
||||
def _hard_wrap(text: str) -> list[str]:
|
||||
"""Last-resort split. Never used on fenced code blocks."""
|
||||
if len(text) <= TARGET_CHARS:
|
||||
return [text]
|
||||
parts = re.split(r"\n\s*\n", text)
|
||||
if len(parts) <= 1:
|
||||
return _split_len(text)
|
||||
packed: list[str] = []
|
||||
buf: list[str] = []
|
||||
n = 0
|
||||
for part in parts:
|
||||
if len(part) > TARGET_CHARS:
|
||||
if buf:
|
||||
packed.append("\n\n".join(buf))
|
||||
buf, n = [], 0
|
||||
packed.extend(_split_len(part))
|
||||
continue
|
||||
if n + len(part) > TARGET_CHARS and buf:
|
||||
packed.append("\n\n".join(buf))
|
||||
buf, n = [], 0
|
||||
buf.append(part)
|
||||
n += len(part)
|
||||
if buf:
|
||||
packed.append("\n\n".join(buf))
|
||||
return packed
|
||||
|
||||
|
||||
def _pack_paragraphs(blocks: list[str], path: str) -> list[tuple[str, str]]:
|
||||
out: list[tuple[str, str]] = []
|
||||
buf: list[str] = []
|
||||
buf_chars = 0
|
||||
|
||||
def flush() -> None:
|
||||
nonlocal buf, buf_chars
|
||||
if buf:
|
||||
out.append(("\n\n".join(buf), path))
|
||||
buf, buf_chars = [], 0
|
||||
|
||||
for p in blocks:
|
||||
if _is_fence(p) and len(p) > TARGET_CHARS:
|
||||
flush()
|
||||
out.append((p, path)) # never slice a fence
|
||||
continue
|
||||
if len(p) > TARGET_CHARS:
|
||||
flush()
|
||||
out.extend((piece, path) for piece in _hard_wrap(p))
|
||||
continue
|
||||
if buf_chars + len(p) > TARGET_CHARS and buf:
|
||||
flush()
|
||||
buf.append(p)
|
||||
buf_chars += len(p)
|
||||
flush()
|
||||
return out
|
||||
|
||||
|
||||
def _split_at_level(blocks: list[str], level: int) -> list[tuple[list[str], list[str]]]:
|
||||
"""Split into (path_suffix, blocks) groups starting at `level` headings.
|
||||
|
||||
path_suffix is [] for preamble, [title] for a section headed at `level`.
|
||||
"""
|
||||
groups: list[tuple[list[str], list[str]]] = []
|
||||
preamble: list[str] = []
|
||||
current_path: list[str] = []
|
||||
current: list[str] = []
|
||||
for b in blocks:
|
||||
lv = _heading_level(b)
|
||||
if lv == level:
|
||||
if current:
|
||||
groups.append((current_path, current))
|
||||
elif preamble:
|
||||
groups.append(([], preamble))
|
||||
preamble = []
|
||||
current_path = [_heading_title(b)]
|
||||
current = [b]
|
||||
elif current:
|
||||
current.append(b)
|
||||
else:
|
||||
preamble.append(b)
|
||||
if current:
|
||||
if preamble and not groups:
|
||||
groups.append(([], preamble))
|
||||
preamble = []
|
||||
groups.append((current_path, current))
|
||||
elif preamble:
|
||||
groups.append(([], preamble))
|
||||
return groups
|
||||
|
||||
|
||||
def _pack_recursive(blocks: list[str], path: list[str]) -> list[tuple[str, str]]:
|
||||
if not blocks:
|
||||
return []
|
||||
path_s = " > ".join(path)
|
||||
if _joined_len(blocks) <= TARGET_CHARS:
|
||||
return [("\n\n".join(blocks), path_s)]
|
||||
|
||||
levels = [_heading_level(b) for b in blocks]
|
||||
heading_levels = [lv for lv in levels if lv > 0]
|
||||
if not heading_levels:
|
||||
return _pack_paragraphs(blocks, path_s)
|
||||
|
||||
split_lv = min(heading_levels)
|
||||
groups = _split_at_level(blocks, split_lv)
|
||||
if len(groups) <= 1:
|
||||
# Can't split at this heading level — descend or pack paragraphs.
|
||||
if levels[0] == split_lv:
|
||||
title = _heading_title(blocks[0])
|
||||
rest = blocks[1:]
|
||||
if not rest:
|
||||
return _pack_paragraphs(blocks, path_s)
|
||||
packed = _pack_recursive(rest, path + [title])
|
||||
if not packed:
|
||||
return [(blocks[0], " > ".join(path + [title]))]
|
||||
glued = blocks[0] + "\n\n" + packed[0][0]
|
||||
if len(glued) <= TARGET_CHARS:
|
||||
packed[0] = (glued, packed[0][1])
|
||||
else:
|
||||
packed.insert(0, (blocks[0], packed[0][1]))
|
||||
return packed
|
||||
return _pack_paragraphs(blocks, path_s)
|
||||
|
||||
out: list[tuple[str, str]] = []
|
||||
for suffix, group in groups:
|
||||
out.extend(_pack_recursive(group, path + suffix))
|
||||
return out
|
||||
|
||||
|
||||
def chunks_from_page(
|
||||
text: str,
|
||||
page_id: str,
|
||||
@@ -112,7 +247,6 @@ def chunks_from_page(
|
||||
if not paragraphs:
|
||||
return
|
||||
|
||||
# ----- Chunk 0: synthetic anchor for dense retrieval ---------
|
||||
title = metadata.get("title") or page_id
|
||||
first_para = next((p for p in paragraphs if not p.startswith("#")), "")
|
||||
chunk0_body = (
|
||||
@@ -127,28 +261,15 @@ def chunks_from_page(
|
||||
"metadata": {**metadata, "ordinal": 0},
|
||||
}
|
||||
|
||||
# ----- Body chunks: pack paragraphs up to TARGET_CHARS -------
|
||||
ordinal = 1
|
||||
|
||||
def emit(buf: list[str]) -> Iterator[dict]:
|
||||
nonlocal ordinal
|
||||
merged = "\n\n".join(buf)
|
||||
for piece in _hard_split(merged):
|
||||
for body, heading_path in _pack_recursive(paragraphs, []):
|
||||
for piece in _hard_split(body):
|
||||
meta = {**metadata, "ordinal": ordinal}
|
||||
if heading_path:
|
||||
meta["heading_path"] = heading_path
|
||||
yield {
|
||||
"id": f"{metadata['bundle_id']}::{page_id}::{ordinal}",
|
||||
"text": piece,
|
||||
"metadata": {**metadata, "ordinal": ordinal},
|
||||
"metadata": meta,
|
||||
}
|
||||
ordinal += 1
|
||||
|
||||
buf: list[str] = []
|
||||
buf_chars = 0
|
||||
for p in paragraphs:
|
||||
if buf_chars + len(p) > TARGET_CHARS and buf:
|
||||
yield from emit(buf)
|
||||
buf = []
|
||||
buf_chars = 0
|
||||
buf.append(p)
|
||||
buf_chars += len(p)
|
||||
if buf:
|
||||
yield from emit(buf)
|
||||
|
||||
@@ -41,6 +41,21 @@ def _resolve_urls() -> list[str]:
|
||||
OLLAMA_URLS = _resolve_urls()
|
||||
EMBED_MODEL = os.environ.get("EMBED_MODEL", "nomic-embed-text")
|
||||
EMBED_DIM = int(os.environ.get("EMBED_DIM", "768"))
|
||||
# nomic asymmetric prefixes — embed time only. Empty string disables.
|
||||
EMBED_QUERY_PREFIX = os.environ.get("EMBED_QUERY_PREFIX", "search_query: ")
|
||||
EMBED_DOC_PREFIX = os.environ.get("EMBED_DOC_PREFIX", "search_document: ")
|
||||
|
||||
|
||||
def apply_prefix(text: str, prefix: str) -> str:
|
||||
if not prefix:
|
||||
return text
|
||||
return prefix + text
|
||||
|
||||
|
||||
def embed_texts(texts: list[str], *, prefix: str, ef: Any | None = None) -> list[list[float]]:
|
||||
"""Embed texts after applying prefix. Callers store the original strings."""
|
||||
fn = ef if ef is not None else embedding_function()
|
||||
return list(fn([apply_prefix(t, prefix) for t in texts]))
|
||||
|
||||
|
||||
class OllamaEmbeddings(EmbeddingFunction):
|
||||
|
||||
+5
-2
@@ -18,7 +18,7 @@ import chromadb
|
||||
from chromadb.config import Settings
|
||||
|
||||
from .chunk import chunks_from_page
|
||||
from .embeddings import embedding_function
|
||||
from .embeddings import EMBED_DOC_PREFIX, embed_texts, embedding_function
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
|
||||
@@ -78,9 +78,12 @@ def upsert_to_chroma(records: list[dict]) -> int:
|
||||
total = 0
|
||||
for i in range(0, len(records), BATCH):
|
||||
chunk = records[i:i + BATCH]
|
||||
texts = [r["text"] for r in chunk]
|
||||
vectors = embed_texts(texts, prefix=EMBED_DOC_PREFIX, ef=embedding_function())
|
||||
col.upsert(
|
||||
ids=[r["id"] for r in chunk],
|
||||
documents=[r["text"] for r in chunk],
|
||||
documents=texts,
|
||||
embeddings=vectors,
|
||||
metadatas=[r["metadata"] for r in chunk],
|
||||
)
|
||||
total += len(chunk)
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
"""Heading-recursive chunker tests. No corpus, no Chroma."""
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
import rag.chunk as chunk
|
||||
|
||||
|
||||
META = {"bundle_id": "Admin.10.0", "title": "Admin Guide"}
|
||||
|
||||
|
||||
def _bodies(text: str) -> list[str]:
|
||||
return [c["text"] for c in chunk.chunks_from_page(text, "Page", META)]
|
||||
|
||||
|
||||
class ChunkTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self._saved = chunk.TARGET_CHARS
|
||||
|
||||
def tearDown(self) -> None:
|
||||
chunk.TARGET_CHARS = self._saved
|
||||
|
||||
def test_chunk0_always_has_title_on_short_page(self) -> None:
|
||||
page = "# Install\n\nJust one short paragraph about installing."
|
||||
chunks = list(chunk.chunks_from_page(page, "Install", META))
|
||||
self.assertGreaterEqual(len(chunks), 1)
|
||||
self.assertEqual(chunks[0]["metadata"]["ordinal"], 0)
|
||||
self.assertIn("# Admin Guide", chunks[0]["text"])
|
||||
|
||||
def test_sibling_h2_sections_do_not_mix(self) -> None:
|
||||
chunk.TARGET_CHARS = 80
|
||||
page = (
|
||||
"# A\n\n"
|
||||
"intro text for A\n\n"
|
||||
"## A.1\n\n"
|
||||
"alpha content lives here only\n\n"
|
||||
"## A.2\n\n"
|
||||
"beta content lives here only\n"
|
||||
)
|
||||
bodies = _bodies(page)
|
||||
self.assertTrue(any("alpha content" in b and "beta content" not in b for b in bodies[1:]),
|
||||
bodies)
|
||||
self.assertTrue(any("beta content" in b and "alpha content" not in b for b in bodies[1:]),
|
||||
bodies)
|
||||
|
||||
def test_giant_section_splits_on_paragraphs(self) -> None:
|
||||
chunk.TARGET_CHARS = 40
|
||||
paras = [f"Paragraph number {i} with enough words." for i in range(8)]
|
||||
page = "# Giant\n\n" + "\n\n".join(paras)
|
||||
bodies = _bodies(page)
|
||||
# skip chunk 0
|
||||
for b in bodies[1:]:
|
||||
# may exceed by at most one paragraph (the one that filled the buf)
|
||||
self.assertLessEqual(len(b), chunk.TARGET_CHARS + len(paras[0]) + 2, b)
|
||||
|
||||
def test_fenced_code_never_sliced(self) -> None:
|
||||
chunk.TARGET_CHARS = 30
|
||||
fence = "```\n" + ("x" * 80) + "\n```"
|
||||
page = "# Code\n\nBefore.\n\n" + fence + "\n\nAfter."
|
||||
bodies = _bodies(page)
|
||||
joined = "\n".join(bodies)
|
||||
self.assertIn(fence, joined)
|
||||
# no body chunk should contain a half-fence
|
||||
for b in bodies:
|
||||
if "```" in b:
|
||||
self.assertTrue(b.strip().startswith("```") or "```\n" in b)
|
||||
self.assertGreaterEqual(b.count("```"), 2, b)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Prefix helpers — no Ollama, no Chroma."""
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from rag.embeddings import apply_prefix, embed_texts
|
||||
|
||||
|
||||
class FakeEmbedder:
|
||||
def __init__(self) -> None:
|
||||
self.seen: list[str] = []
|
||||
|
||||
def __call__(self, texts: list[str]) -> list[list[float]]:
|
||||
self.seen.extend(texts)
|
||||
return [[0.0, 1.0] for _ in texts]
|
||||
|
||||
|
||||
class PrefixTests(unittest.TestCase):
|
||||
def test_index_sends_doc_prefix_keeps_raw_text(self) -> None:
|
||||
fake = FakeEmbedder()
|
||||
raw = ["hello"]
|
||||
vecs = embed_texts(raw, prefix="search_document: ", ef=fake)
|
||||
self.assertEqual(fake.seen, ["search_document: hello"])
|
||||
self.assertEqual(raw, ["hello"]) # caller storage unchanged
|
||||
self.assertEqual(len(vecs), 1)
|
||||
|
||||
def test_query_sends_query_prefix(self) -> None:
|
||||
fake = FakeEmbedder()
|
||||
embed_texts(["q"], prefix="search_query: ", ef=fake)
|
||||
self.assertEqual(fake.seen, ["search_query: q"])
|
||||
self.assertNotIn("q", fake.seen)
|
||||
|
||||
def test_empty_prefix_is_raw(self) -> None:
|
||||
self.assertEqual(apply_prefix("hello", ""), "hello")
|
||||
fake = FakeEmbedder()
|
||||
embed_texts(["hello"], prefix="", ef=fake)
|
||||
self.assertEqual(fake.seen, ["hello"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user