first-aid-rag-assistant / src / evaluation / evaluation.py
evaluation.py
Raw
import json

from tqdm.auto import tqdm

from config import MODEL_NAME
from evaluation.evaluation_utils import llm_structured_retry
from search import rrf, text_search, vector_search  # noqa: F401 (re-exported for notebooks)


def generate_ground_truth(
        doc,
        llm_client,
        data_gen_instructions,
        Questions
    ):
    user_prompt = json.dumps(doc)

    out, usage = llm_structured_retry(
        llm_client,
        data_gen_instructions,
        user_prompt,
        Questions
    )

    results = []

    for q in out.questions:
        results.append({
            "question": q,
            "answer": doc['answer'],
            "document": doc["id"]
        })

    return results, usage

def compute_relevance(q, search_function):
    doc_id = q["document"]
    results = search_function(query=q["question"])

    relevance = []
    for d in results:
        relevance.append(int(d["id"] == doc_id))

    return relevance

def compute_relevance_total(ground_truth, search_function):
    relevance_total = []

    for q in tqdm(ground_truth):
        relevance = compute_relevance(q, search_function)
        relevance_total.append(relevance)

    return relevance_total

def hit_rate(relevance):
    cnt = 0

    for line in relevance:
        if 1 in line:
            cnt = cnt + 1

    return cnt / len(relevance)

def mrr(relevance):
    total_score = 0.0

    for line in relevance:
        for rank in range(len(line)):
            if line[rank] == 1:
                total_score = total_score + 1 / (rank + 1)
                break

    return total_score / len(relevance)

def evaluate(ground_truth, search_function):
    relevance_total = compute_relevance_total(ground_truth, search_function)

    return {
        "hit_rate": hit_rate(relevance_total),
        "mrr": mrr(relevance_total),
    }

def search_boosts(query, question_boost, answer_boost, index, num_results=5):
    boost_dict = {
        "question": question_boost,
        "answer": answer_boost,
    }

    return index.search(
        query,
        num_results=num_results,
        boost_dict=boost_dict,
    )


def generate_rag_answer(rec, doc_idx, assistant):

    question = rec["question"]
    doc_id = rec["document"]
    original_doc = doc_idx[doc_id]

    answer_llm = assistant.rag(question)
    answer_orig = original_doc["answer"]

    result = {
        "question": question,
        "answer_llm": answer_llm,
        "answer_orig": answer_orig,
        "document": doc_id,
    }

    return result

def evaluate_aqa(
        question,
        answer_orig,
        answer_llm,
        llm_client,
        aqa_judge_instructions,
        aqa_judge_prompt,
        AnswerEvaluation,
        model=MODEL_NAME
    ):
        prompt = aqa_judge_prompt.format(
            question=question,
            answer_orig=answer_orig,
            answer_llm=answer_llm
        )

        result, usage = llm_structured_retry(
            llm_client,
            aqa_judge_instructions,
            prompt,
            AnswerEvaluation,
            model=model,
        )

        return result, usage

def judge_record(
        rec,
        llm_client,
        aqa_judge_instructions,
        aqa_judge_prompt,
        AnswerEvaluation,
    ):
    eval_result, usage = evaluate_aqa(
        rec["question"],
        rec["answer_orig"],
        rec["answer_llm"],
        llm_client,
        aqa_judge_instructions,
        aqa_judge_prompt,
        AnswerEvaluation,
    )

    result = {
        "question": rec["question"],
        "document": rec["document"],
        "score": eval_result.score,
        "reasoning": eval_result.reasoning,
    }

    return result, usage