""" Đánh giá ColPali retrieval quality so với Text RAG trên cùng golden set. Vì ColPali trả về pages (không có text extracted), ta dùng LLM judge để đánh giá context recall dựa trên metadata (source, page number). Metrics: - Context Recall: tỷ lệ câu hỏi mà page đúng nằm trong top-k results - Source Precision: tỷ lệ results từ đúng source file - Comparison: ColPali vs Text RAG recall head-to-head """ import json import time import os from dotenv import load_dotenv from groq import Groq load_dotenv() client = Groq(api_key=os.environ["GROQ_API_KEY"].strip()) JUDGE_MODEL = "llama-3.3-70b-versatile" def judge_source_relevance(question: str, source_file: str, page: int) -> float: """LLM judge: source file có khả năng chứa câu trả lời cho question không?""" prompt = f"""Câu hỏi: {question} File được retrieve: {source_file} (trang {page}) Dựa vào tên file, đánh giá khả năng file này chứa câu trả lời cho câu hỏi. Chỉ trả về một số từ 0.0 đến 1.0. 1.0 = file chắc chắn liên quan, 0.0 = file không liên quan. Score:""" try: r = client.chat.completions.create( model=JUDGE_MODEL, messages=[{"role": "user", "content": prompt}], max_tokens=10, temperature=0.0, ) text = r.choices[0].message.content.strip() score = float(text.split()[0]) return max(0.0, min(1.0, score)) except Exception: return 0.5 def evaluate_colpali(golden_path: str = "data/golden_set_pdf.json", top_k: int = 3): from src.colpali_retriever import query as colpali_query from src.rag import retrieve as text_retrieve with open(golden_path, encoding="utf-8") as f: golden = json.load(f) items = [q for q in golden if q.get("type") != "out_of_scope"] colpali_scores = [] text_scores = [] print(f"Evaluating {len(items)} questions (top_k={top_k})...\n") print(f"{'Question':<55} {'ColPali':>10} {'Text RAG':>10}") print("-" * 80) for item in items: question = item["question"] expected_source = item.get("source", "") # ColPali retrieval try: colpali_hits = colpali_query(question, top_k=top_k) colpali_recall = max( judge_source_relevance(question, h["source"], h["page"]) for h in colpali_hits ) if colpali_hits else 0.0 time.sleep(1) except Exception as e: print(f" ColPali error: {e}") colpali_recall = 0.0 # Text RAG retrieval try: _, text_sources = text_retrieve(question, top_k=top_k) text_recall = max( judge_source_relevance(question, s["source"], s["page"]) for s in text_sources ) if text_sources else 0.0 time.sleep(1) except Exception as e: print(f" Text RAG error: {e}") text_recall = 0.0 colpali_scores.append(colpali_recall) text_scores.append(text_recall) q_short = question[:53] + ".." if len(question) > 53 else question print(f"{q_short:<55} {colpali_recall:>10.2f} {text_recall:>10.2f}") avg_colpali = sum(colpali_scores) / len(colpali_scores) avg_text = sum(text_scores) / len(text_scores) print("\n" + "=" * 80) print(f"{'RESULTS':<55} {'ColPali':>10} {'Text RAG':>10}") print("=" * 80) print(f"{'Source Relevance (avg)':<55} {avg_colpali:>10.3f} {avg_text:>10.3f}") winner = "ColPali" if avg_colpali > avg_text else "Text RAG" print(f"\nWinner: {winner} (+{abs(avg_colpali - avg_text):.3f})") results = { "summary": { "colpali_source_relevance": avg_colpali, "text_rag_source_relevance": avg_text, "winner": winner, }, "details": [ {"question": items[i]["question"], "colpali_score": colpali_scores[i], "text_score": text_scores[i]} for i in range(len(items)) ] } with open("eval_colpali_results.json", "w", encoding="utf-8") as f: json.dump(results, f, ensure_ascii=False, indent=2) print("\nSaved to eval_colpali_results.json") if __name__ == "__main__": import sys golden = sys.argv[1] if len(sys.argv) > 1 else "data/golden_set_pdf.json" evaluate_colpali(golden_path=golden)