chatbotai_rag / scripts /benchmark_rag.py
quoctrong's picture
perf: Optimize Latency (Remove query translation, async cache, tune LLM threads, limit tokens)
9d3288b
Raw
History Blame Contribute Delete
4.21 kB
import os
import time
import json
import logging
from pathlib import Path
# Cấu hình log để dễ nhìn output
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
# Ensure working directory is project root
if not Path("src").exists():
os.chdir("..")
from src.core.rag_pipeline import RAGPipeline
def run_benchmark():
# Tắt logging của một số thư viện con để console sạch hơn
logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("sentence_transformers").setLevel(logging.WARNING)
print("\n" + "="*50)
print("🚀 BẮT ĐẦU CHẠY BENCHMARK RAG PIPELINE")
print("="*50 + "\n")
# Khởi tạo pipeline
start_init = time.time()
pipeline = RAGPipeline()
print(f"[+] Khởi tạo RAGPipeline mất: {time.time() - start_init:.2f}s")
session_id = "benchmark_session"
# Xoá cache để đo thời gian chạy LLM thực tế
cache_file = Path("data/cache/semantic_cache.pkl")
if cache_file.exists():
cache_file.unlink()
print("[+] Đã xoá semantic cache cũ.")
test_file = "tai_lieu_test_tom_tat.txt"
if Path(test_file).exists():
print(f"[+] Đang nạp tài liệu {test_file} vào session {session_id}...")
pipeline.process_and_ingest_file_for_session(session_id, test_file)
else:
print(f"[-] Không tìm thấy {test_file}, chạy chay.")
# Danh sách 5 câu hỏi mẫu về báo cáo Vinamilk (bao gồm tiếng Việt và tiếng Anh)
test_queries = [
"Tổng doanh thu thuần của Vinamilk năm 2025 là bao nhiêu?", # Tiếng Việt, liên quan bảng
"Lợi nhuận sau thuế của công ty có tăng trưởng không?", # Tiếng Việt, liên quan số liệu
"What is the total revenue of Vinamilk in 2025?", # Tiếng Anh (test cross-lingual)
"Kế hoạch phát triển bền vững (ESG) của công ty là gì?", # Tiếng Việt, nội dung dạng chữ
"Tóm tắt các rủi ro tài chính chính trong năm qua." # Tiếng Việt, nội dung dài
]
results = []
total_time = 0.0
for i, query in enumerate(test_queries, 1):
print(f"\n--- Câu hỏi {i}: {query}")
start_q = time.time()
# Bỏ qua HTTP, gọi trực tiếp ask()
res = pipeline.ask(query, session_id=session_id)
exec_time = res.get("execution_time_sec", time.time() - start_q)
total_time += exec_time
answer_preview = res.get("answer", "").replace("\n", " ")[:150] + "..."
sources = [s["source_file"] for s in res.get("sources", [])]
print(f" [Latency] {exec_time:.2f}s")
print(f" [Sources] {len(sources)} chunk(s)")
print(f" [Answer] {answer_preview}")
results.append({
"query": query,
"exec_time_sec": exec_time,
"answer_preview": answer_preview,
"sources_count": len(sources)
})
print("\n" + "="*50)
print("📊 KẾT QUẢ TỔNG HỢP")
print("="*50)
print(f"Tổng thời gian cho {len(test_queries)} câu: {total_time:.2f}s")
print(f"Thời gian trung bình/câu: {(total_time / len(test_queries)):.2f}s\n")
# Lưu kết quả ra file JSON để so sánh sau
out_file = Path("benchmark_results.json")
if out_file.exists():
with open(out_file, "r", encoding="utf-8") as f:
old_data = json.load(f)
old_avg = old_data.get("avg_time_sec", 0)
new_avg = total_time / len(test_queries)
if old_avg > 0:
diff = ((old_avg - new_avg) / old_avg) * 100
print(f"📈 So với lần chạy trước (Baseline {old_avg:.2f}s): TỐC ĐỘ CẢI THIỆN {diff:.1f}%")
with open(out_file, "w", encoding="utf-8") as f:
json.dump({
"total_time_sec": total_time,
"avg_time_sec": total_time / len(test_queries),
"details": results
}, f, ensure_ascii=False, indent=2)
if __name__ == "__main__":
run_benchmark()