feat: port template upgrades (citations, eval, chunking, nomic prefixes)
mcp 2.x already on origin/main (#10). Does not change BM25-first search_docs — n=6 is too small to flip the default. - Numbered [1] citations via docs_mcp/format.py - Eval P@1 + JSONL sidecar + eval.pvalue + eval.trace - Heading-recursive chunker, keep chunk-0 and MAX_CHARS=4000 - Nomic prefixes at embed time only; stored text unprefixed Closes #12
This commit is contained in:
@@ -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"
|
||||||
+29
-17
@@ -31,6 +31,7 @@ from mcp.server.mcpserver import MCPServer
|
|||||||
from mcp.server.transport_security import TransportSecuritySettings
|
from mcp.server.transport_security import TransportSecuritySettings
|
||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
|
|
||||||
|
from .format import format_search_hits
|
||||||
from .usage import TimedCall
|
from .usage import TimedCall
|
||||||
|
|
||||||
log = logging.getLogger(__name__)
|
log = logging.getLogger(__name__)
|
||||||
@@ -145,6 +146,17 @@ def _collection():
|
|||||||
return _CHROMA
|
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():
|
def _bm25():
|
||||||
"""Lazy BM25Index handle. None if the FTS5 db isn't built."""
|
"""Lazy BM25Index handle. None if the FTS5 db isn't built."""
|
||||||
global _BM25
|
global _BM25
|
||||||
@@ -258,9 +270,10 @@ def search_docs(
|
|||||||
"""Search the HPE Morpheus Enterprise (Morpheus) docs corpus.
|
"""Search the HPE Morpheus Enterprise (Morpheus) docs corpus.
|
||||||
|
|
||||||
Returns the top-k most relevant chunks (with full source page URLs)
|
Returns the top-k most relevant chunks (with full source page URLs)
|
||||||
given a natural-language query. Optional filters narrow the search
|
given a natural-language query. Hits are numbered [1]… for citation.
|
||||||
to one version, one platform, or one bundle. Use list_versions()
|
Optional filters narrow the search to one version, one platform, or
|
||||||
first if you need to discover the available facet values.
|
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
|
Call this tool whenever the user asks anything that should be
|
||||||
answerable from the official product documentation — install,
|
answerable from the official product documentation — install,
|
||||||
@@ -296,7 +309,7 @@ def search_docs(
|
|||||||
|
|
||||||
if HYBRID_SEARCH and bm is not None:
|
if HYBRID_SEARCH and bm is not None:
|
||||||
try:
|
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]
|
dense_ids = (dense_res.get("ids") or [[]])[0]
|
||||||
bm_hits = bm.query(query, n=pool, where=bm25_where)
|
bm_hits = bm.query(query, n=pool, where=bm25_where)
|
||||||
bm_ids = [cid for cid, _s in bm_hits]
|
bm_ids = [cid for cid, _s in bm_hits]
|
||||||
@@ -326,7 +339,7 @@ def search_docs(
|
|||||||
log.warning("BM25 retrieval failed, falling back to dense: %s", e)
|
log.warning("BM25 retrieval failed, falling back to dense: %s", e)
|
||||||
|
|
||||||
if not docs:
|
if not docs:
|
||||||
res = col.query(query_texts=[query], n_results=k, where=where)
|
res = _query_dense(col, query, k, where)
|
||||||
docs = (res.get("documents") or [[]])[0]
|
docs = (res.get("documents") or [[]])[0]
|
||||||
metas = (res.get("metadatas") or [[]])[0]
|
metas = (res.get("metadatas") or [[]])[0]
|
||||||
dists = (res.get("distances") or [[]])[0]
|
dists = (res.get("distances") or [[]])[0]
|
||||||
@@ -342,7 +355,7 @@ def search_docs(
|
|||||||
extra = bm.query(query, n=pool_size, where=bm25_where) if bm else []
|
extra = bm.query(query, n=pool_size, where=bm25_where) if bm else []
|
||||||
extra_ids = [cid for cid, _s in extra]
|
extra_ids = [cid for cid, _s in extra]
|
||||||
else:
|
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]
|
extra_ids = (extra_res.get("ids") or [[]])[0]
|
||||||
if extra_ids:
|
if extra_ids:
|
||||||
d2, m2, _ = _enrich_from_chroma(col, extra_ids, None)
|
d2, m2, _ = _enrich_from_chroma(col, extra_ids, None)
|
||||||
@@ -371,22 +384,21 @@ def search_docs(
|
|||||||
if not docs:
|
if not docs:
|
||||||
return f"_No matches for `{query}`._"
|
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):
|
for doc, meta, dist in zip(docs, metas, dists):
|
||||||
bid = meta.get("bundle_id", "")
|
bid = meta.get("bundle_id", "")
|
||||||
pid = meta.get("page_id", "")
|
pid = meta.get("page_id", "")
|
||||||
title = meta.get("title") or pid
|
title = meta.get("title") or pid
|
||||||
ver = meta.get("version") or ""
|
ver = meta.get("version") or ""
|
||||||
url = _source_url(bid, pid)
|
url = _source_url(bid, pid)
|
||||||
header = f"## {title}"
|
extra = f"score={1 - dist:.3f}"
|
||||||
if ver:
|
if ver:
|
||||||
header += f" _(v{ver})_"
|
extra = f"v{ver} " + extra
|
||||||
out.append(header)
|
if pid:
|
||||||
out.append(f"[{bid}/{pid}]({url}) · score={1 - dist:.3f}")
|
extra += f" page_id: `{pid}`"
|
||||||
out.append("")
|
hits.append((title, url, doc, extra))
|
||||||
out.append(doc.strip())
|
header = f"# {len(docs)} result(s) for `{query}`\n\n"
|
||||||
out.append("")
|
return header + format_search_hits(hits)
|
||||||
return "\n".join(out)
|
|
||||||
|
|
||||||
|
|
||||||
@mcp.tool()
|
@mcp.tool()
|
||||||
@@ -1073,8 +1085,8 @@ def find_doc_inconsistencies(
|
|||||||
return f"Couldn't open Chroma collection: {e}"
|
return f"Couldn't open Chroma collection: {e}"
|
||||||
where = _build_where(version, platform, bundle_id)
|
where = _build_where(version, platform, bundle_id)
|
||||||
try:
|
try:
|
||||||
res = col.query(query_texts=[scope_query], n_results=max_pages * 3,
|
res = _query_dense(col, scope_query, max_pages * 3, where,
|
||||||
where=where, include=["metadatas"])
|
include=["metadatas"])
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
_call.set(error=f"query: {e}")
|
_call.set(error=f"query: {e}")
|
||||||
return f"Scope query failed: {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())
|
||||||
+3
-1
@@ -51,7 +51,9 @@ class DenseRetriever:
|
|||||||
self.pool = pool
|
self.pool = pool
|
||||||
|
|
||||||
def retrieve(self, query: str, k: int = 10) -> list[tuple[str, str]]:
|
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]
|
ids = (res.get("ids") or [[]])[0]
|
||||||
return _collapse_to_pages(ids, k)
|
return _collapse_to_pages(ids, k)
|
||||||
|
|
||||||
|
|||||||
+38
-8
@@ -34,6 +34,12 @@ def load_queries(path: Path) -> list[dict]:
|
|||||||
return [json.loads(line) for line in fh if line.strip()]
|
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:
|
def reciprocal_rank(retrieved: list[tuple[str, str]], expected: list[tuple[str, str]]) -> float:
|
||||||
expected_set = set(expected)
|
expected_set = set(expected)
|
||||||
for i, page in enumerate(retrieved, start=1):
|
for i, page in enumerate(retrieved, start=1):
|
||||||
@@ -65,8 +71,12 @@ def main() -> int:
|
|||||||
p = argparse.ArgumentParser()
|
p = argparse.ArgumentParser()
|
||||||
p.add_argument("--queries", type=Path, default=Path("eval/queries.jsonl"))
|
p.add_argument("--queries", type=Path, default=Path("eval/queries.jsonl"))
|
||||||
p.add_argument("--k", type=int, default=5)
|
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("--output", type=Path, default=Path("eval/results/baseline.md"))
|
||||||
|
p.add_argument("--compare", type=Path, default=None)
|
||||||
args = p.parse_args()
|
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():
|
if not args.queries.exists():
|
||||||
print(f"queries file not found: {args.queries}")
|
print(f"queries file not found: {args.queries}")
|
||||||
@@ -83,7 +93,7 @@ def main() -> int:
|
|||||||
from rag.bm25 import BM25Index
|
from rag.bm25 import BM25Index
|
||||||
from eval.retrievers import DenseRetriever, BM25Retriever, HybridRetriever
|
from eval.retrievers import DenseRetriever, BM25Retriever, HybridRetriever
|
||||||
|
|
||||||
product = os.environ.get("PRODUCT_NAME", "hvm")
|
product = os.environ.get("PRODUCT_NAME", "morpheus")
|
||||||
repo_root = Path(__file__).resolve().parent.parent
|
repo_root = Path(__file__).resolve().parent.parent
|
||||||
client = chromadb.PersistentClient(path=str(repo_root / "chroma"),
|
client = chromadb.PersistentClient(path=str(repo_root / "chroma"),
|
||||||
settings=Settings(anonymized_telemetry=False))
|
settings=Settings(anonymized_telemetry=False))
|
||||||
@@ -109,34 +119,39 @@ def main() -> int:
|
|||||||
rows: dict[str, dict[str, float]] = {}
|
rows: dict[str, dict[str, float]] = {}
|
||||||
per_query: list[dict] = []
|
per_query: list[dict] = []
|
||||||
for r in retrievers:
|
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
|
elapsed_sum = 0.0
|
||||||
for q in queries:
|
for q in queries:
|
||||||
expected = [(e["bundle_id"], e["page_id"]) for e in q["expected"]]
|
expected = [(e["bundle_id"], e["page_id"]) for e in q["expected"]]
|
||||||
t0 = time.time()
|
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
|
elapsed = time.time() - t0
|
||||||
mrr = reciprocal_rank(retrieved, expected)
|
mrr = reciprocal_rank(retrieved, expected)
|
||||||
|
p1 = p_at_1(retrieved, expected)
|
||||||
recall = recall_at_k(retrieved, expected, args.k)
|
recall = recall_at_k(retrieved, expected, args.k)
|
||||||
ndcg = ndcg_at_k(retrieved, expected, args.k)
|
ndcg = ndcg_at_k(retrieved, expected, args.k)
|
||||||
mrr_sum += mrr
|
mrr_sum += mrr
|
||||||
|
p1_sum += p1
|
||||||
recall_sum += recall
|
recall_sum += recall
|
||||||
ndcg_sum += ndcg
|
ndcg_sum += ndcg
|
||||||
elapsed_sum += elapsed
|
elapsed_sum += elapsed
|
||||||
per_query.append({
|
per_query.append({
|
||||||
"retriever": r.name, "query": q["query"],
|
"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,
|
"top1": list(retrieved[0]) if retrieved else None,
|
||||||
|
"ranked": [list(p) for p in retrieved],
|
||||||
"elapsed_s": round(elapsed, 3),
|
"elapsed_s": round(elapsed, 3),
|
||||||
})
|
})
|
||||||
n = len(queries)
|
n = len(queries)
|
||||||
rows[r.name] = {
|
rows[r.name] = {
|
||||||
|
"P@1": p1_sum / n,
|
||||||
"MRR": mrr_sum / n,
|
"MRR": mrr_sum / n,
|
||||||
f"Recall@{args.k}": recall_sum / n,
|
f"Recall@{args.k}": recall_sum / n,
|
||||||
f"nDCG@{args.k}": ndcg_sum / n,
|
f"nDCG@{args.k}": ndcg_sum / n,
|
||||||
"avg_latency_s": elapsed_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"Recall@{args.k}={rows[r.name][f'Recall@{args.k}']:.3f} "
|
||||||
f"nDCG@{args.k}={rows[r.name][f'nDCG@{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")
|
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)
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||||
md = [f"# Retrieval eval — k={args.k}", "",
|
md = [f"# Retrieval eval — k={args.k}", "",
|
||||||
f"_{len(queries)} hand-curated queries, generated {time.strftime('%Y-%m-%d %H:%M:%S')}_", "",
|
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():
|
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 |")
|
f"| {m[f'nDCG@{args.k}']:.3f} | {m['avg_latency_s']*1000:.0f}ms |")
|
||||||
md += ["", "## Per-query results", "",
|
md += ["", "## Per-query results", "",
|
||||||
"| Retriever | Query | MRR | top-1 |", "| --- | --- | ---: | --- |"]
|
"| 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 "—"
|
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} |")
|
md.append(f"| `{r['retriever']}` | {r['query'][:60]} | {r['mrr']:.3f} | {top1} |")
|
||||||
args.output.write_text("\n".join(md) + "\n")
|
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 {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
|
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", "morpheus")
|
||||||
|
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 = bm25_pages or dense_pages # default retrieval is BM25-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
|
Chunk by semantic section (ATX headings), not raw page/length. A
|
||||||
significantly from prose. The output shape (id, text, metadata) is
|
synthetic chunk 0 (title + first paragraph + optional keyword bag) is
|
||||||
fixed by the downstream Chroma + BM25 indexing in rag/index.py — don't
|
always emitted first — dense retrieval lands on it. Do not drop it.
|
||||||
change that.
|
|
||||||
|
|
||||||
The key knob you'll tune per product is chunk-0. Dense retrieval lands
|
No chonkie / langchain dependency. The output shape (id, text, metadata)
|
||||||
on chunk 0 first for most queries. Make it a synthetic chunk built
|
is fixed by rag/index.py — don't change that. `heading_path` is optional
|
||||||
from:
|
metadata (e.g. "Install > Linux") and is not required for indexing.
|
||||||
|
|
||||||
- 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.
|
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -26,17 +14,12 @@ import re
|
|||||||
from typing import Iterator
|
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
|
CHARS_PER_TOKEN = 4
|
||||||
TARGET_TOKENS = 500
|
TARGET_TOKENS = 500
|
||||||
TARGET_CHARS = TARGET_TOKENS * CHARS_PER_TOKEN
|
TARGET_CHARS = TARGET_TOKENS * CHARS_PER_TOKEN
|
||||||
# Hard cap: nomic-embed-text's context is 2048 tokens. Anything larger
|
# nomic-embed-text context is 2048 tokens. Markdown tables with lots of
|
||||||
# 400s the entire embed batch. 6000 chars works for prose but markdown
|
# `|` tokenize ~1.4× denser than prose; 4000 chars stays under 2048 even
|
||||||
# tables with lots of `|` separators tokenize ~1.4× denser; a 5839-char
|
# for dense table chunks.
|
||||||
# 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.
|
|
||||||
MAX_CHARS = 4000
|
MAX_CHARS = 4000
|
||||||
|
|
||||||
|
|
||||||
@@ -97,6 +80,158 @@ def split_paragraphs(md: str) -> list[str]:
|
|||||||
return [b for b in blocks if b]
|
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(
|
def chunks_from_page(
|
||||||
text: str,
|
text: str,
|
||||||
page_id: str,
|
page_id: str,
|
||||||
@@ -112,7 +247,6 @@ def chunks_from_page(
|
|||||||
if not paragraphs:
|
if not paragraphs:
|
||||||
return
|
return
|
||||||
|
|
||||||
# ----- Chunk 0: synthetic anchor for dense retrieval ---------
|
|
||||||
title = metadata.get("title") or page_id
|
title = metadata.get("title") or page_id
|
||||||
first_para = next((p for p in paragraphs if not p.startswith("#")), "")
|
first_para = next((p for p in paragraphs if not p.startswith("#")), "")
|
||||||
chunk0_body = (
|
chunk0_body = (
|
||||||
@@ -127,28 +261,15 @@ def chunks_from_page(
|
|||||||
"metadata": {**metadata, "ordinal": 0},
|
"metadata": {**metadata, "ordinal": 0},
|
||||||
}
|
}
|
||||||
|
|
||||||
# ----- Body chunks: pack paragraphs up to TARGET_CHARS -------
|
|
||||||
ordinal = 1
|
ordinal = 1
|
||||||
|
for body, heading_path in _pack_recursive(paragraphs, []):
|
||||||
def emit(buf: list[str]) -> Iterator[dict]:
|
for piece in _hard_split(body):
|
||||||
nonlocal ordinal
|
meta = {**metadata, "ordinal": ordinal}
|
||||||
merged = "\n\n".join(buf)
|
if heading_path:
|
||||||
for piece in _hard_split(merged):
|
meta["heading_path"] = heading_path
|
||||||
yield {
|
yield {
|
||||||
"id": f"{metadata['bundle_id']}::{page_id}::{ordinal}",
|
"id": f"{metadata['bundle_id']}::{page_id}::{ordinal}",
|
||||||
"text": piece,
|
"text": piece,
|
||||||
"metadata": {**metadata, "ordinal": ordinal},
|
"metadata": meta,
|
||||||
}
|
}
|
||||||
ordinal += 1
|
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)
|
|
||||||
|
|||||||
+23
-2
@@ -22,8 +22,14 @@ import logging
|
|||||||
import time
|
import time
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
try:
|
||||||
from chromadb import EmbeddingFunction, Documents, Embeddings
|
import httpx
|
||||||
|
from chromadb import Documents, EmbeddingFunction, Embeddings
|
||||||
|
except ImportError:
|
||||||
|
httpx = None # type: ignore[assignment]
|
||||||
|
EmbeddingFunction = object # type: ignore[misc,assignment]
|
||||||
|
Documents = list # type: ignore[misc,assignment]
|
||||||
|
Embeddings = list # type: ignore[misc,assignment]
|
||||||
|
|
||||||
log = logging.getLogger(__name__)
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -41,6 +47,21 @@ def _resolve_urls() -> list[str]:
|
|||||||
OLLAMA_URLS = _resolve_urls()
|
OLLAMA_URLS = _resolve_urls()
|
||||||
EMBED_MODEL = os.environ.get("EMBED_MODEL", "nomic-embed-text")
|
EMBED_MODEL = os.environ.get("EMBED_MODEL", "nomic-embed-text")
|
||||||
EMBED_DIM = int(os.environ.get("EMBED_DIM", "768"))
|
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):
|
class OllamaEmbeddings(EmbeddingFunction):
|
||||||
|
|||||||
+5
-2
@@ -18,7 +18,7 @@ import chromadb
|
|||||||
from chromadb.config import Settings
|
from chromadb.config import Settings
|
||||||
|
|
||||||
from .chunk import chunks_from_page
|
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__)
|
log = logging.getLogger(__name__)
|
||||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
|
||||||
@@ -78,9 +78,12 @@ def upsert_to_chroma(records: list[dict]) -> int:
|
|||||||
total = 0
|
total = 0
|
||||||
for i in range(0, len(records), BATCH):
|
for i in range(0, len(records), BATCH):
|
||||||
chunk = records[i:i + 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(
|
col.upsert(
|
||||||
ids=[r["id"] for r in chunk],
|
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],
|
metadatas=[r["metadata"] for r in chunk],
|
||||||
)
|
)
|
||||||
total += len(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