Files
akyldash/backend/search/evaluate_relevance.py

152 lines
5.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Measure document search Recall@K and MRR@K against a relevance set."""
from __future__ import annotations
import argparse
import json
import sys
import urllib.parse
from pathlib import Path
from search.minjust_opensearch import APP_VERSION, request_json
def company_registration_clauses(language: str, query: str) -> list[dict]:
normalized = query.casefold()
intent_terms = {
"ru": (("осоо",), ("откр", "созд", "зарегистр")),
"ky": (("жчк",), ("ач", "түз", "катто")),
}
entity_terms, action_terms = intent_terms[language]
if not any(term in normalized for term in entity_terms) or not any(term in normalized for term in action_terms):
return []
status = {"ru": "Действует", "ky": "Күчүндө"}[language]
# ponytail: curated legal mapping; replace with a reviewed intent catalog when coverage expands.
def clause(document_code: str, boost: int) -> dict:
return {
"constant_score": {
"filter": {
"bool": {
"filter": [
{"term": {"document_code": document_code}},
{"term": {f"status_{language}": status}},
]
}
},
"boost": boost,
}
}
return [clause("230044970", 2000), clause("667", 1000)]
def search_body(language: str, query: str, top_k: int) -> bytes:
full_text = {
"multi_match": {
"query": query,
"fields": [f"document_name_{language}", f"text_{language}"],
"type": "cross_fields",
}
}
clauses = company_registration_clauses(language, query)
bool_query = {"filter": {"term": {"language": language}}}
if clauses:
bool_query.update({"should": [full_text, *clauses], "minimum_should_match": 1})
else:
bool_query["must"] = full_text
return json.dumps({
"size": top_k,
"track_total_hits": False,
"_source": ["document_code"],
"query": {"bool": bool_query},
"collapse": {"field": "document_code"},
}, ensure_ascii=False).encode()
def load_queries(path: Path) -> list[dict]:
with path.open(encoding="utf-8") as source:
queries = json.load(source)
if not isinstance(queries, list) or not queries:
raise ValueError("Relevance set must be a non-empty JSON array")
seen = set()
for item in queries:
if not isinstance(item, dict) or set(item) != {
"id", "language", "query", "relevant_document_codes"
}:
raise ValueError("Each query must contain id, language, query and relevant_document_codes")
codes = item["relevant_document_codes"]
if (
not isinstance(item["id"], str)
or not item["id"].strip()
or item["id"] in seen
or not isinstance(item["language"], str)
or item["language"] not in {"ru", "ky"}
or not isinstance(item["query"], str)
or not item["query"].strip()
or not isinstance(codes, list)
or not codes
or any(not isinstance(code, str) or not code for code in codes)
or len(codes) != len(set(codes))
):
raise ValueError(f"Invalid relevance query: {item.get('id', '<unknown>')}")
seen.add(item["id"])
return queries
def search(base_url: str, index: str, item: dict, top_k: int) -> list[str]:
language = item["language"]
body = search_body(language, item["query"], top_k)
url = f"{base_url.rstrip('/')}/{urllib.parse.quote(index, safe='')}/_search"
response = request_json(url, "POST", body, "application/json")
try:
return [hit["_source"]["document_code"] for hit in response["hits"]["hits"]]
except (KeyError, TypeError) as error:
raise RuntimeError("OpenSearch search response is incomplete") from error
def evaluate(queries: list[dict], base_url: str, index: str, top_k: int) -> dict:
results = []
for item in queries:
retrieved = search(base_url, index, item, top_k)
relevant = set(item["relevant_document_codes"])
matches = [rank for rank, code in enumerate(retrieved, 1) if code in relevant]
results.append({
"id": item["id"],
"language": item["language"],
"query": item["query"],
"retrieved_document_codes": retrieved,
f"recall_at_{top_k}": len(relevant.intersection(retrieved)) / len(relevant),
f"reciprocal_rank_at_{top_k}": 1 / matches[0] if matches else 0.0,
})
return {
"summary": {
"query_count": len(results),
f"recall_at_{top_k}": sum(item[f"recall_at_{top_k}"] for item in results) / len(results),
f"mrr_at_{top_k}": sum(item[f"reciprocal_rank_at_{top_k}"] for item in results) / len(results),
},
"queries": results,
}
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("relevance_set", type=Path)
parser.add_argument("--url", default="http://127.0.0.1:9200")
parser.add_argument("--index", default="akyldash-fragments-v1")
parser.add_argument("--top-k", type=int, default=10)
parser.add_argument("--version", action="version", version=APP_VERSION)
arguments = parser.parse_args()
if arguments.top_k <= 0:
raise SystemExit("--top-k must be greater than zero")
result = evaluate(load_queries(arguments.relevance_set), arguments.url, arguments.index, arguments.top_k)
print(json.dumps(result, ensure_ascii=False, indent=2))
print(f"Akyldash Backend v{APP_VERSION} · Frontend — not created", file=sys.stderr)
return 0
if __name__ == "__main__":
raise SystemExit(main())