ForcePilot/backend/package/yuxi/services/evaluation_service.py

651 lines
27 KiB
Python
Raw Normal View History

import json
import os
import re
import uuid
from datetime import datetime
from typing import Any
from yuxi.knowledge import knowledge_base
from yuxi.knowledge.eval.benchmark_generation import dump_benchmark_item, iter_generated_benchmark_items
from yuxi.knowledge.eval.evaluator import aggregate_metrics, evaluate_question
from yuxi.models import select_model
from yuxi.repositories.evaluation_repository import EvaluationRepository
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
from yuxi.services.task_service import TaskContext, tasker
from yuxi.utils import logger
class EvaluationService:
"""RAG评估服务"""
def __init__(self):
self.eval_repo = EvaluationRepository()
self.kb_repo = KnowledgeBaseRepository()
async def _get_benchmark_dir(self, db_id: str) -> str:
"""获取评估基准目录"""
kb_instance = await knowledge_base.aget_kb(db_id)
base_dir = os.path.join(kb_instance.work_dir, db_id)
path = os.path.join(base_dir, "benchmarks")
os.makedirs(path, exist_ok=True)
return path
async def _get_result_dir(self, db_id: str) -> str:
"""获取评估结果目录"""
kb_instance = await knowledge_base.aget_kb(db_id)
base_dir = os.path.join(kb_instance.work_dir, db_id)
path = os.path.join(base_dir, "results")
os.makedirs(path, exist_ok=True)
return path
# 已移除基准回退逻辑,统一使用集中元数据
# 已移除结果回退逻辑,统一通过 db_id 定位
async def upload_benchmark(
self, db_id: str, file_content: bytes, filename: str, name: str, description: str, created_by: str
) -> dict[str, Any]:
"""上传评估基准文件"""
try:
content_str = file_content.decode("utf-8")
questions = []
has_gold_chunks = False
has_gold_answers = False
# 解析 JSONL
for line_num, line in enumerate(content_str.strip().split("\n"), 1):
if not line.strip():
continue
try:
item = json.loads(line)
if "query" not in item:
raise ValueError(f"{line_num}行缺少必需的'query'字段")
if item.get("gold_chunk_ids"):
has_gold_chunks = True
if item.get("gold_answer"):
has_gold_answers = True
questions.append(item)
except json.JSONDecodeError as e:
raise ValueError(f"{line_num}行JSON格式错误: {str(e)}")
if not questions:
raise ValueError("文件中没有有效的问题数据")
benchmark_id = f"benchmark_{uuid.uuid4().hex[:8]}"
benchmark_dir = await self._get_benchmark_dir(db_id)
# 保存数据文件 (.jsonl)
data_file_path = os.path.join(benchmark_dir, f"{benchmark_id}.jsonl")
with open(data_file_path, "w", encoding="utf-8") as f:
f.write(content_str)
meta = {
"id": benchmark_id,
"benchmark_id": benchmark_id,
"name": name,
"description": description,
"db_id": db_id,
"question_count": len(questions),
"has_gold_chunks": has_gold_chunks,
"has_gold_answers": has_gold_answers,
"benchmark_file": data_file_path,
"created_by": created_by,
"created_at": datetime.utcnow().isoformat(),
"updated_at": datetime.utcnow().isoformat(),
}
await self.eval_repo.create_benchmark(
{
"benchmark_id": benchmark_id,
"db_id": db_id,
"name": name,
"description": description,
"question_count": len(questions),
"has_gold_chunks": has_gold_chunks,
"has_gold_answers": has_gold_answers,
"data_file_path": data_file_path,
"created_by": created_by,
}
)
return meta
except Exception as e:
logger.error(f"上传评估基准失败: {e}")
raise
async def get_benchmarks(self, db_id: str) -> list[dict[str, Any]]:
"""获取知识库的评估基准列表"""
try:
rows = await self.eval_repo.list_benchmarks(db_id)
return [
{
"id": row.benchmark_id,
"benchmark_id": row.benchmark_id,
"name": row.name,
"description": row.description,
"db_id": row.db_id,
"question_count": row.question_count,
"has_gold_chunks": row.has_gold_chunks,
"has_gold_answers": row.has_gold_answers,
"benchmark_file": row.data_file_path,
"created_by": row.created_by,
"created_at": row.created_at.isoformat() if row.created_at else None,
"updated_at": row.updated_at.isoformat() if row.updated_at else None,
}
for row in rows
]
except Exception as e:
logger.error(f"获取评估基准列表失败: {e}")
raise
async def get_benchmark_detail(self, benchmark_id: str) -> dict[str, Any]:
"""获取评估基准详情 (包含问题列表)"""
try:
row = await self.eval_repo.get_benchmark(benchmark_id)
if row is None:
raise ValueError("Benchmark not found")
questions = []
if row.data_file_path and os.path.exists(row.data_file_path):
with open(row.data_file_path, encoding="utf-8") as f:
for line in f:
if line.strip():
questions.append(json.loads(line))
return {
"id": row.benchmark_id,
"benchmark_id": row.benchmark_id,
"name": row.name,
"description": row.description,
"db_id": row.db_id,
"question_count": row.question_count,
"has_gold_chunks": row.has_gold_chunks,
"has_gold_answers": row.has_gold_answers,
"benchmark_file": row.data_file_path,
"created_by": row.created_by,
"created_at": row.created_at.isoformat() if row.created_at else None,
"updated_at": row.updated_at.isoformat() if row.updated_at else None,
"questions": questions,
}
except Exception as e:
logger.error(f"获取评估基准详情失败: {e}")
raise
2025-12-14 22:44:18 +08:00
async def get_benchmark_detail_by_db(
self, db_id: str, benchmark_id: str, page: int = 1, page_size: int = 10
) -> dict[str, Any]:
"""根据 db_id 获取评估基准详情(支持分页)"""
try:
row = await self.eval_repo.get_benchmark(benchmark_id)
if row is None or row.db_id != db_id:
raise ValueError("Benchmark not found")
data_file_path = row.data_file_path
total_questions = row.question_count or 0
questions = []
if data_file_path and os.path.exists(data_file_path):
# 计算分页范围
start_index = (page - 1) * page_size
end_index = start_index + page_size
# 读取指定范围的问题
with open(data_file_path, encoding="utf-8") as f:
current_index = 0
for line in f:
if not line.strip():
continue
# 只处理指定范围内的问题
if current_index >= start_index and current_index < end_index:
questions.append(json.loads(line))
elif current_index >= end_index:
break # 已经读取到足够的问题,停止读取
current_index += 1
# 计算分页信息
total_pages = (total_questions + page_size - 1) // page_size
return {
"id": row.benchmark_id,
"benchmark_id": row.benchmark_id,
"name": row.name,
"description": row.description,
"db_id": row.db_id,
"question_count": row.question_count,
"has_gold_chunks": row.has_gold_chunks,
"has_gold_answers": row.has_gold_answers,
"benchmark_file": data_file_path,
"created_by": row.created_by,
"created_at": row.created_at.isoformat() if row.created_at else None,
"updated_at": row.updated_at.isoformat() if row.updated_at else None,
"questions": questions,
"pagination": {
"current_page": page,
"page_size": page_size,
"total_questions": total_questions,
"total_pages": total_pages,
"has_next": page < total_pages,
"has_prev": page > 1,
},
}
except Exception as e:
logger.error(f"获取评估基准详情失败: {e}")
raise
async def get_benchmark_download_info(self, benchmark_id: str) -> dict[str, str]:
"""获取评估基准下载信息"""
row = await self.eval_repo.get_benchmark(benchmark_id)
if row is None:
raise ValueError("Benchmark not found")
data_file_path = row.data_file_path or ""
if not data_file_path or not os.path.exists(data_file_path):
raise ValueError("Benchmark file not found")
filename_base = (row.name or "").strip()
if not filename_base:
filename_base = row.benchmark_id
filename_base = re.sub(r"[\\/:*?\"<>|]+", "_", filename_base).strip()
if not filename_base or filename_base in {".", ".."}:
filename_base = row.benchmark_id
if not filename_base.endswith(".jsonl"):
filename_base = f"{filename_base}.jsonl"
return {"file_path": data_file_path, "filename": filename_base}
async def delete_benchmark(self, benchmark_id: str) -> None:
"""删除评估基准"""
try:
row = await self.eval_repo.get_benchmark(benchmark_id)
if row is None:
raise ValueError("Benchmark not found")
if row.data_file_path and os.path.exists(row.data_file_path):
os.remove(row.data_file_path)
await self.eval_repo.delete_benchmark(benchmark_id)
logger.info(f"成功删除评估基准: {benchmark_id}")
return
except Exception as e:
logger.error(f"删除评估基准失败: {e}")
raise
async def delete_evaluation_result(self, task_id: str, db_id: str) -> None:
"""删除评估结果"""
if not task_id:
raise ValueError("task_id is required")
await self.delete_evaluation_result_by_db(db_id, task_id)
async def generate_benchmark(self, db_id: str, params: dict[str, Any], created_by: str) -> dict[str, Any]:
task_id = f"gen_benchmark_{uuid.uuid4().hex[:8]}"
await tasker.enqueue(
name="生成评估基准",
task_type="benchmark_generation",
payload={"task_id": task_id, "db_id": db_id, "created_by": created_by, **params},
coroutine=self._generate_benchmark_task,
)
return {"task_id": task_id, "message": "基准生成任务已提交"}
async def _generate_benchmark_task(self, context: TaskContext):
await context.set_progress(0, "初始化")
task = context._tasker._tasks.get(context.task_id)
payload = task.payload if task else {}
db_id = payload.get("db_id")
name = payload.get("name", "自动生成评估基准")
description = payload.get("description", "")
count = int(payload.get("count", 10))
neighbors_count = int(payload.get("neighbors_count", 1))
llm_model_spec = payload.get("llm_model_spec")
kb_instance = await knowledge_base.aget_kb(db_id)
if not kb_instance:
await context.set_message("知识库不存在")
raise ValueError("Knowledge Base not found")
if kb_instance.kb_type != "milvus":
await context.set_message("仅支持 commonrag/Milvus 类型知识库生成评估基准")
raise ValueError("Unsupported KB type for benchmark generation")
benchmark_id = f"benchmark_{uuid.uuid4().hex[:8]}"
bench_dir = await self._get_benchmark_dir(db_id)
data_file_path = os.path.join(bench_dir, f"{benchmark_id}.jsonl")
generated = 0
try:
with open(data_file_path, "w", encoding="utf-8") as f:
async for item in iter_generated_benchmark_items(
kb_instance=kb_instance,
db_id=db_id,
count=count,
neighbors_count=neighbors_count,
llm_model_spec=llm_model_spec,
progress_cb=context.set_progress,
):
f.write(dump_benchmark_item(item))
generated += 1
except ValueError as e:
if str(e) == "No chunks found in knowledge base":
await context.set_message("知识库为空或未解析到chunks")
raise
await self.eval_repo.create_benchmark(
{
"benchmark_id": benchmark_id,
"db_id": db_id,
"name": name,
"description": description,
"question_count": generated,
"has_gold_chunks": True,
"has_gold_answers": True,
"data_file_path": data_file_path,
"created_by": payload.get("created_by"),
}
)
await context.set_progress(100, "完成")
async def run_evaluation(
self, db_id: str, benchmark_id: str, model_config: dict[str, Any] = None, created_by: str = "system"
) -> str:
"""运行RAG评估"""
try:
task_id = f"eval_{uuid.uuid4().hex[:8]}"
benchmark_row = await self.eval_repo.get_benchmark(benchmark_id)
if benchmark_row is None or benchmark_row.db_id != db_id:
raise ValueError("Benchmark not found")
# 从知识库元数据中获取检索配置
retrieval_config = {}
try:
kb_row = await self.kb_repo.get_by_id(db_id)
query_params = (kb_row.query_params if kb_row else None) or {}
retrieval_config = query_params.get("options", {}) if isinstance(query_params, dict) else {}
if not retrieval_config:
kb_instance = await knowledge_base.aget_kb(db_id)
if kb_instance:
retrieval_config = kb_instance._get_default_query_params(db_id).get("options", {})
logger.info(f"从知识库 {db_id} 加载检索配置: {list(retrieval_config.keys())}")
except Exception as e:
logger.error(f"获取知识库检索配置失败: {e}")
# 使用空配置作为默认值
# 合并前端传递的模型配置
if model_config:
retrieval_config.update(model_config)
await self.eval_repo.create_result(
{
"task_id": task_id,
"db_id": db_id,
"benchmark_id": benchmark_id,
"status": "running",
"retrieval_config": retrieval_config,
"metrics": {},
"overall_score": None,
"total_questions": benchmark_row.question_count or 0,
"completed_questions": 0,
"started_at": datetime.utcnow(),
"completed_at": None,
"created_by": created_by,
}
)
await tasker.enqueue(
name=f"RAG评估({benchmark_row.name})",
task_type="rag_evaluation",
payload={
"task_id": task_id,
"db_id": db_id,
"benchmark_id": benchmark_id,
"retrieval_config": retrieval_config,
"created_by": created_by,
},
coroutine=self._run_evaluation_task,
)
return task_id
except Exception as e:
logger.error(f"启动评估失败: {e}")
raise
async def _run_evaluation_task(self, context: TaskContext):
"""运行评估任务"""
try:
task = context._tasker._tasks.get(context.task_id)
if not task:
raise ValueError("Task not found")
payload = task.payload
task_id = payload["task_id"]
db_id = payload["db_id"]
benchmark_id = payload["benchmark_id"]
retrieval_config = payload["retrieval_config"]
# 加载基准数据
await context.set_progress(5, "加载基准数据")
benchmark_row = await self.eval_repo.get_benchmark(benchmark_id)
if benchmark_row is None or benchmark_row.db_id != db_id:
raise ValueError("Benchmark not found")
data_path = benchmark_row.data_file_path
if not data_path or not os.path.exists(data_path):
raise ValueError("Benchmark file not found")
benchmark_data = []
with open(data_path, encoding="utf-8") as f:
for line in f:
if line.strip():
benchmark_data.append(json.loads(line))
# 开始评估
kb_instance = await knowledge_base.aget_kb(db_id)
if not kb_instance:
raise ValueError(f"Knowledge Base {db_id} not found")
if kb_instance.kb_type == "lightrag":
raise ValueError("暂不支持对 LightRAG 类型的知识库进行 RAG 评估")
# 初始化 Judge LLM
judge_llm = None
if benchmark_row.has_gold_answers:
# 优先使用配置中的 judge_llm否则回退到 answer_llm或者默认
judge_model_spec = retrieval_config.get("judge_llm") or retrieval_config.get("answer_llm")
if judge_model_spec:
try:
logger.debug(f"Initializing Judge LLM: {judge_model_spec}")
judge_llm = select_model(model_spec=judge_model_spec)
except Exception as e:
logger.error(f"Failed to load judge LLM: {e}")
total_questions = len(benchmark_data)
all_retrieval_metrics = []
all_answer_metrics = []
async def update_result_db(
status: str | None = None, completed: int | None = None, metrics=None, final_score=None
):
payload = {}
if status is not None:
payload["status"] = status
if status in ["completed", "failed"]:
payload["completed_at"] = datetime.utcnow()
if completed is not None:
payload["completed_questions"] = completed
if metrics is not None:
payload["metrics"] = metrics
if final_score is not None:
payload["overall_score"] = final_score
if payload:
await self.eval_repo.update_result(task_id, payload)
for i, question_data in enumerate(benchmark_data):
await context.raise_if_cancelled()
progress = 10 + (i / total_questions) * 80
await context.set_progress(progress, f"评估 {i + 1}/{total_questions}")
question_result = await evaluate_question(
kb_instance=kb_instance,
db_id=db_id,
question_data=question_data,
retrieval_config=retrieval_config,
has_gold_chunks=benchmark_row.has_gold_chunks,
has_gold_answers=benchmark_row.has_gold_answers,
judge_llm=judge_llm,
select_model_fn=select_model,
)
if benchmark_row.has_gold_chunks and question_data.get("gold_chunk_ids"):
all_retrieval_metrics.append(question_result["retrieval_scores"])
if benchmark_row.has_gold_answers and question_data.get("gold_answer") and judge_llm:
all_answer_metrics.append(question_result["answer_scores"])
await self.eval_repo.upsert_result_detail(
task_id=task_id,
query_index=i,
data=question_result["detail"],
)
current_overall_metrics, _ = aggregate_metrics(all_retrieval_metrics, all_answer_metrics)
await context.set_result(
{
"current_metrics": current_overall_metrics,
"completed_questions": i + 1,
"total_questions": total_questions,
}
)
if (i + 1) % 5 == 0 or (i + 1) == total_questions:
await update_result_db(completed=i + 1)
await context.set_progress(95, "计算最终指标")
overall_metrics, overall_score = aggregate_metrics(
all_retrieval_metrics, all_answer_metrics, include_overall_score=True
)
await update_result_db(
status="completed",
completed=total_questions,
metrics=overall_metrics,
final_score=overall_score,
)
await context.set_progress(100, "完成")
except Exception as e:
logger.error(f"Task failed: {e}")
try:
if "payload" in locals():
await self.eval_repo.update_result(
payload["task_id"],
{"status": "failed", "metrics": {"error": str(e)}, "completed_at": datetime.utcnow()},
)
except Exception as exc:
logger.error(f"Error updating result record: {exc}")
await context.set_message(f"Error: {str(e)}")
raise
async def get_evaluation_results(self, task_id: str, db_id: str) -> dict[str, Any]:
"""获取评估结果"""
if not task_id:
raise ValueError("task_id is required")
return await self.get_evaluation_results_by_db(db_id, task_id)
async def get_evaluation_history(self, db_id: str) -> list[dict[str, Any]]:
"""获取知识库的评估历史记录"""
try:
rows = await self.eval_repo.list_results(db_id)
return [
{
"task_id": row.task_id,
"benchmark_id": row.benchmark_id,
"status": row.status,
"started_at": row.started_at.isoformat() if row.started_at else None,
"completed_at": row.completed_at.isoformat() if row.completed_at else None,
"total_questions": row.total_questions,
"completed_questions": row.completed_questions,
"overall_score": row.overall_score,
"retrieval_config": row.retrieval_config or {},
"metrics": row.metrics or {},
}
for row in rows
]
except Exception as e:
logger.error(f"获取评估历史失败: {e}")
raise
# 索引与回退逻辑已移除,统一通过 db_id 定位
async def get_evaluation_results_by_db(
2025-12-14 22:44:18 +08:00
self, db_id: str, task_id: str, page: int = 1, page_size: int = 20, error_only: bool = False
) -> dict[str, Any]:
if not re.match(r"^eval_[a-f0-9]{8}$", task_id):
raise ValueError("Invalid task_id format")
row = await self.eval_repo.get_result(task_id)
if row is None or row.db_id != db_id:
task = await tasker.get_task(task_id)
if task:
return {"task_id": task_id, "status": task.status, "progress": task.progress, "message": task.message}
raise ValueError(f"Result not found for task {task_id}")
details = await self.eval_repo.list_result_details(task_id)
all_results = [
{
"query": d.query_text,
"gold_chunk_ids": d.gold_chunk_ids,
"gold_answer": d.gold_answer,
"generated_answer": d.generated_answer,
"retrieved_chunks": d.retrieved_chunks,
"metrics": d.metrics or {},
}
for d in details
]
if error_only:
filtered_results = []
for item in all_results:
if item.get("metrics", {}).get("score", 1.0) <= 0.5:
filtered_results.append(item)
continue
metrics = item.get("metrics", {})
has_low_recall = any(metrics.get(k, 1.0) < 0.3 for k in metrics if k.startswith("recall@"))
if has_low_recall:
filtered_results.append(item)
all_results = filtered_results
total = len(all_results)
start_idx = (page - 1) * page_size
end_idx = start_idx + page_size
paged_results = all_results[start_idx:end_idx]
return {
"task_id": row.task_id,
"status": row.status,
"started_at": row.started_at.isoformat() if row.started_at else None,
"completed_at": row.completed_at.isoformat() if row.completed_at else None,
"total_questions": row.total_questions or 0,
"completed_questions": row.completed_questions or 0,
"overall_score": row.overall_score,
"retrieval_config": row.retrieval_config or {},
"interim_results": paged_results,
"pagination": {
"current_page": page,
"page_size": page_size,
"total": total,
"total_pages": (total + page_size - 1) // page_size,
"error_only": error_only,
},
}
async def delete_evaluation_result_by_db(self, db_id: str, task_id: str) -> None:
if not re.match(r"^eval_[a-f0-9]{8}$", task_id):
raise ValueError("Invalid task_id format")
row = await self.eval_repo.get_result(task_id)
if row is None or row.db_id != db_id:
raise ValueError("Result not found")
await self.eval_repo.delete_result(task_id)
logger.info(f"成功删除评估结果: {task_id}")
return