rag-vietnamese / evaluate_colpali.py
thaidinhz1's picture
Claude Sonnet 4.6
feat: add ColPali eval script + update comparison results in README
0165cfe
Raw History Blame Contribute Delete
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)