Spaces:
Sleeping
Sleeping
Download evaluate_colpali.py from thaidinhz1/rag-vietnamese: direct link, hf CLI and curl.
- Browser
- Download file 4.45 kB
-
https://huggingface.co/spaces/thaidinhz1/rag-vietnamese/resolve/main/evaluate_colpali.py
- Command line
-
hf download hf://spaces/thaidinhz1/rag-vietnamese/evaluate_colpali.py
-
curl -L -o evaluate_colpali.py https://huggingface.co/spaces/thaidinhz1/rag-vietnamese/resolve/main/evaluate_colpali.py
4.45 kB
| """ | |
| Đá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) | |