File size: 2,421 Bytes
f5628ad
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39cba11
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
import json

from server.retriever import retrieve_with_scores
from server.utils import setup_logger

logger = setup_logger(__name__)


def compute_precision_at_k(query: str, retrieved_chunks: list[dict], ground_truth: dict, k: int = 5) -> float:
    """
    Precision@K = (relevant chunks in top-K) / K

    A chunk is "relevant" if:
    - Its source matches ground_truth["relevant_sources"], OR
    - Its content contains any keyword from ground_truth["relevant_chunk_keywords"]

    Return float between 0 and 1.
    """
    relevant_sources = ground_truth.get("relevant_sources", [])
    keywords = ground_truth.get("relevant_chunk_keywords", [])

    top_k = retrieved_chunks[:k]
    relevant_count = 0

    for chunk in top_k:
        source_match = chunk.get("source", "") in relevant_sources
        keyword_match = any(
            kw.lower() in chunk.get("content", "").lower() for kw in keywords
        )
        if source_match or keyword_match:
            relevant_count += 1

    precision = relevant_count / k if k > 0 else 0.0
    return round(precision, 4)


def run_batch_precision_eval(eval_pairs_path: str, k: int = 5) -> dict:
    """
    Run precision@K for all queries in eval_pairs.json.
    Return dict with mean_precision_at_k and per_query_results.
    """
    with open(eval_pairs_path, "r") as f:
        eval_pairs = json.load(f)

    per_query_results = []

    for pair in eval_pairs:
        query = pair["query"]
        retrieved = retrieve_with_scores(query, k=k)
        precision = compute_precision_at_k(query, retrieved, pair, k=k)
        retrieved_sources = [c["source"] for c in retrieved]

        per_query_results.append({
            "query": query,
            "precision_at_k": precision,
            "retrieved_sources": retrieved_sources,
        })

        logger.info(f"P@{k}={precision:.2f} | {query[:60]}...")

    mean_precision = sum(r["precision_at_k"] for r in per_query_results) / len(per_query_results)

    return {
        "mean_precision_at_k": round(mean_precision, 4),
        "per_query_results": per_query_results,
    }


def run_batch_precision_eval_multi_k(eval_pairs_path: str, ks: list[int]) -> dict:
    """
    Run precision evaluation for multiple K values.
    Return a dict keyed by precision@K.
    """
    results = {}
    for k in ks:
        results[f"precision@{k}"] = run_batch_precision_eval(eval_pairs_path, k=k)
    return results