medbillcodes-api / app /eval_retrieval.py
Zikorao's picture
Upload folder using huggingface_hub
9c16dbf verified
Raw
History Blame Contribute Delete
14.6 kB
"""Retrieval evaluation + RRF/normalization tuning harness.
Compares BM25-only, k-NN-only, and hybrid fusion against labelled eval sets
(Paediatrics + Family Practice by default). Reports pooled and per-specialty
top-1 / top-3 / MRR so promote decisions do not overfit one specialty.
Usage:
python -m app.eval_retrieval
python -m app.eval_retrieval --sets pediatric,family --smoke 12
"""
from __future__ import annotations
import argparse
import json
import logging
from collections import defaultdict
from pathlib import Path
from .config import settings
from .embeddings import embed_text
from .opensearch_client import get_client
logger = logging.getLogger(__name__)
DATA_DIR = Path(__file__).resolve().parent.parent / "data"
# Named sets used by --sets and embed_experiment.
EVAL_SETS: dict[str, tuple[Path, str]] = {
"pediatric": (DATA_DIR / "pediatric_eval.json", "26"),
"family": (DATA_DIR / "family_eval.json", "00"),
}
TOP_K = 10
def index_vector_dim(client=None) -> int | None:
"""Return current ``code_vector`` dimension, or None if index missing."""
client = client or get_client()
idx = settings.opensearch_index
if not client.indices.exists(index=idx):
return None
mapping = client.indices.get_mapping(index=idx)
props = mapping[idx]["mappings"]["properties"]
vec = props.get("code_vector") or {}
return vec.get("dimension")
def _empty_metrics(cases: list[dict], *, skipped: str) -> dict:
by_spec: dict[str, list[int | None]] = defaultdict(list)
for case in cases:
by_spec[case.get("specialty") or "unknown"].append(None)
out = _metrics_from_ranks([None] * len(cases))
out["by_specialty"] = {
spec: _metrics_from_ranks(rs) for spec, rs in sorted(by_spec.items())
}
out["skipped"] = skipped
return out
def vector_dims_compatible(client=None) -> tuple[bool, str | None]:
"""True when query embed dim matches the OpenSearch ``code_vector`` dim."""
index_dim = index_vector_dim(client)
if index_dim is None:
return False, "index missing or has no code_vector dimension"
# Prefer a live probe so remote (Featherless) dim mismatches surface even
# when EMBEDDING_DIM in env is wrong/stale.
try:
probe_dim = len(embed_text("dimension probe"))
except Exception as exc: # noqa: BLE001
return False, f"embed probe failed: {exc}"
if probe_dim != index_dim:
return (
False,
f"query embed dim {probe_dim} != index code_vector dim {index_dim} "
f"(settings.embedding_dim={settings.embedding_dim})",
)
if settings.embedding_dim != index_dim:
logger.warning(
"settings.embedding_dim=%s differs from index/probe dim=%s",
settings.embedding_dim,
index_dim,
)
return True, None
def _specialty_filters(specialty: str | None) -> list[dict]:
"""Production-like section filters (desk retrieval is specialty-scoped)."""
if specialty == "00":
from .family_scope import family_practice_00_filter # noqa: PLC0415
return family_practice_00_filter()
if specialty == "26":
from .specialty_scope import paediatrics_26_filter # noqa: PLC0415
return paediatrics_26_filter()
return []
def _bm25_query(text: str, filters: list[dict] | None = None) -> dict:
bm25: dict = {
"bool": {
"must": [
{
"multi_match": {
"query": text,
"fields": [
"billing_code^2",
"description_text",
"rules_and_constraints",
],
}
}
]
}
}
if filters:
bm25["bool"]["filter"] = filters
return {
"size": TOP_K,
"_source": ["billing_code"],
"query": bm25,
}
def _knn_query(text: str, filters: list[dict] | None = None) -> dict:
knn_field: dict = {"vector": embed_text(text), "k": TOP_K}
if filters:
knn_field["filter"] = {"bool": {"filter": filters}}
return {
"size": TOP_K,
"_source": ["billing_code"],
"query": {"knn": {"code_vector": knn_field}},
}
def _hybrid_query(text: str, filters: list[dict] | None = None) -> dict:
bm25: dict = {
"bool": {
"must": [
{
"multi_match": {
"query": text,
"fields": [
"billing_code^2",
"description_text",
"rules_and_constraints",
],
}
}
]
}
}
if filters:
bm25["bool"]["filter"] = filters
knn_field: dict = {"vector": embed_text(text), "k": TOP_K}
if filters:
knn_field["filter"] = {"bool": {"filter": filters}}
return {
"size": TOP_K,
"_source": ["billing_code"],
"query": {
"hybrid": {
"queries": [bm25, {"knn": {"code_vector": knn_field}}],
}
},
}
def _codes(resp: dict) -> list[str]:
return [h["_source"]["billing_code"] for h in resp["hits"]["hits"]]
def _ensure_rrf_pipeline(client, name: str, rank_constant: int) -> None:
client.transport.perform_request(
"PUT",
f"/_search/pipeline/{name}",
body={
"phase_results_processors": [
{
"score-ranker-processor": {
"combination": {
"technique": "rrf",
"rank_constant": rank_constant,
}
}
}
]
},
)
def _ensure_norm_pipeline(client, name: str, weights: list[float]) -> None:
client.transport.perform_request(
"PUT",
f"/_search/pipeline/{name}",
body={
"phase_results_processors": [
{
"normalization-processor": {
"normalization": {"technique": "min_max"},
"combination": {
"technique": "arithmetic_mean",
"parameters": {"weights": weights},
},
}
}
]
},
)
def load_cases(
*,
sets: list[str] | None = None,
smoke: int | None = None,
) -> list[dict]:
"""Load labelled cases; each row gains ``specialty`` + ``eval_set``."""
names = sets or ["pediatric", "family"]
cases: list[dict] = []
for name in names:
key = name.strip().lower()
if key not in EVAL_SETS:
raise ValueError(
f"Unknown eval set '{name}'. Known: {', '.join(EVAL_SETS)}"
)
path, specialty = EVAL_SETS[key]
if not path.exists():
raise FileNotFoundError(path)
rows = json.loads(path.read_text())
for row in rows:
cases.append(
{
**row,
"specialty": specialty,
"eval_set": key,
}
)
if smoke is not None and smoke > 0:
# Round-robin across sets so smoke is not all paediatric.
by_set: dict[str, list[dict]] = defaultdict(list)
for c in cases:
by_set[c["eval_set"]].append(c)
picked: list[dict] = []
i = 0
while len(picked) < smoke:
progressed = False
for key in by_set:
bucket = by_set[key]
if i < len(bucket):
picked.append(bucket[i])
progressed = True
if len(picked) >= smoke:
break
if not progressed:
break
i += 1
cases = picked
return cases
def _metrics_from_ranks(ranks: list[int | None]) -> dict:
n = len(ranks)
if n == 0:
return {"top1": 0.0, "top3": 0.0, "mrr": 0.0, "n": 0}
top1 = sum(1 for r in ranks if r == 1)
top3 = sum(1 for r in ranks if r is not None and r <= 3)
mrr = sum(1.0 / r for r in ranks if r is not None)
return {
"top1": round(top1 / n, 3),
"top3": round(top3 / n, 3),
"mrr": round(mrr / n, 3),
"n": n,
}
def _score(cases: list[dict], runner) -> dict:
"""Pooled metrics plus per-specialty breakdown.
``runner`` receives the full case dict (query + specialty).
"""
ranks: list[int | None] = []
by_spec: dict[str, list[int | None]] = defaultdict(list)
for case in cases:
codes = runner(case)
expected = set(case["expected"])
rank = next((i + 1 for i, c in enumerate(codes) if c in expected), None)
ranks.append(rank)
by_spec[case.get("specialty") or "unknown"].append(rank)
out = _metrics_from_ranks(ranks)
out["by_specialty"] = {
spec: _metrics_from_ranks(rs) for spec, rs in sorted(by_spec.items())
}
return out
def run(
*,
sets: list[str] | None = None,
smoke: int | None = None,
bm25_only: bool = False,
# Default off: paediatric eval labels still include many FP office codes
# (A001/A007/…) that the Paeds (26) section filter correctly excludes.
# Use specialty_scope=True for Family Practice desk-like scoring.
specialty_scope: bool = False,
) -> dict:
cases = load_cases(sets=sets, smoke=smoke)
client = get_client()
idx = settings.opensearch_index
results: dict[str, dict] = {}
notes: list[str] = []
if specialty_scope:
notes.append("specialty_scope=on (Family/Paeds section filters, desk-like)")
else:
notes.append("specialty_scope=off (whole-index BM25)")
def _filters(case: dict) -> list[dict] | None:
if not specialty_scope:
return None
filt = _specialty_filters(case.get("specialty"))
return filt or None
results["bm25_only"] = _score(
cases,
lambda case: _codes(
client.search(
index=idx,
body=_bm25_query(case["query"], _filters(case)),
)
),
)
vectors_ok, dim_note = (
(False, "bm25_only requested") if bm25_only else vector_dims_compatible(client)
)
if not vectors_ok:
skip = dim_note or "vector search unavailable"
notes.append(f"SKIP knn/hybrid: {skip}")
logger.warning("%s", notes[-1])
results["knn_only"] = _empty_metrics(cases, skipped=skip)
for rc in (20, 60, 100):
results[f"hybrid_rrf_rc{rc}"] = _empty_metrics(cases, skipped=skip)
for w in (0.7, 0.8, 0.9, 0.95):
results[f"hybrid_norm_bm25{w}"] = _empty_metrics(cases, skipped=skip)
else:
results["knn_only"] = _score(
cases,
lambda case: _codes(
client.search(
index=idx,
body=_knn_query(case["query"], _filters(case)),
)
),
)
for rc in (20, 60, 100):
name = f"rrf-{rc}"
_ensure_rrf_pipeline(client, name, rc)
results[f"hybrid_rrf_rc{rc}"] = _score(
cases,
lambda case, n=name: _codes(
client.search(
index=idx,
body=_hybrid_query(case["query"], _filters(case)),
params={"search_pipeline": n},
)
),
)
for w in ([0.7, 0.3], [0.8, 0.2], [0.9, 0.1], [0.95, 0.05]):
name = f"norm-{int(w[0] * 100)}{int(w[1] * 100):02d}"
_ensure_norm_pipeline(client, name, w)
results[f"hybrid_norm_bm25{w[0]}"] = _score(
cases,
lambda case, n=name: _codes(
client.search(
index=idx,
body=_hybrid_query(case["query"], _filters(case)),
params={"search_pipeline": n},
)
),
)
set_counts: dict[str, int] = defaultdict(int)
for c in cases:
set_counts[c["eval_set"]] += 1
return {
"n_cases": len(cases),
"sets": dict(set_counts),
"smoke": smoke,
"specialty_scope": specialty_scope,
"index_dim": index_vector_dim(client),
"embedding_dim": settings.embedding_dim,
"vectors_ok": vectors_ok,
"notes": notes,
"results": results,
}
def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
p = argparse.ArgumentParser(description="OHIP retrieval fusion eval")
p.add_argument(
"--sets",
type=str,
default="pediatric,family",
help="Comma-separated eval set names (default: pediatric,family)",
)
p.add_argument(
"--smoke",
type=int,
default=None,
help="Evaluate only N cases (round-robin across sets)",
)
p.add_argument(
"--bm25-only",
action="store_true",
help="Skip k-NN / hybrid (useful when index dim mismatches embed dim)",
)
p.add_argument(
"--specialty-scope",
action="store_true",
help="Apply Family/Paeds section filters (desk-like; use for family set)",
)
return p.parse_args(argv)
if __name__ == "__main__":
logging.basicConfig(level=logging.WARNING)
args = _parse_args()
names = [s.strip() for s in args.sets.split(",") if s.strip()]
out = run(
sets=names,
smoke=args.smoke,
bm25_only=args.bm25_only,
specialty_scope=args.specialty_scope,
)
print(
f"\nRetrieval eval over {out['n_cases']} cases "
f"(sets={out['sets']}"
+ (f", smoke={out['smoke']}" if out.get("smoke") else "")
+ "):\n"
)
for note in out.get("notes") or []:
print(f"NOTE: {note}")
print(f"{'config':<28} {'top1':>6} {'top3':>6} {'mrr':>6}")
print("-" * 50)
for cfg, m in out["results"].items():
print(f"{cfg:<28} {m['top1']:>6} {m['top3']:>6} {m['mrr']:>6}")
# Per-specialty snapshot for BM25 vs best hybrid top3
bm25 = out["results"].get("bm25_only", {})
print("\nPer-specialty top3 (bm25_only):")
for spec, m in (bm25.get("by_specialty") or {}).items():
print(f" specialty {spec}: top3={m['top3']} n={m['n']}")