Port docs-mcp-template upgrades (eval, citations, prefixes) (#26)
Image rebuild (skip scrape) / build (push) Successful in 3m58s
Image rebuild (skip scrape) / build (push) Successful in 3m58s
Co-authored-by: claude <[email protected]>
This commit was merged in pull request #26.
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"
|
||||
+38
-45
@@ -29,6 +29,7 @@ from mcp.server.mcpserver import MCPServer
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from pydantic import Field
|
||||
|
||||
from .format import format_search_hit
|
||||
from .usage import TimedCall
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -182,6 +183,17 @@ def _collection():
|
||||
return _chroma_collection
|
||||
|
||||
|
||||
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_index():
|
||||
"""Return the BM25 index, or None if it doesn't exist on disk."""
|
||||
global _bm25
|
||||
@@ -257,10 +269,8 @@ def _read_markdown(source: str, source_key: str) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _format_hit(doc: str, meta: dict, distance: float | None = None) -> str:
|
||||
"""Render one retrieval hit as a fenced markdown block with full
|
||||
provenance attached. ``doc`` is the chunk text; ``meta`` is the
|
||||
chunk's metadata dict."""
|
||||
def _format_hit(n: int, doc: str, meta: dict, distance: float | None = None) -> str:
|
||||
"""Render one numbered retrieval hit. ``n`` is 1-based display order."""
|
||||
src_url = meta.get("source_url") or ""
|
||||
src_key = meta.get("source_key") or ""
|
||||
src = meta.get("source") or ""
|
||||
@@ -268,15 +278,10 @@ def _format_hit(doc: str, meta: dict, distance: float | None = None) -> str:
|
||||
brand = meta.get("brand") or ""
|
||||
crop = meta.get("crop") or ""
|
||||
name = meta.get("product_name") or src_key
|
||||
|
||||
header = (
|
||||
f"### {name} \n"
|
||||
f"`{src}::{src_key}` — {vendor} / {brand} / {crop} \n"
|
||||
f"<{src_url}>"
|
||||
)
|
||||
extra = f"`{src}::{src_key}` — {vendor} / {brand} / {crop}"
|
||||
if distance is not None:
|
||||
header += f" \n_(distance={distance:.4f})_"
|
||||
return f"{header}\n\n{doc.strip()}\n"
|
||||
extra += f" distance={distance:.4f}"
|
||||
return format_search_hit(n, name, src_url, doc, extra)
|
||||
|
||||
|
||||
def _rrf_fuse(rankings: list[list[str]], k: int = RRF_K) -> list[str]:
|
||||
@@ -548,7 +553,8 @@ def search_docs(
|
||||
"""Search the seed-variety corpus for hybrids/varieties matching a query.
|
||||
|
||||
Returns the top-k variety chunks with their full source URLs, ratings,
|
||||
maturity, traits, and regional listings. Optional filters narrow to one
|
||||
maturity, traits, and regional listings. Hits are numbered [1]… for
|
||||
citation. Optional filters narrow to one
|
||||
crop, brand, vendor, or scraper source. Use **list_versions()** first
|
||||
to discover valid facet values. Use **lookup_variety()** on any
|
||||
candidate the user is serious about — that returns the canonical
|
||||
@@ -593,11 +599,7 @@ def search_docs(
|
||||
)
|
||||
|
||||
try:
|
||||
dense = col.query(
|
||||
query_texts=[query],
|
||||
n_results=pool_size,
|
||||
where=where,
|
||||
)
|
||||
dense = _query_dense(col, query, pool_size, where)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
_call.set(error_dense=str(exc), hits_returned=0)
|
||||
return f"_(retrieval failed: {exc})_"
|
||||
@@ -702,20 +704,21 @@ def search_docs(
|
||||
)
|
||||
|
||||
blocks: list[str] = []
|
||||
for cid in final_ids:
|
||||
for n, cid in enumerate(final_ids, start=1):
|
||||
doc = id_to_doc.get(cid, "")
|
||||
meta = id_to_meta.get(cid, {})
|
||||
dist = id_to_dist.get(cid) if not used_hybrid else None
|
||||
blocks.append(_format_hit(doc, meta, dist))
|
||||
blocks.append(_format_hit(n, doc, meta, dist))
|
||||
|
||||
header = (
|
||||
f"# Search results — {len(final_ids)} variety chunk"
|
||||
f"{'s' if len(final_ids) != 1 else ''}"
|
||||
f"{' (dense + BM25 hybrid)' if used_hybrid else ' (dense only)'}\n"
|
||||
f"_Use `lookup_variety(source_key=...)` on any candidate "
|
||||
f"to fact-check ratings from the canonical sidecar._\n\n---\n\n"
|
||||
f"_Hits are numbered [1]… for citation. "
|
||||
f"Use `lookup_variety(source_key=...)` on any candidate "
|
||||
f"to fact-check ratings from the canonical sidecar._\n\n"
|
||||
)
|
||||
return header + "\n---\n\n".join(blocks)
|
||||
return header + "\n\n".join(blocks) + "\n"
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
@@ -975,11 +978,7 @@ def search_trials(
|
||||
full_query = f"{query} {product}"
|
||||
|
||||
try:
|
||||
dense = col.query(
|
||||
query_texts=[full_query],
|
||||
n_results=pool_size,
|
||||
where=where,
|
||||
)
|
||||
dense = _query_dense(col, full_query, pool_size, where)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
_call.set(error_dense=str(exc), hits_returned=0)
|
||||
return f"_(trial retrieval failed: {exc})_"
|
||||
@@ -1093,28 +1092,27 @@ def search_trials(
|
||||
)
|
||||
|
||||
blocks: list[str] = []
|
||||
for cid in final_ids:
|
||||
for n, cid in enumerate(final_ids, start=1):
|
||||
doc = id_to_doc.get(cid, "")
|
||||
meta = id_to_meta.get(cid, {})
|
||||
dist = id_to_dist.get(cid) if not used_hybrid else None
|
||||
blocks.append(_format_trial_hit(doc, meta, dist))
|
||||
blocks.append(_format_trial_hit(n, doc, meta, dist))
|
||||
|
||||
header = (
|
||||
f"# Trial search results — {len(final_ids)} trial document"
|
||||
f"{'s' if len(final_ids) != 1 else ''}"
|
||||
f"{' (dense + BM25 hybrid)' if used_hybrid else ' (dense only)'}\n"
|
||||
f"_Use `get_page(source=..., source_key=...)` to read the "
|
||||
f"_Hits are numbered [1]… for citation. "
|
||||
f"Use `get_page(source=..., source_key=...)` to read the "
|
||||
f"full trial body. Use `lookup_variety(source_key=...)` on "
|
||||
f"any product code to verify its identity (RM, traits, "
|
||||
f"disease ratings)._\n\n---\n\n"
|
||||
f"disease ratings)._\n\n"
|
||||
)
|
||||
return header + "\n---\n\n".join(blocks)
|
||||
return header + "\n\n".join(blocks) + "\n"
|
||||
|
||||
|
||||
def _format_trial_hit(doc: str, meta: dict, distance: float | None = None) -> str:
|
||||
"""Trial-specific result header. Highlights crop/state/year and
|
||||
sources URL (vs variety hits which emphasize brand + product
|
||||
identity)."""
|
||||
def _format_trial_hit(n: int, doc: str, meta: dict, distance: float | None = None) -> str:
|
||||
"""Trial-specific numbered hit. Highlights crop/state/year."""
|
||||
src_url = meta.get("source_url") or ""
|
||||
src_key = meta.get("source_key") or ""
|
||||
src = meta.get("source") or ""
|
||||
@@ -1125,15 +1123,10 @@ def _format_trial_hit(doc: str, meta: dict, distance: float | None = None) -> st
|
||||
|
||||
title_bits = [b for b in [crop.title(), region or state, str(year) if year else ""] if b]
|
||||
title = " · ".join(title_bits) if title_bits else src_key
|
||||
|
||||
header = (
|
||||
f"### Trial: {title} \n"
|
||||
f"`{src}::{src_key}` — {meta.get('vendor', '')} / {meta.get('brand', '')} \n"
|
||||
f"<{src_url}>"
|
||||
)
|
||||
extra = f"`{src}::{src_key}` — {meta.get('vendor', '')} / {meta.get('brand', '')}"
|
||||
if distance is not None:
|
||||
header += f" \n_(distance={distance:.4f})_"
|
||||
return f"{header}\n\n{doc.strip()}\n"
|
||||
extra += f" distance={distance:.4f}"
|
||||
return format_search_hit(n, f"Trial: {title}", src_url, doc, extra)
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
|
||||
@@ -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,41 @@
|
||||
# seed-mcp retrieval eval — k=5
|
||||
|
||||
_21 golden queries × 4 retrievers_
|
||||
|
||||
## Summary
|
||||
|
||||
| Retriever | Passed | Recall | P@1 | MRR | Avg ms |
|
||||
|---|---|---|---|---|---|
|
||||
| **hybrid+rerank** | 21/21 | 100.00% | 90.48% | 0.905 | 2064 |
|
||||
| **bm25** | 20/21 | 95.24% | 80.95% | 0.833 | 5 |
|
||||
| **hybrid** | 15/21 | 71.43% | 61.90% | 0.619 | 73 |
|
||||
| **dense** | 14/21 | 66.67% | 38.10% | 0.440 | 79 |
|
||||
|
||||
**Recall** = % of queries where ≥1 top-k chunk satisfied the spec. **P@1** = % where the very first result satisfied it. **MRR** = mean of `1 / rank-of-first-satisfying-result` (0 if missed).
|
||||
|
||||
## Per-query results
|
||||
|
||||
| Query | bm25 | dense | hybrid | hybrid+rerank |
|
||||
|---|---|---|---|---|
|
||||
| `DKC62-08RIB ratings` | ✅ #1 | ❌ | ❌ | ✅ #1 |
|
||||
| `AG29XF4 disease ratings` | ✅ #1 | ❌ | ❌ | ✅ #1 |
|
||||
| `WB6430 westbred wheat` | ✅ #1 | ❌ | ❌ | ✅ #1 |
|
||||
| `E085Z5 corn` | ✅ #1 | ❌ | ❌ | ✅ #1 |
|
||||
| `AP Iliad wheat performance` | ✅ #1 | ❌ | ❌ | ✅ #1 |
|
||||
| `drought tolerant corn for sandy soil short season Iowa` | ✅ #2 | ✅ #1 | ✅ #1 | ✅ #1 |
|
||||
| `soybean cyst nematode SCN resistant variety` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 |
|
||||
| `Phytophthora resistance Rps3a soybean` | ✅ #1 | ✅ #2 | ✅ #1 | ✅ #1 |
|
||||
| `XtendFlex soybean Northern Plains` | ❌ | ✅ #1 | ✅ #1 | ✅ #1 |
|
||||
| `Hard Red Spring wheat stripe rust resistance` | ✅ #1 | ✅ #3 | ✅ #1 | ✅ #1 |
|
||||
| `Soft White Winter wheat Pacific Northwest` | ✅ #1 | ✅ #5 | ✅ #1 | ✅ #1 |
|
||||
| `Goss's Wilt resistance corn` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 |
|
||||
| `best corn 2024 Iowa` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 |
|
||||
| `Indiana corn yield comparison 2024` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 |
|
||||
| `AP Iliad Idaho wheat trial` | ✅ #1 | ✅ #5 | ✅ #1 | ✅ #1 |
|
||||
| `DKC65-95 corn yield in trials` | ✅ #1 | ❌ | ✅ #1 | ✅ #1 |
|
||||
| `NK1701 corn trials head to head` | ✅ #1 | ❌ | ❌ | ✅ #1 |
|
||||
| `silage corn high milk per acre dairy` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 |
|
||||
| `soybean 2025 Minnesota top performers` | ✅ #1 | ✅ #1 | ✅ #1 | ✅ #1 |
|
||||
| `Pioneer P1142 hybrid recommendation` | ✅ | ✅ | ✅ | ✅ |
|
||||
| `DKC65-20 yield Alabama trial` | ✅ | ✅ | ✅ | ✅ |
|
||||
|
||||
+9
-3
@@ -71,8 +71,10 @@ class DenseRetriever:
|
||||
def retrieve(self, query: str, k: int, filters: dict | None) -> list[str]:
|
||||
where = _build_where(filters)
|
||||
try:
|
||||
from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts
|
||||
qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0]
|
||||
r = self.col.query(
|
||||
query_texts=[query], n_results=max(k, self.pool), where=where,
|
||||
query_embeddings=[qvec], n_results=max(k, self.pool), where=where,
|
||||
)
|
||||
except Exception:
|
||||
return []
|
||||
@@ -105,7 +107,9 @@ class HybridRetriever:
|
||||
def retrieve(self, query: str, k: int, filters: dict | None) -> list[str]:
|
||||
where = _build_where(filters)
|
||||
try:
|
||||
d = self.col.query(query_texts=[query], n_results=self.pool, where=where)
|
||||
from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts
|
||||
qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0]
|
||||
d = self.col.query(query_embeddings=[qvec], n_results=self.pool, where=where)
|
||||
dense_ids = (d.get("ids") or [[]])[0]
|
||||
except Exception:
|
||||
dense_ids = []
|
||||
@@ -138,8 +142,10 @@ class HybridRerankRetriever:
|
||||
def retrieve(self, query: str, k: int, filters: dict | None) -> list[str]:
|
||||
where = _build_where(filters)
|
||||
try:
|
||||
from rag.embeddings import EMBED_QUERY_PREFIX, embed_texts
|
||||
qvec = embed_texts([query], prefix=EMBED_QUERY_PREFIX)[0]
|
||||
d = self.col.query(
|
||||
query_texts=[query], n_results=self.pool, where=where,
|
||||
query_embeddings=[qvec], n_results=self.pool, where=where,
|
||||
include=["documents"],
|
||||
)
|
||||
dense_ids = (d.get("ids") or [[]])[0]
|
||||
|
||||
+41
-1
@@ -265,10 +265,15 @@ 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)
|
||||
p.add_argument("--trace", action="store_true")
|
||||
p.add_argument("--rerank-url", default=os.environ.get("RERANK_URL", ""))
|
||||
p.add_argument("--product-name", default=os.environ.get("PRODUCT_NAME", "crop_seed"))
|
||||
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}")
|
||||
@@ -300,14 +305,40 @@ def main() -> int:
|
||||
for r in retrievers:
|
||||
print(f"running {r.name}...")
|
||||
for q in queries:
|
||||
res = _evaluate_one(r, q, args.k, col)
|
||||
res = _evaluate_one(r, q, max_k, col)
|
||||
all_results.append(res)
|
||||
|
||||
summary = _aggregate(all_results)
|
||||
md = _emit_markdown(queries, all_results, summary, args.k)
|
||||
# k-curve from rank_first_match at max_k
|
||||
md += "\n## k-curve (P@1 / recall from rank_first_match)\n\n"
|
||||
md += "| Retriever | " + " | ".join(f"P@1@k={k}" for k in ks) + " |\n"
|
||||
md += "|" + "---|" * (len(ks) + 1) + "\n"
|
||||
by_r: dict[str, list[dict]] = {}
|
||||
for row in all_results:
|
||||
by_r.setdefault(row["retriever"], []).append(row)
|
||||
for name, rows in by_r.items():
|
||||
cells = []
|
||||
for kk in ks:
|
||||
hits = sum(1 for r in rows if r.get("rank_first_match") and r["rank_first_match"] <= kk)
|
||||
cells.append(f"{hits / len(rows):.3f}" if rows else "0")
|
||||
md += f"| `{name}` | " + " | ".join(cells) + " |\n"
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(md, encoding="utf-8")
|
||||
sidecar = args.output.with_suffix(".jsonl")
|
||||
with open(sidecar, "w") as fh:
|
||||
for r in all_results:
|
||||
rank = r.get("rank_first_match")
|
||||
fh.write(json.dumps({
|
||||
"query": r["query"],
|
||||
"retriever": r["retriever"],
|
||||
"passed": bool(r.get("passed")),
|
||||
"p_at_1": 1 if rank == 1 else 0,
|
||||
"rr": (1.0 / rank) if rank else 0.0,
|
||||
"rank_first_match": rank,
|
||||
}) + "\n")
|
||||
print(f"\nreport: {args.output}")
|
||||
print(f"sidecar: {sidecar}")
|
||||
print()
|
||||
# Print summary to stdout too
|
||||
for line in md.split("\n"):
|
||||
@@ -315,6 +346,15 @@ def main() -> int:
|
||||
print(line)
|
||||
if line.startswith("## Per-query"):
|
||||
break
|
||||
if args.compare:
|
||||
from eval.pvalue import compare, load_sidecar, render
|
||||
print()
|
||||
print(render(compare(load_sidecar(sidecar), load_sidecar(args.compare))), end="")
|
||||
if args.trace:
|
||||
from eval.trace import render_misses
|
||||
misses_path = args.output.with_name("misses.md")
|
||||
misses_path.write_text(render_misses(all_results))
|
||||
print(f"misses: {misses_path}")
|
||||
return 0
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Stdlib tests for the permutation test. No Chroma."""
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from eval.pvalue import paired_permutation
|
||||
|
||||
|
||||
class PermutationTests(unittest.TestCase):
|
||||
def test_identical_is_not_significant(self) -> None:
|
||||
a = [1.0, 0.5, 0.0, 1.0]
|
||||
r = paired_permutation(a, list(a), n_resamples=200, seed=0)
|
||||
self.assertEqual(r["Diff(A-B)"], 0.0)
|
||||
self.assertFalse(r["significant"])
|
||||
|
||||
def test_large_shift_is_significant(self) -> None:
|
||||
a = [1.0] * 20
|
||||
b = [0.0] * 20
|
||||
r = paired_permutation(a, b, n_resamples=500, seed=0)
|
||||
self.assertTrue(r["significant"])
|
||||
self.assertLess(r["p_value"], 0.05)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,50 @@
|
||||
"""Miss dump from an eval sidecar or in-memory rows.
|
||||
|
||||
python -m eval.trace --sidecar eval/results/baseline.jsonl
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def render_misses(rows: list[dict]) -> str:
|
||||
misses = [r for r in rows if not r.get("passed") and r.get("p_at_1", 1) == 0]
|
||||
# passed=False is the seed schema; p_at_1==0 covers pvalue sidecars
|
||||
if not misses:
|
||||
misses = [r for r in rows if not r.get("passed")]
|
||||
if not misses:
|
||||
return "# Eval misses\n\n_(none)_\n"
|
||||
lines = [f"# Eval misses ({len(misses)})", ""]
|
||||
for row in misses:
|
||||
lines += [
|
||||
f"## {row.get('query', '')}",
|
||||
"",
|
||||
f"- retriever: `{row.get('retriever')}`",
|
||||
f"- rank_first_match: `{row.get('rank_first_match')}`",
|
||||
f"- kind: `{row.get('kind', '')}`",
|
||||
"",
|
||||
]
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--sidecar", type=Path, default=Path("eval/results/baseline.jsonl"))
|
||||
p.add_argument("--misses-out", type=Path, default=Path("eval/results/misses.md"))
|
||||
args = p.parse_args()
|
||||
if not args.sidecar.exists():
|
||||
args.misses_out.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.misses_out.write_text("# Eval misses\n\nno sidecar\n")
|
||||
print(f"no sidecar ({args.sidecar}); wrote empty misses")
|
||||
return 0
|
||||
rows = [json.loads(line) for line in args.sidecar.read_text().splitlines() if line.strip()]
|
||||
args.misses_out.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.misses_out.write_text(render_misses(rows))
|
||||
print(f"wrote {args.misses_out} ({len(rows)} rows)")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
+23
-2
@@ -14,8 +14,14 @@ import os
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from chromadb import EmbeddingFunction, Documents, Embeddings
|
||||
try:
|
||||
import httpx
|
||||
from chromadb import Documents, EmbeddingFunction, Embeddings
|
||||
except ImportError: # unit tests for apply_prefix / embed_texts
|
||||
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__)
|
||||
|
||||
@@ -23,6 +29,21 @@ OLLAMA_URLS = [u.strip() for u in os.environ.get("OLLAMA_URL",
|
||||
"http://localhost:11434").split(",") if u.strip()]
|
||||
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
@@ -23,7 +23,7 @@ import chromadb
|
||||
from chromadb.config import Settings
|
||||
|
||||
from .chunk import chunks_from_variety, chunks_from_trial
|
||||
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")
|
||||
@@ -83,9 +83,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,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