refactor(eval): 重构评估指标和基准生成
- 引入了一个新的指标计算模块,用于检索和答案评估。 -通过整合基准生成和评估指标来简化评估服务。 -为新指标和基准生成功能添加了单元测试。 - 更新了前端组件,以利用新的图标,并改进了评估基准的样式。 -将与评估相关的逻辑整合为专门的“知识/评估”结构,以更好地组织。
This commit is contained in:
parent
0890e504d8
commit
179b048f07
1
backend/package/yuxi/knowledge/eval/__init__.py
Normal file
1
backend/package/yuxi/knowledge/eval/__init__.py
Normal file
@ -0,0 +1 @@
|
||||
"""知识库评估核心能力。"""
|
||||
141
backend/package/yuxi/knowledge/eval/benchmark_generation.py
Normal file
141
backend/package/yuxi/knowledge/eval/benchmark_generation.py
Normal file
@ -0,0 +1,141 @@
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from typing import Any
|
||||
|
||||
import json_repair
|
||||
|
||||
from yuxi import config
|
||||
from yuxi.models import select_embedding_model, select_model
|
||||
from yuxi.utils import logger
|
||||
|
||||
|
||||
async def collect_kb_chunks(kb_instance: Any, db_id: str) -> list[dict[str, Any]]:
|
||||
chunks = []
|
||||
for fid, finfo in kb_instance.files_meta.items():
|
||||
if finfo.get("database_id") != db_id:
|
||||
continue
|
||||
try:
|
||||
content_info = await kb_instance.get_file_content(db_id, fid)
|
||||
for line in content_info.get("lines", []):
|
||||
chunks.append(
|
||||
{
|
||||
"id": line.get("id"),
|
||||
"content": line.get("content", ""),
|
||||
"file_id": fid,
|
||||
"chunk_index": line.get("chunk_order_index"),
|
||||
}
|
||||
)
|
||||
except Exception:
|
||||
continue
|
||||
return chunks
|
||||
|
||||
|
||||
def clamp_neighbors_count(neighbors_count: int) -> int:
|
||||
return min(max(neighbors_count, 0), 10)
|
||||
|
||||
|
||||
def cosine_similarity(a: list[float], b: list[float], na: float, nb: float) -> float:
|
||||
s = 0.0
|
||||
for i in range(len(a)):
|
||||
s += a[i] * b[i]
|
||||
return s / (na * nb)
|
||||
|
||||
|
||||
def select_neighbor_indices(
|
||||
anchor_idx: int, embeddings: list[list[float]], norms: list[float], neighbors_count: int
|
||||
) -> list[int]:
|
||||
anchor_embedding = embeddings[anchor_idx]
|
||||
anchor_norm = norms[anchor_idx]
|
||||
sims = []
|
||||
for idx in range(len(embeddings)):
|
||||
if idx == anchor_idx:
|
||||
continue
|
||||
score = cosine_similarity(anchor_embedding, embeddings[idx], anchor_norm, norms[idx])
|
||||
sims.append((idx, score))
|
||||
sims.sort(key=lambda x: x[1], reverse=True)
|
||||
return [idx for idx, _ in sims[:neighbors_count]]
|
||||
|
||||
|
||||
def build_benchmark_generation_prompt(ctx_items: list[tuple[str, str]]) -> str:
|
||||
context_text = "\n\n".join([f"片段ID={cid}\n{content}" for cid, content in ctx_items])
|
||||
return (
|
||||
"你将基于以下上下文生成一个可由上下文准确回答的问题与标准答案。"
|
||||
"仅返回一个JSON对象,不要包含其他文字。"
|
||||
"键为 query、gold_answer、gold_chunk_ids。gold_chunk_ids 必须是上述上下文片段的ID子集。\n\n"
|
||||
"上下文:\n" + context_text + "\n"
|
||||
)
|
||||
|
||||
|
||||
async def iter_generated_benchmark_items(
|
||||
*,
|
||||
kb_instance: Any,
|
||||
db_id: str,
|
||||
count: int,
|
||||
neighbors_count: int,
|
||||
embedding_model_id: str,
|
||||
llm_model_spec: Any,
|
||||
progress_cb: Callable[[int, str], Any] | None = None,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
if progress_cb:
|
||||
await progress_cb(5, "加载chunks")
|
||||
|
||||
all_chunks = await collect_kb_chunks(kb_instance, db_id)
|
||||
if not all_chunks:
|
||||
raise ValueError("No chunks found in knowledge base")
|
||||
|
||||
if progress_cb:
|
||||
await progress_cb(15, "向量化")
|
||||
|
||||
contents = [chunk["content"] for chunk in all_chunks]
|
||||
embed_model = select_embedding_model(embedding_model_id)
|
||||
batch_size = int(getattr(embed_model, "batch_size", 40) or 40)
|
||||
if embedding_model_id in config.embed_model_names:
|
||||
batch_size = config.embed_model_names[embedding_model_id].batch_size
|
||||
|
||||
embeddings = await embed_model.abatch_encode(contents, batch_size=batch_size)
|
||||
norms = [math.sqrt(sum(x * x for x in vec)) or 1.0 for vec in embeddings]
|
||||
llm = select_model(model_spec=llm_model_spec)
|
||||
neighbors_count = clamp_neighbors_count(neighbors_count)
|
||||
generated = 0
|
||||
attempts = 0
|
||||
max_attempts = max(count * 5, 50)
|
||||
|
||||
if progress_cb:
|
||||
await progress_cb(0, "准备生成样本")
|
||||
|
||||
while generated < count and attempts < max_attempts:
|
||||
attempts += 1
|
||||
anchor_idx = random.randrange(len(all_chunks))
|
||||
neighbor_indices = select_neighbor_indices(anchor_idx, embeddings, norms, neighbors_count)
|
||||
ctx_items = [(all_chunks[anchor_idx]["id"], all_chunks[anchor_idx]["content"])]
|
||||
ctx_items.extend((all_chunks[idx]["id"], all_chunks[idx]["content"]) for idx in neighbor_indices)
|
||||
allowed_ids = {cid for cid, _ in ctx_items}
|
||||
|
||||
try:
|
||||
resp = await llm.call(build_benchmark_generation_prompt(ctx_items), False)
|
||||
obj = json_repair.loads(resp.content if resp else "")
|
||||
query = obj.get("query")
|
||||
answer = obj.get("gold_answer")
|
||||
gold_ids = obj.get("gold_chunk_ids")
|
||||
if not query or not answer or not isinstance(gold_ids, list):
|
||||
logger.warning(f"Generated JSON missing fields or invalid format: {obj}")
|
||||
continue
|
||||
|
||||
gold_ids = [str(item) for item in gold_ids if str(item) in allowed_ids]
|
||||
if not gold_ids:
|
||||
logger.warning("Generated gold_chunk_ids not found in allowed context")
|
||||
continue
|
||||
|
||||
generated += 1
|
||||
if progress_cb:
|
||||
await progress_cb(0 + int(99 * generated / max(count, 1)), f"已生成 {generated}/{count}")
|
||||
yield {"query": query, "gold_chunk_ids": gold_ids, "gold_answer": answer}
|
||||
except Exception as e:
|
||||
logger.warning(f"Benchmark generation failed for one item: {e}")
|
||||
continue
|
||||
|
||||
|
||||
def dump_benchmark_item(item: dict[str, Any]) -> str:
|
||||
return json.dumps(item, ensure_ascii=False) + "\n"
|
||||
137
backend/package/yuxi/knowledge/eval/evaluator.py
Normal file
137
backend/package/yuxi/knowledge/eval/evaluator.py
Normal file
@ -0,0 +1,137 @@
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from yuxi.knowledge.eval.metrics import EvaluationMetricsCalculator
|
||||
from yuxi.utils import logger
|
||||
|
||||
|
||||
def normalize_query_result(query_result: Any) -> tuple[str, list[dict[str, Any]]]:
|
||||
if isinstance(query_result, dict):
|
||||
return query_result.get("answer", ""), query_result.get("retrieved_chunks", [])
|
||||
if isinstance(query_result, list):
|
||||
return "", query_result
|
||||
return "", []
|
||||
|
||||
|
||||
def build_answer_prompt(query: str, retrieved_chunks: list[dict[str, Any]], max_docs: int = 5) -> str:
|
||||
context_docs = []
|
||||
for idx, chunk in enumerate(retrieved_chunks[:max_docs]):
|
||||
content = chunk.get("content", "")
|
||||
if content:
|
||||
context_docs.append(f"文档 {idx + 1}:\n{content}")
|
||||
|
||||
context_text = "\\n\\n".join(context_docs)
|
||||
return (
|
||||
f"基于以下上下文信息,请回答用户的问题。\n\n"
|
||||
f"上下文信息:{context_text}\n\n"
|
||||
f"用户问题:{query}\n\n"
|
||||
"请根据上下文信息准确回答问题。\n\n"
|
||||
"如果上下文中缺少相关信息,请回答“信息不足,无法回答”。\n\n"
|
||||
)
|
||||
|
||||
|
||||
async def generate_answer_if_needed(
|
||||
*,
|
||||
query: str,
|
||||
generated_answer: str,
|
||||
retrieved_chunks: list[dict[str, Any]],
|
||||
retrieval_config: dict[str, Any],
|
||||
select_model_fn: Callable[..., Any],
|
||||
) -> str:
|
||||
if generated_answer:
|
||||
return generated_answer
|
||||
if not retrieved_chunks or not retrieval_config.get("answer_llm"):
|
||||
return ""
|
||||
|
||||
logger.debug(f"使用 LLM {retrieval_config.get('answer_llm')} 生成答案...")
|
||||
try:
|
||||
llm = select_model_fn(model_spec=retrieval_config["answer_llm"])
|
||||
response = await llm.call(build_answer_prompt(query, retrieved_chunks), stream=False)
|
||||
generated_answer = response.content if response else ""
|
||||
logger.debug(f"LLM 生成的答案长度: {len(generated_answer) if generated_answer else 0}")
|
||||
return generated_answer
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 生成答案失败: {e}")
|
||||
return ""
|
||||
|
||||
|
||||
async def evaluate_question(
|
||||
*,
|
||||
kb_instance: Any,
|
||||
db_id: str,
|
||||
question_data: dict[str, Any],
|
||||
retrieval_config: dict[str, Any],
|
||||
has_gold_chunks: bool,
|
||||
has_gold_answers: bool,
|
||||
judge_llm: Any | None,
|
||||
select_model_fn: Callable[..., Any],
|
||||
) -> dict[str, Any]:
|
||||
query = question_data["query"]
|
||||
query_result = await kb_instance.aquery(query, db_id, **retrieval_config)
|
||||
generated_answer, retrieved_chunks = normalize_query_result(query_result)
|
||||
generated_answer = await generate_answer_if_needed(
|
||||
query=query,
|
||||
generated_answer=generated_answer,
|
||||
retrieved_chunks=retrieved_chunks,
|
||||
retrieval_config=retrieval_config,
|
||||
select_model_fn=select_model_fn,
|
||||
)
|
||||
|
||||
current_metrics = {}
|
||||
retrieval_scores = {}
|
||||
answer_scores = {}
|
||||
|
||||
if has_gold_chunks and question_data.get("gold_chunk_ids"):
|
||||
retrieval_scores = EvaluationMetricsCalculator.calculate_retrieval_metrics(
|
||||
retrieved_chunks, question_data["gold_chunk_ids"]
|
||||
)
|
||||
current_metrics.update(retrieval_scores)
|
||||
|
||||
if has_gold_answers and question_data.get("gold_answer"):
|
||||
if judge_llm:
|
||||
answer_scores = await EvaluationMetricsCalculator.calculate_answer_metrics(
|
||||
query=query,
|
||||
generated_answer=generated_answer,
|
||||
gold_answer=question_data["gold_answer"],
|
||||
judge_llm=judge_llm,
|
||||
)
|
||||
current_metrics.update(answer_scores)
|
||||
else:
|
||||
logger.warning("需要计算答案指标但未配置 Judge LLM")
|
||||
|
||||
return {
|
||||
"detail": {
|
||||
"query_text": query,
|
||||
"gold_chunk_ids": question_data.get("gold_chunk_ids"),
|
||||
"gold_answer": question_data.get("gold_answer"),
|
||||
"generated_answer": generated_answer,
|
||||
"retrieved_chunks": retrieved_chunks,
|
||||
"metrics": current_metrics,
|
||||
},
|
||||
"retrieval_scores": retrieval_scores,
|
||||
"answer_scores": answer_scores,
|
||||
}
|
||||
|
||||
|
||||
def aggregate_metrics(
|
||||
retrieval_metrics_list: list[dict[str, float]],
|
||||
answer_metrics_list: list[dict[str, Any]],
|
||||
*,
|
||||
include_overall_score: bool = False,
|
||||
) -> tuple[dict[str, Any], float]:
|
||||
overall_metrics = {}
|
||||
|
||||
if retrieval_metrics_list:
|
||||
keys = retrieval_metrics_list[0].keys()
|
||||
for key in keys:
|
||||
overall_metrics[key] = sum(m.get(key, 0) for m in retrieval_metrics_list) / len(retrieval_metrics_list)
|
||||
|
||||
if answer_metrics_list:
|
||||
scores = [m.get("score", 0) for m in answer_metrics_list]
|
||||
overall_metrics["answer_correctness"] = sum(scores) / len(scores) if scores else 0.0
|
||||
|
||||
overall_score = EvaluationMetricsCalculator.calculate_overall_score(retrieval_metrics_list, answer_metrics_list)
|
||||
if include_overall_score:
|
||||
overall_metrics["overall_score"] = overall_score
|
||||
|
||||
return overall_metrics, overall_score
|
||||
152
backend/package/yuxi/knowledge/eval/metrics.py
Normal file
152
backend/package/yuxi/knowledge/eval/metrics.py
Normal file
@ -0,0 +1,152 @@
|
||||
"""
|
||||
RAG评估指标计算工具
|
||||
简化版:只保留Recall/F1(检索)和 LLM Judge(答案准确性)
|
||||
"""
|
||||
|
||||
import json
|
||||
import textwrap
|
||||
from typing import Any
|
||||
|
||||
from yuxi.utils import logger
|
||||
|
||||
|
||||
class RetrievalMetrics:
|
||||
"""检索评估指标计算"""
|
||||
|
||||
@staticmethod
|
||||
def precision_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
|
||||
"""计算Precision@K"""
|
||||
if not retrieved_ids[:k]:
|
||||
return 0.0
|
||||
retrieved_set = set(retrieved_ids[:k])
|
||||
relevant_set = set(relevant_ids)
|
||||
return len(retrieved_set & relevant_set) / k
|
||||
|
||||
@staticmethod
|
||||
def recall_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
|
||||
"""计算Recall@K"""
|
||||
if not relevant_ids:
|
||||
return 0.0
|
||||
retrieved_set = set(retrieved_ids[:k])
|
||||
relevant_set = set(relevant_ids)
|
||||
return len(retrieved_set & relevant_set) / len(relevant_set)
|
||||
|
||||
@staticmethod
|
||||
def f1_score_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
|
||||
"""计算F1@K"""
|
||||
precision = RetrievalMetrics.precision_at_k(retrieved_ids, relevant_ids, k)
|
||||
recall = RetrievalMetrics.recall_at_k(retrieved_ids, relevant_ids, k)
|
||||
if precision + recall == 0:
|
||||
return 0.0
|
||||
return 2 * precision * recall / (precision + recall)
|
||||
|
||||
|
||||
class AnswerMetrics:
|
||||
"""答案评估指标计算"""
|
||||
|
||||
@staticmethod
|
||||
async def judge_correctness(query: str, generated_answer: str, gold_answer: str, judge_llm: Any) -> dict[str, Any]:
|
||||
"""
|
||||
使用LLM判断生成的答案是否正确
|
||||
"""
|
||||
if not generated_answer:
|
||||
return {"score": 0.0, "reasoning": "未生成答案"}
|
||||
if not gold_answer:
|
||||
return {"score": 0.0, "reasoning": "无参考答案"}
|
||||
|
||||
prompt = textwrap.dedent(f"""你是一个公正的评判者,请评估AI生成的答案相对于标准答案的准确性。
|
||||
|
||||
问题:{query}
|
||||
|
||||
标准答案:
|
||||
{gold_answer}
|
||||
|
||||
AI生成的答案:
|
||||
{generated_answer}
|
||||
|
||||
请判断AI生成的答案是否在事实层面与标准答案一致。
|
||||
忽略措辞、标点符号或格式上的细微差异。
|
||||
只关注核心事实是否准确包含。
|
||||
|
||||
请返回以下JSON格式的结果(不要包含其他文本):
|
||||
{{
|
||||
"score": 1.0, // 如果答案正确返回 1.0, 错误返回 0.0
|
||||
"reasoning": "简要说明判定理由"
|
||||
}}
|
||||
""")
|
||||
try:
|
||||
response = await judge_llm.call(prompt, stream=False)
|
||||
content = response.content.strip()
|
||||
|
||||
# 尝试清理可能的 markdown 代码块
|
||||
if content.startswith("```json"):
|
||||
content = content[7:]
|
||||
if content.endswith("```"):
|
||||
content = content[:-3]
|
||||
content = content.strip()
|
||||
|
||||
result = json.loads(content)
|
||||
return {"score": float(result.get("score", 0.0)), "reasoning": result.get("reasoning", "")}
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 评判失败: {e}")
|
||||
return {"score": 0.0, "reasoning": f"评判出错: {str(e)}"}
|
||||
|
||||
|
||||
class EvaluationMetricsCalculator:
|
||||
"""综合评估指标计算器"""
|
||||
|
||||
@staticmethod
|
||||
def calculate_retrieval_metrics(
|
||||
retrieved_chunks: list[dict[str, Any]], gold_chunk_ids: list[str], k_values: list[int] = [1, 3, 5, 10]
|
||||
) -> dict[str, float]:
|
||||
"""计算检索指标 (Recall, F1)"""
|
||||
if not retrieved_chunks or not gold_chunk_ids:
|
||||
return {}
|
||||
|
||||
# 提取 ID
|
||||
retrieved_ids = []
|
||||
for chunk in retrieved_chunks:
|
||||
chunk_id = chunk.get("chunk_id") or chunk.get("metadata", {}).get("chunk_id")
|
||||
retrieved_ids.append(str(chunk_id) if chunk_id else "")
|
||||
|
||||
metrics = {}
|
||||
for k in k_values:
|
||||
metrics[f"recall@{k}"] = RetrievalMetrics.recall_at_k(retrieved_ids, gold_chunk_ids, k)
|
||||
metrics[f"f1@{k}"] = RetrievalMetrics.f1_score_at_k(retrieved_ids, gold_chunk_ids, k)
|
||||
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
async def calculate_answer_metrics(
|
||||
query: str, generated_answer: str, gold_answer: str, judge_llm: Any = None
|
||||
) -> dict[str, Any]:
|
||||
"""计算答案指标 (LLM Judge)"""
|
||||
if not judge_llm:
|
||||
return {}
|
||||
|
||||
return await AnswerMetrics.judge_correctness(query, generated_answer, gold_answer, judge_llm)
|
||||
|
||||
@staticmethod
|
||||
def calculate_overall_score(
|
||||
retrieval_metrics_list: list[dict[str, float]], answer_metrics_list: list[dict[str, Any]]
|
||||
) -> float:
|
||||
"""计算整体平均分"""
|
||||
total_score = 0.0
|
||||
count = 0
|
||||
|
||||
# 简单的平均策略:将所有retrieval metric的值和answer metric的score一起平均
|
||||
# 用户可能希望分开看,但calculate_overall_score返回一个单值。
|
||||
|
||||
# 计算检索平均分
|
||||
for m in retrieval_metrics_list:
|
||||
if m:
|
||||
total_score += sum(m.values()) / len(m)
|
||||
count += 1
|
||||
|
||||
# 计算答案平均分
|
||||
for m in answer_metrics_list:
|
||||
if "score" in m:
|
||||
total_score += m["score"]
|
||||
count += 1
|
||||
|
||||
return total_score / count if count > 0 else 0.0
|
||||
@ -5,14 +5,14 @@ import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from yuxi import config
|
||||
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
|
||||
from yuxi.utils.evaluation_metrics import EvaluationMetricsCalculator
|
||||
|
||||
|
||||
class EvaluationService:
|
||||
@ -288,9 +288,6 @@ class EvaluationService:
|
||||
return {"task_id": task_id, "message": "基准生成任务已提交"}
|
||||
|
||||
async def _generate_benchmark_task(self, context: TaskContext):
|
||||
import math
|
||||
import random
|
||||
|
||||
await context.set_progress(0, "初始化")
|
||||
|
||||
task = context._tasker._tasks.get(context.task_id)
|
||||
@ -304,11 +301,6 @@ class EvaluationService:
|
||||
embedding_model_id = payload.get("embedding_model_id")
|
||||
llm_model_spec = payload.get("llm_model_spec") or (payload.get("llm_config") or {}).get("model_spec")
|
||||
|
||||
if neighbors_count < 0:
|
||||
neighbors_count = 0
|
||||
if neighbors_count > 10:
|
||||
neighbors_count = 10
|
||||
|
||||
kb_instance = await knowledge_base.aget_kb(db_id)
|
||||
if not kb_instance:
|
||||
await context.set_message("知识库不存在")
|
||||
@ -317,35 +309,6 @@ class EvaluationService:
|
||||
await context.set_message("暂不支持该类型知识库生成评估基准")
|
||||
raise ValueError("Unsupported KB type for benchmark generation")
|
||||
|
||||
await context.set_progress(5, "加载chunks")
|
||||
|
||||
all_chunks = []
|
||||
for fid, finfo in kb_instance.files_meta.items():
|
||||
if finfo.get("database_id") != db_id:
|
||||
continue
|
||||
try:
|
||||
content_info = await kb_instance.get_file_content(db_id, fid)
|
||||
lines = content_info.get("lines", [])
|
||||
for line in lines:
|
||||
all_chunks.append(
|
||||
{
|
||||
"id": line.get("id"),
|
||||
"content": line.get("content", ""),
|
||||
"file_id": fid,
|
||||
"chunk_index": line.get("chunk_order_index"),
|
||||
}
|
||||
)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
if not all_chunks:
|
||||
await context.set_message("知识库为空或未解析到chunks")
|
||||
raise ValueError("No chunks found in knowledge base")
|
||||
|
||||
contents = [c["content"] for c in all_chunks]
|
||||
|
||||
await context.set_progress(15, "向量化")
|
||||
|
||||
db_meta = kb_instance.databases_meta.get(db_id, {})
|
||||
embed_info = db_meta.get("embed_info", {})
|
||||
if not embedding_model_id:
|
||||
@ -353,94 +316,28 @@ class EvaluationService:
|
||||
if not embedding_model_id:
|
||||
raise ValueError("Embedding model not specified")
|
||||
|
||||
from yuxi.models import select_embedding_model, select_model
|
||||
|
||||
embed_model = select_embedding_model(embedding_model_id)
|
||||
batch_size = int(getattr(embed_model, "batch_size", 40) or 40)
|
||||
if embedding_model_id in config.embed_model_names:
|
||||
batch_size = config.embed_model_names[embedding_model_id].batch_size
|
||||
# TODO: Performance Optimization
|
||||
# Currently, we re-calculate embeddings for ALL chunks in the KB for every benchmark generation.
|
||||
# This is inefficient for large KBs (O(N) embedding calls).
|
||||
# Optimization: Reuse existing embeddings from Vector DB if embedding_model_id matches the KB's embedding model.
|
||||
embeddings = await embed_model.abatch_encode(contents, batch_size=batch_size)
|
||||
norms = [math.sqrt(sum(x * x for x in vec)) or 1.0 for vec in embeddings]
|
||||
|
||||
def cosine(a, b, na, nb):
|
||||
s = 0.0
|
||||
for i in range(len(a)):
|
||||
s += a[i] * b[i]
|
||||
return s / (na * nb)
|
||||
|
||||
llm = select_model(model_spec=llm_model_spec)
|
||||
|
||||
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
|
||||
attempts = 0
|
||||
|
||||
await context.set_progress(0, "准备生成样本")
|
||||
|
||||
with open(data_file_path, "w", encoding="utf-8") as f:
|
||||
# Allow more attempts to generate enough questions
|
||||
max_attempts = max(count * 5, 50)
|
||||
while generated < count and attempts < max_attempts:
|
||||
attempts += 1
|
||||
i0 = random.randrange(len(all_chunks))
|
||||
e0 = embeddings[i0]
|
||||
n0 = norms[i0]
|
||||
|
||||
sims = []
|
||||
for j in range(len(all_chunks)):
|
||||
if j == i0:
|
||||
continue
|
||||
s = cosine(e0, embeddings[j], n0, norms[j])
|
||||
sims.append((j, s))
|
||||
sims.sort(key=lambda x: x[1], reverse=True)
|
||||
top_js = [j for j, _ in sims[:neighbors_count]]
|
||||
|
||||
ctx_items = []
|
||||
ctx_items.append((all_chunks[i0]["id"], all_chunks[i0]["content"]))
|
||||
for j in top_js:
|
||||
ctx_items.append((all_chunks[j]["id"], all_chunks[j]["content"]))
|
||||
allowed_ids = {cid for cid, _ in ctx_items}
|
||||
context_text = "\n\n".join([f"片段ID={cid}\n{content}" for cid, content in ctx_items])
|
||||
|
||||
prompt = (
|
||||
"你将基于以下上下文生成一个可由上下文准确回答的问题与标准答案。"
|
||||
"仅返回一个JSON对象,不要包含其他文字。"
|
||||
"键为 query、gold_answer、gold_chunk_ids。gold_chunk_ids 必须是上述上下文片段的ID子集。\n\n"
|
||||
"上下文:\n" + context_text + "\n"
|
||||
)
|
||||
|
||||
try:
|
||||
resp = await llm.call(prompt, False)
|
||||
content = resp.content if resp else ""
|
||||
|
||||
import json_repair
|
||||
|
||||
obj = json_repair.loads(content)
|
||||
q = obj.get("query")
|
||||
a = obj.get("gold_answer")
|
||||
gids = obj.get("gold_chunk_ids")
|
||||
if not q or not a or not isinstance(gids, list):
|
||||
logger.warning(f"Generated JSON missing fields or invalid format: {obj}")
|
||||
continue
|
||||
|
||||
gids = [str(x) for x in gids if str(x) in allowed_ids]
|
||||
if not gids:
|
||||
logger.warning("Generated gold_chunk_ids not found in allowed context")
|
||||
continue
|
||||
|
||||
line = {"query": q, "gold_chunk_ids": gids, "gold_answer": a}
|
||||
f.write(json.dumps(line, ensure_ascii=False) + "\n")
|
||||
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,
|
||||
embedding_model_id=embedding_model_id,
|
||||
llm_model_spec=llm_model_spec,
|
||||
progress_cb=context.set_progress,
|
||||
):
|
||||
f.write(dump_benchmark_item(item))
|
||||
generated += 1
|
||||
await context.set_progress(0 + int(99 * generated / max(count, 1)), f"已生成 {generated}/{count}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Benchmark generation failed for one item: {e}")
|
||||
continue
|
||||
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(
|
||||
{
|
||||
@ -590,109 +487,33 @@ class EvaluationService:
|
||||
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}")
|
||||
|
||||
# 执行查询
|
||||
query_result = await kb_instance.aquery(question_data["query"], db_id, **retrieval_config)
|
||||
|
||||
# 处理结果
|
||||
if isinstance(query_result, dict):
|
||||
generated_answer = query_result.get("answer", "")
|
||||
retrieved_chunks = query_result.get("retrieved_chunks", [])
|
||||
else:
|
||||
retrieved_chunks = query_result if isinstance(query_result, list) else []
|
||||
generated_answer = ""
|
||||
|
||||
# 如果没有生成的答案,但有检索结果且配置了 LLM,则生成答案
|
||||
if not generated_answer and retrieved_chunks and retrieval_config.get("answer_llm"):
|
||||
logger.debug(f"使用 LLM {retrieval_config.get('answer_llm')} 生成答案...")
|
||||
try:
|
||||
# 从配置中获取 LLM
|
||||
model_spec = retrieval_config["answer_llm"]
|
||||
llm = select_model(model_spec=model_spec)
|
||||
|
||||
# 构建上下文
|
||||
context_docs = []
|
||||
for idx, chunk in enumerate(retrieved_chunks[:5]): # 使用前5个最相关的文档
|
||||
content = chunk.get("content", "")
|
||||
if content:
|
||||
context_docs.append(f"文档 {idx + 1}:\n{content}")
|
||||
|
||||
context_text = "\\n\\n".join(context_docs)
|
||||
|
||||
# 构建提示词
|
||||
prompt = (
|
||||
f"基于以下上下文信息,请回答用户的问题。\n\n"
|
||||
f"上下文信息:{context_text}\n\n"
|
||||
f"用户问题:{question_data['query']}\n\n"
|
||||
"请根据上下文信息准确回答问题。\n\n"
|
||||
"如果上下文中缺少相关信息,请回答“信息不足,无法回答”。\n\n"
|
||||
)
|
||||
|
||||
# 生成答案
|
||||
response = await llm.call(prompt, stream=False)
|
||||
generated_answer = response.content if response else ""
|
||||
logger.debug(f"LLM 生成的答案长度: {len(generated_answer) if generated_answer else 0}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 生成答案失败: {e}")
|
||||
generated_answer = ""
|
||||
|
||||
# 计算指标
|
||||
current_metrics = {}
|
||||
retrieval_scores = {}
|
||||
answer_scores = {}
|
||||
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"):
|
||||
retrieval_scores = EvaluationMetricsCalculator.calculate_retrieval_metrics(
|
||||
retrieved_chunks, question_data["gold_chunk_ids"]
|
||||
)
|
||||
current_metrics.update(retrieval_scores)
|
||||
all_retrieval_metrics.append(retrieval_scores)
|
||||
|
||||
if benchmark_row.has_gold_answers and question_data.get("gold_answer"):
|
||||
if judge_llm:
|
||||
# 评判过程包含 LLM 调用
|
||||
answer_scores = await EvaluationMetricsCalculator.calculate_answer_metrics(
|
||||
query=question_data["query"],
|
||||
generated_answer=generated_answer,
|
||||
gold_answer=question_data["gold_answer"],
|
||||
judge_llm=judge_llm,
|
||||
)
|
||||
current_metrics.update(answer_scores)
|
||||
all_answer_metrics.append(answer_scores)
|
||||
else:
|
||||
logger.warning("需要计算答案指标但未配置 Judge LLM")
|
||||
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={
|
||||
"query_text": question_data["query"],
|
||||
"gold_chunk_ids": question_data.get("gold_chunk_ids"),
|
||||
"gold_answer": question_data.get("gold_answer"),
|
||||
"generated_answer": generated_answer,
|
||||
"retrieved_chunks": retrieved_chunks,
|
||||
"metrics": current_metrics,
|
||||
},
|
||||
data=question_result["detail"],
|
||||
)
|
||||
|
||||
# 计算当前累计指标
|
||||
current_overall_metrics = {}
|
||||
if all_retrieval_metrics:
|
||||
keys = all_retrieval_metrics[0].keys()
|
||||
for k in keys:
|
||||
current_overall_metrics[k] = sum(m.get(k, 0) for m in all_retrieval_metrics) / len(
|
||||
all_retrieval_metrics
|
||||
)
|
||||
if all_answer_metrics:
|
||||
scores = [m.get("score", 0) for m in all_answer_metrics]
|
||||
current_overall_metrics["answer_correctness"] = sum(scores) / len(scores) if scores else 0.0
|
||||
|
||||
# 更新 Tasker 的 result 以便实时获取当前指标
|
||||
current_overall_metrics, _ = aggregate_metrics(all_retrieval_metrics, all_answer_metrics)
|
||||
await context.set_result(
|
||||
{
|
||||
"current_metrics": current_overall_metrics,
|
||||
@ -701,31 +522,13 @@ class EvaluationService:
|
||||
}
|
||||
)
|
||||
|
||||
# 定期更新文件 (每5个或最后一个)
|
||||
if (i + 1) % 5 == 0 or (i + 1) == total_questions:
|
||||
await update_result_db(completed=i + 1)
|
||||
|
||||
# 最终计算
|
||||
await context.set_progress(95, "计算最终指标")
|
||||
|
||||
# 汇总指标
|
||||
overall_metrics = {}
|
||||
|
||||
# 检索指标平均值
|
||||
if all_retrieval_metrics:
|
||||
keys = all_retrieval_metrics[0].keys()
|
||||
for k in keys:
|
||||
overall_metrics[k] = sum(m.get(k, 0) for m in all_retrieval_metrics) / len(all_retrieval_metrics)
|
||||
|
||||
# 答案指标平均值
|
||||
if all_answer_metrics:
|
||||
scores = [m.get("score", 0) for m in all_answer_metrics]
|
||||
overall_metrics["answer_correctness"] = sum(scores) / len(scores) if scores else 0.0
|
||||
|
||||
overall_score = EvaluationMetricsCalculator.calculate_overall_score(
|
||||
all_retrieval_metrics, all_answer_metrics
|
||||
overall_metrics, overall_score = aggregate_metrics(
|
||||
all_retrieval_metrics, all_answer_metrics, include_overall_score=True
|
||||
)
|
||||
overall_metrics["overall_score"] = overall_score
|
||||
|
||||
await update_result_db(
|
||||
status="completed",
|
||||
|
||||
@ -1,152 +1,5 @@
|
||||
"""
|
||||
RAG评估指标计算工具
|
||||
简化版:只保留Recall/F1(检索)和 LLM Judge(答案准确性)
|
||||
"""
|
||||
"""RAG评估指标兼容入口。"""
|
||||
|
||||
import json
|
||||
import textwrap
|
||||
from typing import Any
|
||||
from yuxi.knowledge.eval.metrics import AnswerMetrics, EvaluationMetricsCalculator, RetrievalMetrics
|
||||
|
||||
from yuxi.utils import logger
|
||||
|
||||
|
||||
class RetrievalMetrics:
|
||||
"""检索评估指标计算"""
|
||||
|
||||
@staticmethod
|
||||
def precision_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
|
||||
"""计算Precision@K"""
|
||||
if not retrieved_ids[:k]:
|
||||
return 0.0
|
||||
retrieved_set = set(retrieved_ids[:k])
|
||||
relevant_set = set(relevant_ids)
|
||||
return len(retrieved_set & relevant_set) / k
|
||||
|
||||
@staticmethod
|
||||
def recall_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
|
||||
"""计算Recall@K"""
|
||||
if not relevant_ids:
|
||||
return 0.0
|
||||
retrieved_set = set(retrieved_ids[:k])
|
||||
relevant_set = set(relevant_ids)
|
||||
return len(retrieved_set & relevant_set) / len(relevant_set)
|
||||
|
||||
@staticmethod
|
||||
def f1_score_at_k(retrieved_ids: list[str], relevant_ids: list[str], k: int) -> float:
|
||||
"""计算F1@K"""
|
||||
precision = RetrievalMetrics.precision_at_k(retrieved_ids, relevant_ids, k)
|
||||
recall = RetrievalMetrics.recall_at_k(retrieved_ids, relevant_ids, k)
|
||||
if precision + recall == 0:
|
||||
return 0.0
|
||||
return 2 * precision * recall / (precision + recall)
|
||||
|
||||
|
||||
class AnswerMetrics:
|
||||
"""答案评估指标计算"""
|
||||
|
||||
@staticmethod
|
||||
async def judge_correctness(query: str, generated_answer: str, gold_answer: str, judge_llm: Any) -> dict[str, Any]:
|
||||
"""
|
||||
使用LLM判断生成的答案是否正确
|
||||
"""
|
||||
if not generated_answer:
|
||||
return {"score": 0.0, "reasoning": "未生成答案"}
|
||||
if not gold_answer:
|
||||
return {"score": 0.0, "reasoning": "无参考答案"}
|
||||
|
||||
prompt = textwrap.dedent(f"""你是一个公正的评判者,请评估AI生成的答案相对于标准答案的准确性。
|
||||
|
||||
问题:{query}
|
||||
|
||||
标准答案:
|
||||
{gold_answer}
|
||||
|
||||
AI生成的答案:
|
||||
{generated_answer}
|
||||
|
||||
请判断AI生成的答案是否在事实层面与标准答案一致。
|
||||
忽略措辞、标点符号或格式上的细微差异。
|
||||
只关注核心事实是否准确包含。
|
||||
|
||||
请返回以下JSON格式的结果(不要包含其他文本):
|
||||
{{
|
||||
"score": 1.0, // 如果答案正确返回 1.0, 错误返回 0.0
|
||||
"reasoning": "简要说明判定理由"
|
||||
}}
|
||||
""")
|
||||
try:
|
||||
response = await judge_llm.call(prompt, stream=False)
|
||||
content = response.content.strip()
|
||||
|
||||
# 尝试清理可能的 markdown 代码块
|
||||
if content.startswith("```json"):
|
||||
content = content[7:]
|
||||
if content.endswith("```"):
|
||||
content = content[:-3]
|
||||
content = content.strip()
|
||||
|
||||
result = json.loads(content)
|
||||
return {"score": float(result.get("score", 0.0)), "reasoning": result.get("reasoning", "")}
|
||||
except Exception as e:
|
||||
logger.error(f"LLM 评判失败: {e}")
|
||||
return {"score": 0.0, "reasoning": f"评判出错: {str(e)}"}
|
||||
|
||||
|
||||
class EvaluationMetricsCalculator:
|
||||
"""综合评估指标计算器"""
|
||||
|
||||
@staticmethod
|
||||
def calculate_retrieval_metrics(
|
||||
retrieved_chunks: list[dict[str, Any]], gold_chunk_ids: list[str], k_values: list[int] = [1, 3, 5, 10]
|
||||
) -> dict[str, float]:
|
||||
"""计算检索指标 (Recall, F1)"""
|
||||
if not retrieved_chunks or not gold_chunk_ids:
|
||||
return {}
|
||||
|
||||
# 提取 ID
|
||||
retrieved_ids = []
|
||||
for chunk in retrieved_chunks:
|
||||
chunk_id = chunk.get("chunk_id") or chunk.get("metadata", {}).get("chunk_id")
|
||||
retrieved_ids.append(str(chunk_id) if chunk_id else "")
|
||||
|
||||
metrics = {}
|
||||
for k in k_values:
|
||||
metrics[f"recall@{k}"] = RetrievalMetrics.recall_at_k(retrieved_ids, gold_chunk_ids, k)
|
||||
metrics[f"f1@{k}"] = RetrievalMetrics.f1_score_at_k(retrieved_ids, gold_chunk_ids, k)
|
||||
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
async def calculate_answer_metrics(
|
||||
query: str, generated_answer: str, gold_answer: str, judge_llm: Any = None
|
||||
) -> dict[str, Any]:
|
||||
"""计算答案指标 (LLM Judge)"""
|
||||
if not judge_llm:
|
||||
return {}
|
||||
|
||||
return await AnswerMetrics.judge_correctness(query, generated_answer, gold_answer, judge_llm)
|
||||
|
||||
@staticmethod
|
||||
def calculate_overall_score(
|
||||
retrieval_metrics_list: list[dict[str, float]], answer_metrics_list: list[dict[str, Any]]
|
||||
) -> float:
|
||||
"""计算整体平均分"""
|
||||
total_score = 0.0
|
||||
count = 0
|
||||
|
||||
# 简单的平均策略:将所有retrieval metric的值和answer metric的score一起平均
|
||||
# 用户可能希望分开看,但calculate_overall_score返回一个单值。
|
||||
|
||||
# 计算检索平均分
|
||||
for m in retrieval_metrics_list:
|
||||
if m:
|
||||
total_score += sum(m.values()) / len(m)
|
||||
count += 1
|
||||
|
||||
# 计算答案平均分
|
||||
for m in answer_metrics_list:
|
||||
if "score" in m:
|
||||
total_score += m["score"]
|
||||
count += 1
|
||||
|
||||
return total_score / count if count > 0 else 0.0
|
||||
__all__ = ["AnswerMetrics", "EvaluationMetricsCalculator", "RetrievalMetrics"]
|
||||
|
||||
@ -0,0 +1,54 @@
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
os.environ.setdefault("OPENAI_API_KEY", "test-key")
|
||||
|
||||
from yuxi.knowledge.eval.benchmark_generation import (
|
||||
build_benchmark_generation_prompt,
|
||||
clamp_neighbors_count,
|
||||
collect_kb_chunks,
|
||||
cosine_similarity,
|
||||
select_neighbor_indices,
|
||||
)
|
||||
|
||||
|
||||
class FakeKnowledgeBase:
|
||||
files_meta = {
|
||||
"file_a": {"database_id": "db_1"},
|
||||
"file_b": {"database_id": "db_2"},
|
||||
}
|
||||
|
||||
async def get_file_content(self, db_id, fid):
|
||||
return {
|
||||
"lines": [
|
||||
{"id": f"{fid}_chunk", "content": "内容", "chunk_order_index": 0},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def test_clamp_neighbors_count():
|
||||
assert clamp_neighbors_count(-1) == 0
|
||||
assert clamp_neighbors_count(3) == 3
|
||||
assert clamp_neighbors_count(11) == 10
|
||||
|
||||
|
||||
def test_select_neighbor_indices_orders_by_cosine_similarity():
|
||||
embeddings = [[1.0, 0.0], [0.9, 0.1], [0.0, 1.0]]
|
||||
norms = [1.0, cosine_similarity(embeddings[1], embeddings[1], 1.0, 1.0) ** 0.5, 1.0]
|
||||
|
||||
assert select_neighbor_indices(0, embeddings, norms, 1) == [1]
|
||||
|
||||
|
||||
def test_build_benchmark_generation_prompt_contains_required_schema():
|
||||
prompt = build_benchmark_generation_prompt([("chunk_1", "片段内容")])
|
||||
|
||||
assert "片段ID=chunk_1" in prompt
|
||||
assert "query、gold_answer、gold_chunk_ids" in prompt
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_collect_kb_chunks_filters_database_id():
|
||||
chunks = await collect_kb_chunks(FakeKnowledgeBase(), "db_1")
|
||||
|
||||
assert chunks == [{"id": "file_a_chunk", "content": "内容", "file_id": "file_a", "chunk_index": 0}]
|
||||
39
backend/test/unit/knowledge/eval/test_evaluator.py
Normal file
39
backend/test/unit/knowledge/eval/test_evaluator.py
Normal file
@ -0,0 +1,39 @@
|
||||
import os
|
||||
|
||||
os.environ.setdefault("OPENAI_API_KEY", "test-key")
|
||||
|
||||
from yuxi.knowledge.eval.evaluator import aggregate_metrics, build_answer_prompt, normalize_query_result
|
||||
|
||||
|
||||
def test_normalize_query_result_supports_dict_and_list():
|
||||
answer, chunks = normalize_query_result({"answer": "A", "retrieved_chunks": [{"content": "C"}]})
|
||||
assert answer == "A"
|
||||
assert chunks == [{"content": "C"}]
|
||||
|
||||
answer, chunks = normalize_query_result([{"content": "C"}])
|
||||
assert answer == ""
|
||||
assert chunks == [{"content": "C"}]
|
||||
|
||||
|
||||
def test_build_answer_prompt_uses_first_five_non_empty_chunks():
|
||||
chunks = [{"content": f"内容{i}"} for i in range(6)] + [{"content": ""}]
|
||||
|
||||
prompt = build_answer_prompt("问题", chunks)
|
||||
|
||||
assert "用户问题:问题" in prompt
|
||||
assert "内容0" in prompt
|
||||
assert "内容4" in prompt
|
||||
assert "内容5" not in prompt
|
||||
|
||||
|
||||
def test_aggregate_metrics_matches_service_output_shape():
|
||||
metrics, overall_score = aggregate_metrics(
|
||||
[{"recall@1": 1.0, "f1@1": 0.0}, {"recall@1": 0.0, "f1@1": 1.0}],
|
||||
[{"score": 1.0}, {"score": 0.0}],
|
||||
include_overall_score=True,
|
||||
)
|
||||
|
||||
assert metrics["recall@1"] == 0.5
|
||||
assert metrics["f1@1"] == 0.5
|
||||
assert metrics["answer_correctness"] == 0.5
|
||||
assert metrics["overall_score"] == overall_score
|
||||
28
backend/test/unit/knowledge/eval/test_metrics.py
Normal file
28
backend/test/unit/knowledge/eval/test_metrics.py
Normal file
@ -0,0 +1,28 @@
|
||||
import os
|
||||
|
||||
os.environ.setdefault("OPENAI_API_KEY", "test-key")
|
||||
|
||||
from yuxi.knowledge.eval.metrics import EvaluationMetricsCalculator, RetrievalMetrics
|
||||
|
||||
|
||||
def test_retrieval_metrics_use_metadata_chunk_id():
|
||||
retrieved_chunks = [
|
||||
{"metadata": {"chunk_id": "chunk_a"}},
|
||||
{"metadata": {"chunk_id": "chunk_b"}},
|
||||
]
|
||||
|
||||
metrics = EvaluationMetricsCalculator.calculate_retrieval_metrics(
|
||||
retrieved_chunks, ["chunk_b", "chunk_c"], k_values=[1, 3]
|
||||
)
|
||||
|
||||
assert metrics["recall@1"] == 0.0
|
||||
assert metrics["recall@3"] == 0.5
|
||||
assert metrics["f1@3"] == RetrievalMetrics.f1_score_at_k(["chunk_a", "chunk_b"], ["chunk_b", "chunk_c"], 3)
|
||||
|
||||
|
||||
def test_overall_score_keeps_existing_average_strategy():
|
||||
score = EvaluationMetricsCalculator.calculate_overall_score(
|
||||
[{"recall@1": 1.0, "f1@1": 0.5}], [{"score": 0.25}]
|
||||
)
|
||||
|
||||
assert score == 0.5
|
||||
@ -40,6 +40,7 @@
|
||||
- 修复知识库文档入库状态回退:当已解析文件缺失 `markdown_file` 解析产物时,索引流程会将文件状态恢复为未解析,便于重新解析而不是停留在索引失败。
|
||||
- 优化 Agent 输入框文件 mention:用户级 workspace 文件候选改为从独立 workspace API 递归加载,不再依赖 active thread;插入时仍转换为 `/home/gem/user-data/workspace/` 沙盒虚拟路径,并修复附件上传后未立即刷新 mention 候选的问题。
|
||||
- 调整知识库思维导图后端结构:将思维导图路由文件重命名为知识库语义更明确的 router,并把文件列表整理、提示词构建、AI JSON 解析等纯逻辑下沉到知识库 utils。
|
||||
- 收敛知识库评估后端结构:将评估指标、单题评估、答案生成提示词和自动基准生成算法下沉到 `knowledge/eval`,`EvaluationService` 保留任务、文件和持久化编排职责。
|
||||
- 新增个人工作区预览与管理:提供独立于对话 thread 的用户级 workspace API,并增加“工作区”页面,用于浏览个人 workspace 文件、预览 Markdown/文本/代码/图片/PDF;支持新建文件夹、上传文件、下载文件、删除文件/文件夹和多选删除;工作区预览支持 Markdown/TXT 在右侧预览框内切换编辑并保存,其他格式和非工作区预览默认只读;知识库与团队空间入口先展示到占位层级;默认创建 `agents/AGENTS.md`,并在 Agent 执行时将其内容追加到系统提示词。
|
||||
- 加固 JWT 鉴权安全:移除历史默认密钥回退,初始化脚本支持生成并持久化 `JWT_SECRET_KEY` 与 `YUXI_INSTANCE_ID`,签发和验证令牌时校验 `iss/aud`,并在鉴权阶段拒绝已删除或登录锁定用户继续使用旧令牌访问系统。
|
||||
- 扩展管理界面交互逻辑重构:将 MCP / Subagents / Skills 三个标签页从「左侧边栏 + 右侧详情面板」布局重构为「卡片式网格布局 + 路由跳转二级页面」布局,工具标签页改为卡片网格布局 + 弹窗详情(保持弹窗内容不变)。新增共享组件 `ExtensionCard`、`ExtensionCardGrid`、`ExtensionToolbar`、`ExtensionDetailLayout`,详情页(`McpDetailView`、`SubagentDetailView`、`SkillDetailView`)使用居中宽度限制,路由规划为 `/extensions/mcp/:name`、`/extensions/subagent/:name`、`/extensions/skill/:slug`。
|
||||
|
||||
@ -6,16 +6,16 @@
|
||||
<span class="total-count">{{ benchmarks.length }} 个基准</span>
|
||||
</div>
|
||||
<div class="header-right">
|
||||
<a-button @click="loadBenchmarks">
|
||||
<template #icon><ReloadOutlined /></template>
|
||||
<a-button class="lucide-icon-btn" @click="loadBenchmarks">
|
||||
<template #icon><RefreshCw :size="16" /></template>
|
||||
刷新
|
||||
</a-button>
|
||||
<a-button type="primary" @click="showUploadModal">
|
||||
<template #icon><UploadOutlined /></template>
|
||||
<a-button type="primary" class="lucide-icon-btn" @click="showUploadModal">
|
||||
<template #icon><Upload :size="16" /></template>
|
||||
上传基准
|
||||
</a-button>
|
||||
<a-button @click="showGenerateModal">
|
||||
<template #icon><RobotOutlined /></template>
|
||||
<a-button class="lucide-icon-btn" @click="showGenerateModal">
|
||||
<template #icon><Bot :size="16" /></template>
|
||||
自动生成
|
||||
</a-button>
|
||||
</div>
|
||||
@ -45,19 +45,31 @@
|
||||
<div class="benchmark-header">
|
||||
<h4 class="benchmark-name">{{ benchmark.name }}</h4>
|
||||
<div class="benchmark-actions">
|
||||
<a-button type="text" size="small" @click.stop="previewBenchmark(benchmark)">
|
||||
<EyeOutlined />
|
||||
<a-button
|
||||
type="text"
|
||||
size="small"
|
||||
class="lucide-icon-btn"
|
||||
@click.stop="previewBenchmark(benchmark)"
|
||||
>
|
||||
<Eye :size="15" />
|
||||
</a-button>
|
||||
<a-button
|
||||
type="text"
|
||||
size="small"
|
||||
class="lucide-icon-btn"
|
||||
:loading="!!downloadingBenchmarkMap[benchmark.benchmark_id]"
|
||||
@click.stop="downloadBenchmark(benchmark)"
|
||||
>
|
||||
<DownloadOutlined />
|
||||
<Download :size="15" />
|
||||
</a-button>
|
||||
<a-button type="text" size="small" danger @click.stop="deleteBenchmark(benchmark)">
|
||||
<DeleteOutlined />
|
||||
<a-button
|
||||
type="text"
|
||||
size="small"
|
||||
class="lucide-icon-btn"
|
||||
danger
|
||||
@click.stop="deleteBenchmark(benchmark)"
|
||||
>
|
||||
<Trash2 :size="15" />
|
||||
</a-button>
|
||||
</div>
|
||||
</div>
|
||||
@ -69,23 +81,23 @@
|
||||
<div class="meta-row">
|
||||
<span
|
||||
v-if="benchmark.has_gold_chunks && benchmark.has_gold_answers"
|
||||
class="type-badge type-both"
|
||||
class="card-tag benchmark-tag tag-purple"
|
||||
>
|
||||
检索 + 问答
|
||||
</span>
|
||||
<span v-else-if="benchmark.has_gold_chunks" class="type-badge type-retrieval">
|
||||
<span v-else-if="benchmark.has_gold_chunks" class="card-tag benchmark-tag tag-blue">
|
||||
检索评估
|
||||
</span>
|
||||
<span v-else-if="benchmark.has_gold_answers" class="type-badge type-answer">
|
||||
<span v-else-if="benchmark.has_gold_answers" class="card-tag benchmark-tag tag-gold">
|
||||
问答评估
|
||||
</span>
|
||||
<span v-else class="type-badge type-query">仅查询</span>
|
||||
<span v-else class="card-tag benchmark-tag">仅查询</span>
|
||||
|
||||
<span :class="['tag', benchmark.has_gold_chunks ? 'tag-yes' : 'tag-no']">
|
||||
{{ benchmark.has_gold_chunks ? '✓' : '✗' }} 黄金Chunk
|
||||
<span v-if="benchmark.has_gold_chunks" class="card-tag benchmark-tag tag-green">
|
||||
Gold Chunks
|
||||
</span>
|
||||
<span :class="['tag', benchmark.has_gold_answers ? 'tag-yes' : 'tag-no']">
|
||||
{{ benchmark.has_gold_answers ? '✓' : '✗' }} 黄金答案
|
||||
<span v-if="benchmark.has_gold_answers" class="card-tag benchmark-tag tag-green">
|
||||
Gold Answer
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
@ -125,13 +137,13 @@
|
||||
{{ previewData.question_count }}
|
||||
</span>
|
||||
<span class="meta-item">
|
||||
<span class="meta-label">黄金Chunk:</span>
|
||||
<span class="meta-label">Gold Chunks:</span>
|
||||
<span :class="previewData.has_gold_chunks ? 'status-yes' : 'status-no'">
|
||||
{{ previewData.has_gold_chunks ? '有' : '无' }}
|
||||
</span>
|
||||
</span>
|
||||
<span class="meta-item">
|
||||
<span class="meta-label">黄金答案:</span>
|
||||
<span class="meta-label">Gold Answer:</span>
|
||||
<span :class="previewData.has_gold_answers ? 'status-yes' : 'status-no'">
|
||||
{{ previewData.has_gold_answers ? '有' : '无' }}
|
||||
</span>
|
||||
@ -200,14 +212,7 @@
|
||||
<script setup>
|
||||
import { ref, reactive, onMounted, computed } from 'vue'
|
||||
import { message, Modal } from 'ant-design-vue'
|
||||
import {
|
||||
UploadOutlined,
|
||||
RobotOutlined,
|
||||
EyeOutlined,
|
||||
DownloadOutlined,
|
||||
DeleteOutlined,
|
||||
ReloadOutlined
|
||||
} from '@ant-design/icons-vue'
|
||||
import { Bot, Download, Eye, RefreshCw, Trash2, Upload } from 'lucide-vue-next'
|
||||
import { evaluationApi } from '@/apis/knowledge_api'
|
||||
import { useTaskerStore } from '@/stores/tasker'
|
||||
import BenchmarkUploadModal from './modals/BenchmarkUploadModal.vue'
|
||||
@ -256,14 +261,14 @@ const questionColumns = [
|
||||
ellipsis: false
|
||||
},
|
||||
{
|
||||
title: '黄金Chunk',
|
||||
title: 'Gold Chunks',
|
||||
dataIndex: 'gold_chunk_ids',
|
||||
key: 'gold_chunk_ids',
|
||||
width: 200,
|
||||
ellipsis: false
|
||||
},
|
||||
{
|
||||
title: '黄金答案',
|
||||
title: 'Gold Answer',
|
||||
dataIndex: 'gold_answer',
|
||||
key: 'gold_answer',
|
||||
width: 420,
|
||||
@ -611,47 +616,10 @@ onMounted(() => {
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
|
||||
.tag {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
padding: 1px 6px;
|
||||
border-radius: 3px;
|
||||
font-size: 11px;
|
||||
font-weight: 500;
|
||||
background: var(--main-50);
|
||||
color: var(--color-text-tertiary);
|
||||
|
||||
&.tag-yes {
|
||||
// background: var(--color-success-50);
|
||||
color: var(--main-500);
|
||||
}
|
||||
}
|
||||
|
||||
.type-badge {
|
||||
padding: 1px 6px;
|
||||
border-radius: 3px;
|
||||
font-size: 11px;
|
||||
font-weight: 500;
|
||||
|
||||
&.type-both {
|
||||
background: var(--color-accent-50);
|
||||
color: var(--color-accent-700);
|
||||
}
|
||||
|
||||
&.type-retrieval {
|
||||
background: var(--color-info-50);
|
||||
color: var(--color-info-700);
|
||||
}
|
||||
|
||||
&.type-answer {
|
||||
background: var(--color-warning-50);
|
||||
color: var(--color-warning-700);
|
||||
}
|
||||
|
||||
&.type-query {
|
||||
background: var(--gray-100);
|
||||
color: var(--gray-700);
|
||||
}
|
||||
.benchmark-tag {
|
||||
min-height: 22px;
|
||||
padding: 0 8px;
|
||||
border-radius: 4px;
|
||||
}
|
||||
|
||||
.benchmark-footer {
|
||||
|
||||
@ -26,15 +26,13 @@
|
||||
size="middle"
|
||||
:loading="benchmarksLoading"
|
||||
@click="() => loadBenchmarks(true)"
|
||||
:icon="h(ReloadOutlined)"
|
||||
class="refresh-benchmarks-btn"
|
||||
:icon="h(RefreshCw, { size: 16 })"
|
||||
class="refresh-benchmarks-btn lucide-icon-btn"
|
||||
title="刷新评估基准列表"
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<div class="toolbar-right">
|
||||
<!-- 检索配置按钮 -->
|
||||
<a-button size="middle" @click="openSearchConfigModal" :icon="h(SettingOutlined)" />
|
||||
<!-- 开始评估按钮 -->
|
||||
<a-button
|
||||
type="primary"
|
||||
@ -116,8 +114,8 @@
|
||||
size="small"
|
||||
:loading="refreshingHistory"
|
||||
@click="refreshHistory"
|
||||
:icon="h('ReloadOutlined')"
|
||||
class="refresh-btn"
|
||||
:icon="h(RefreshCw, { size: 14 })"
|
||||
class="refresh-btn lucide-icon-btn"
|
||||
>
|
||||
刷新
|
||||
</a-button>
|
||||
@ -130,12 +128,7 @@
|
||||
size="small"
|
||||
>
|
||||
<template #bodyCell="{ column, record }">
|
||||
<template v-if="column.key === 'status'">
|
||||
<a-tag :color="getStatusColor(record.status)">
|
||||
{{ getStatusText(record.status) }}
|
||||
</a-tag>
|
||||
</template>
|
||||
<template v-else-if="column.key === 'overall_score'">
|
||||
<template v-if="column.key === 'overall_score'">
|
||||
<span v-if="record.overall_score !== null">
|
||||
<a-tag :color="getScoreTagColor(record.overall_score)">
|
||||
{{ (record.overall_score * 100).toFixed(0) }}%
|
||||
@ -151,8 +144,11 @@
|
||||
size="small"
|
||||
@click="viewResults(record.task_id)"
|
||||
>
|
||||
查看结果
|
||||
查看
|
||||
</a-button>
|
||||
<a-tag v-else :color="getStatusColor(record.status)">
|
||||
{{ getStatusText(record.status) }}
|
||||
</a-tag>
|
||||
<a-popconfirm
|
||||
title="确定要删除这条评估记录吗?"
|
||||
description="删除后将无法恢复"
|
||||
@ -394,12 +390,6 @@
|
||||
</div>
|
||||
</a-modal>
|
||||
|
||||
<!-- 检索配置弹窗 -->
|
||||
<SearchConfigModal
|
||||
v-model="searchConfigModalVisible"
|
||||
:database-id="databaseId"
|
||||
@save="handleSearchConfigSave"
|
||||
/>
|
||||
</template>
|
||||
|
||||
<script setup>
|
||||
@ -407,8 +397,7 @@ import { ref, reactive, onMounted, computed, h } from 'vue'
|
||||
import { message, Modal } from 'ant-design-vue'
|
||||
import { evaluationApi } from '@/apis/knowledge_api'
|
||||
import ModelSelectorComponent from '@/components/ModelSelectorComponent.vue'
|
||||
import SearchConfigModal from './SearchConfigModal.vue'
|
||||
import { SettingOutlined, ReloadOutlined } from '@ant-design/icons-vue'
|
||||
import { RefreshCw } from 'lucide-vue-next'
|
||||
import { useTaskerStore } from '@/stores/tasker'
|
||||
|
||||
const props = defineProps({
|
||||
@ -435,7 +424,6 @@ const selectedResult = ref(null)
|
||||
const detailedResults = ref([])
|
||||
const evaluationStats = ref({})
|
||||
const resultsLoading = ref(false)
|
||||
const searchConfigModalVisible = ref(false)
|
||||
const refreshingHistory = ref(false)
|
||||
const showErrorsOnly = ref(false)
|
||||
const currentPage = ref(1)
|
||||
@ -511,12 +499,6 @@ const historyColumns = [
|
||||
return benchmark ? benchmark.name : record.benchmark_id?.slice(0, 8) || '-'
|
||||
}
|
||||
},
|
||||
{
|
||||
title: '状态',
|
||||
dataIndex: 'status',
|
||||
key: 'status',
|
||||
width: 100
|
||||
},
|
||||
{
|
||||
title: 'Recall@10',
|
||||
key: 'recall_10',
|
||||
@ -568,7 +550,7 @@ const historyColumns = [
|
||||
{
|
||||
title: '操作',
|
||||
key: 'actions',
|
||||
width: 150
|
||||
width: 100
|
||||
}
|
||||
]
|
||||
|
||||
@ -659,17 +641,6 @@ const loadResultsWithPagination = async () => {
|
||||
}
|
||||
}
|
||||
|
||||
// 打开检索配置弹窗
|
||||
const openSearchConfigModal = () => {
|
||||
searchConfigModalVisible.value = true
|
||||
}
|
||||
|
||||
// 处理检索配置保存
|
||||
const handleSearchConfigSave = (config) => {
|
||||
console.log('RAG评估中的检索配置已更新:', config)
|
||||
// 可以在这里添加配置更新后的处理逻辑
|
||||
}
|
||||
|
||||
// 加载基准列表
|
||||
const loadBenchmarks = async (showSuccessMessage = false) => {
|
||||
if (!props.databaseId) return
|
||||
|
||||
Loading…
Reference in New Issue
Block a user