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()