aeval-framework 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- aeval_framework-0.1.0.dist-info/METADATA +42 -0
- aeval_framework-0.1.0.dist-info/RECORD +63 -0
- aeval_framework-0.1.0.dist-info/WHEEL +4 -0
- aeval_framework-0.1.0.dist-info/entry_points.txt +2 -0
- agent_eval/__init__.py +14 -0
- agent_eval/api/__init__.py +14 -0
- agent_eval/api/app.py +82 -0
- agent_eval/api/events.py +96 -0
- agent_eval/api/routes/__init__.py +0 -0
- agent_eval/api/routes/datasets.py +441 -0
- agent_eval/api/routes/graders.py +19 -0
- agent_eval/api/routes/metrics.py +49 -0
- agent_eval/api/routes/runs.py +573 -0
- agent_eval/api/routes/suites.py +84 -0
- agent_eval/api/routes/tasks.py +114 -0
- agent_eval/api/standalone.py +105 -0
- agent_eval/cli.py +455 -0
- agent_eval/core/__init__.py +48 -0
- agent_eval/core/contract.py +296 -0
- agent_eval/core/metrics.py +184 -0
- agent_eval/core/runner.py +868 -0
- agent_eval/core/suite.py +60 -0
- agent_eval/core/types.py +227 -0
- agent_eval/dataset/__init__.py +31 -0
- agent_eval/dataset/models.py +199 -0
- agent_eval/dataset/quality.py +194 -0
- agent_eval/dataset/sources/__init__.py +45 -0
- agent_eval/dataset/sources/llm_generator.py +219 -0
- agent_eval/dataset/sources/manual.py +172 -0
- agent_eval/dataset/sources/regression.py +201 -0
- agent_eval/dataset/sources/trace_mining.py +277 -0
- agent_eval/dataset/storage.py +342 -0
- agent_eval/dataset/version.py +72 -0
- agent_eval/examples/__init__.py +0 -0
- agent_eval/examples/basic_usage.py +175 -0
- agent_eval/examples/mock_runner.py +195 -0
- agent_eval/graders/__init__.py +91 -0
- agent_eval/graders/artifact_check.py +114 -0
- agent_eval/graders/code_based.py +101 -0
- agent_eval/graders/human.py +77 -0
- agent_eval/graders/metric.py +142 -0
- agent_eval/graders/model_based.py +179 -0
- agent_eval/graders/state_check.py +106 -0
- agent_eval/graders/step_level.py +116 -0
- agent_eval/graders/tool_calls.py +102 -0
- agent_eval/graders/transcript.py +86 -0
- agent_eval/metrics/__init__.py +110 -0
- agent_eval/metrics/answer_relevancy.py +57 -0
- agent_eval/metrics/base.py +155 -0
- agent_eval/metrics/batch_evaluation.py +267 -0
- agent_eval/metrics/context_precision.py +62 -0
- agent_eval/metrics/context_recall.py +71 -0
- agent_eval/metrics/faithfulness.py +72 -0
- agent_eval/metrics/llm_judge.py +100 -0
- agent_eval/metrics/prompt_metric.py +150 -0
- agent_eval/metrics/pytest_plugin.py +308 -0
- agent_eval/metrics/report.py +149 -0
- agent_eval/metrics/synthetic_data.py +203 -0
- agent_eval/storage/__init__.py +17 -0
- agent_eval/storage/memory.py +95 -0
- agent_eval/storage/sqlite.py +240 -0
- agent_eval/trace/__init__.py +16 -0
- agent_eval/trace/phoenix.py +144 -0
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
"""P0 metric — Context Precision (retrieval_context 中相关文档占比)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from agent_eval.metrics.base import BaseLLMMetric, MetricError, MetricResult
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class ContextPrecisionMetric(BaseLLMMetric):
|
|
9
|
+
"""
|
|
10
|
+
上下文精确率: 检索到的文档中有多少真正与问题相关。
|
|
11
|
+
|
|
12
|
+
缺 retrieval_context 明确失败 (score=0 + 理由)。
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
name = "context_precision"
|
|
16
|
+
|
|
17
|
+
_SYSTEM_PROMPT = """你是一个信息检索评测专家。请评估检索上下文的精确率。
|
|
18
|
+
|
|
19
|
+
步骤:
|
|
20
|
+
1. 对检索上下文中的每个文档, 判断它是否对回答用户问题有用
|
|
21
|
+
2. 计算精确率 = 有用文档数 / 总文档数
|
|
22
|
+
|
|
23
|
+
以 JSON 格式返回:
|
|
24
|
+
{
|
|
25
|
+
"documents": [{"index": 0, "relevant": true, "reason": "..."}],
|
|
26
|
+
"score": 0.67,
|
|
27
|
+
"reason": "一句话理由"
|
|
28
|
+
}"""
|
|
29
|
+
|
|
30
|
+
async def measure(
|
|
31
|
+
self,
|
|
32
|
+
input: str,
|
|
33
|
+
actual_output: str,
|
|
34
|
+
expected_output: str | None = None,
|
|
35
|
+
context: list[str] | None = None,
|
|
36
|
+
retrieval_context: list[str] | None = None,
|
|
37
|
+
) -> MetricResult:
|
|
38
|
+
docs = [d for d in (retrieval_context or []) if str(d).strip()]
|
|
39
|
+
if not docs:
|
|
40
|
+
return MetricResult(
|
|
41
|
+
name=self.name,
|
|
42
|
+
score=0.0,
|
|
43
|
+
reason="需要 retrieval_context 才能评估精确率",
|
|
44
|
+
details={"error": "missing_parameters"},
|
|
45
|
+
threshold=self.threshold,
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
doc_lines = "\n".join(f"[{i}] {doc}" for i, doc in enumerate(docs))
|
|
49
|
+
user_prompt = f"用户问题: {input}\n\n检索文档:\n{doc_lines}"
|
|
50
|
+
data = await self._llm_judge(self._SYSTEM_PROMPT, user_prompt)
|
|
51
|
+
|
|
52
|
+
if "score" not in data:
|
|
53
|
+
raise MetricError(f"judge response missing 'score': {data!r}")
|
|
54
|
+
score = self._score_of(data)
|
|
55
|
+
|
|
56
|
+
return MetricResult(
|
|
57
|
+
name=self.name,
|
|
58
|
+
score=score,
|
|
59
|
+
reason=str(data.get("reason", "")),
|
|
60
|
+
details={"documents": data.get("documents", [])},
|
|
61
|
+
threshold=self.threshold,
|
|
62
|
+
)
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""P0 metric — Context Recall (expected_output 信息点被 retrieval_context 覆盖率)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from agent_eval.metrics.base import BaseLLMMetric, MetricError, MetricResult
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class ContextRecallMetric(BaseLLMMetric):
|
|
9
|
+
"""
|
|
10
|
+
上下文召回率: 检索到的文档是否包含回答所需的全部信息。
|
|
11
|
+
|
|
12
|
+
需要 expected_output (信息点来源) 与 retrieval_context (检索结果);
|
|
13
|
+
缺任一参数明确失败 (score=0 + 理由)。
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
name = "context_recall"
|
|
17
|
+
|
|
18
|
+
_SYSTEM_PROMPT = """你是一个信息检索评测专家。请评估检索上下文的召回率。
|
|
19
|
+
|
|
20
|
+
步骤:
|
|
21
|
+
1. 从期望回答中提取所有关键信息点 (facts/claims)
|
|
22
|
+
2. 对每个信息点, 检查检索上下文中是否包含
|
|
23
|
+
3. 计算召回率 = 被覆盖的信息点数 / 总信息点数
|
|
24
|
+
|
|
25
|
+
以 JSON 格式返回:
|
|
26
|
+
{
|
|
27
|
+
"information_points": ["信息点1", "信息点2"],
|
|
28
|
+
"covered": [true, false],
|
|
29
|
+
"score": 0.75,
|
|
30
|
+
"reason": "一句话理由"
|
|
31
|
+
}"""
|
|
32
|
+
|
|
33
|
+
async def measure(
|
|
34
|
+
self,
|
|
35
|
+
input: str,
|
|
36
|
+
actual_output: str,
|
|
37
|
+
expected_output: str | None = None,
|
|
38
|
+
context: list[str] | None = None,
|
|
39
|
+
retrieval_context: list[str] | None = None,
|
|
40
|
+
) -> MetricResult:
|
|
41
|
+
docs = [d for d in (retrieval_context or []) if str(d).strip()]
|
|
42
|
+
if not expected_output or not expected_output.strip() or not docs:
|
|
43
|
+
return MetricResult(
|
|
44
|
+
name=self.name,
|
|
45
|
+
score=0.0,
|
|
46
|
+
reason="需要 expected_output 和 retrieval_context 才能评估召回率",
|
|
47
|
+
details={"error": "missing_parameters"},
|
|
48
|
+
threshold=self.threshold,
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
retrieval_str = "\n---\n".join(str(d) for d in docs)
|
|
52
|
+
user_prompt = (
|
|
53
|
+
f"期望回答 (包含所有应覆盖的信息):\n{expected_output}\n\n"
|
|
54
|
+
f"检索上下文:\n{retrieval_str}"
|
|
55
|
+
)
|
|
56
|
+
data = await self._llm_judge(self._SYSTEM_PROMPT, user_prompt)
|
|
57
|
+
|
|
58
|
+
if "score" not in data:
|
|
59
|
+
raise MetricError(f"judge response missing 'score': {data!r}")
|
|
60
|
+
score = self._score_of(data)
|
|
61
|
+
|
|
62
|
+
return MetricResult(
|
|
63
|
+
name=self.name,
|
|
64
|
+
score=score,
|
|
65
|
+
reason=str(data.get("reason", "")),
|
|
66
|
+
details={
|
|
67
|
+
"information_points": data.get("information_points", []),
|
|
68
|
+
"covered": data.get("covered", []),
|
|
69
|
+
},
|
|
70
|
+
threshold=self.threshold,
|
|
71
|
+
)
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
"""P0 metric — Faithfulness (回答忠于 context 的程度, 防幻觉)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from agent_eval.metrics.base import BaseLLMMetric, MetricError, MetricResult
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class FaithfulnessMetric(BaseLLMMetric):
|
|
9
|
+
"""
|
|
10
|
+
忠实度: Agent 回答是否完全基于给定上下文, 无幻觉。
|
|
11
|
+
|
|
12
|
+
无 context 时明确失败 (score=0 + 明确理由), 不猜测 (spec 场景:
|
|
13
|
+
"对 faithfulness 计算但不提供 context → score=0 与'无上下文无法评估'")。
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
name = "faithfulness"
|
|
17
|
+
|
|
18
|
+
_SYSTEM_PROMPT = """你是一个事实核查专家。请评估 Agent 回答是否忠于给定上下文。
|
|
19
|
+
|
|
20
|
+
步骤:
|
|
21
|
+
1. 从 Agent 回答中提取所有事实性陈述
|
|
22
|
+
2. 对每个陈述, 在上下文中查找支持证据
|
|
23
|
+
3. 判断: supported / unsupported
|
|
24
|
+
4. 计算忠实度 = 有支持的陈述数 / 总陈述数
|
|
25
|
+
|
|
26
|
+
以 JSON 格式返回:
|
|
27
|
+
{
|
|
28
|
+
"claims": ["陈述1", "陈述2"],
|
|
29
|
+
"verdicts": ["supported", "unsupported"],
|
|
30
|
+
"unsupported_claims": ["不被上下文支持的陈述"],
|
|
31
|
+
"score": 0.8,
|
|
32
|
+
"reason": "一句话理由 (标注不被支持的陈述)"
|
|
33
|
+
}"""
|
|
34
|
+
|
|
35
|
+
async def measure(
|
|
36
|
+
self,
|
|
37
|
+
input: str,
|
|
38
|
+
actual_output: str,
|
|
39
|
+
expected_output: str | None = None,
|
|
40
|
+
context: list[str] | None = None,
|
|
41
|
+
retrieval_context: list[str] | None = None,
|
|
42
|
+
) -> MetricResult:
|
|
43
|
+
context = [c for c in (context or []) if str(c).strip()]
|
|
44
|
+
if not context:
|
|
45
|
+
# 缺参明确失败路径: 返回 score=0 与明确理由 (不猜、不静默)
|
|
46
|
+
return MetricResult(
|
|
47
|
+
name=self.name,
|
|
48
|
+
score=0.0,
|
|
49
|
+
reason="无上下文,无法评估忠实度 — 需要提供 context 或 retrieval_context",
|
|
50
|
+
details={"error": "missing_context"},
|
|
51
|
+
threshold=self.threshold,
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
context_str = "\n---\n".join(str(c) for c in context)
|
|
55
|
+
user_prompt = f"上下文:\n{context_str}\n\nAgent 回答:\n{actual_output}"
|
|
56
|
+
data = await self._llm_judge(self._SYSTEM_PROMPT, user_prompt)
|
|
57
|
+
|
|
58
|
+
if "score" not in data:
|
|
59
|
+
raise MetricError(f"judge response missing 'score': {data!r}")
|
|
60
|
+
score = self._score_of(data)
|
|
61
|
+
|
|
62
|
+
return MetricResult(
|
|
63
|
+
name=self.name,
|
|
64
|
+
score=score,
|
|
65
|
+
reason=str(data.get("reason", "")),
|
|
66
|
+
details={
|
|
67
|
+
"claims": data.get("claims", []),
|
|
68
|
+
"verdicts": data.get("verdicts", []),
|
|
69
|
+
"unsupported_claims": data.get("unsupported_claims", []),
|
|
70
|
+
},
|
|
71
|
+
threshold=self.threshold,
|
|
72
|
+
)
|
|
@@ -0,0 +1,100 @@
|
|
|
1
|
+
"""LLM Judge infrastructure — protocol-injected LLM function + tolerant JSON parsing.
|
|
2
|
+
|
|
3
|
+
D2: the framework core only knows the LLMFn protocol
|
|
4
|
+
(async (system_prompt, user_message) -> raw text); no LLM SDK is bound here.
|
|
5
|
+
Assembly of a concrete implementation lives in eval_integration.config.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import asyncio
|
|
11
|
+
import json
|
|
12
|
+
import re
|
|
13
|
+
from collections.abc import Awaitable, Callable
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
# LLM 函数协议: (system_prompt, user_message) → raw text
|
|
17
|
+
LLMFn = Callable[[str, str], Awaitable[str]]
|
|
18
|
+
|
|
19
|
+
# 解析失败后的最大重试次数 (每次重试重新调用 LLM)
|
|
20
|
+
DEFAULT_MAX_RETRIES = 2
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class LLMJudgeError(Exception):
|
|
24
|
+
"""LLM Judge 调用/解析最终失败 (重试用尽)。"""
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class LLMNotConfiguredError(Exception):
|
|
28
|
+
"""未注入 LLM 函数却调用了依赖 LLM 的指标 — 明确配置错误。"""
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def require_llm_fn(llm_fn: LLMFn | None) -> LLMFn:
|
|
32
|
+
"""断言 LLM 函数已注入, 否则抛出明确配置错误 (而非静默 0 分)。"""
|
|
33
|
+
if llm_fn is None:
|
|
34
|
+
raise LLMNotConfiguredError(
|
|
35
|
+
"LLM function not configured — inject llm_fn (eval_integration "
|
|
36
|
+
"assembles one from AEVAL_JUDGE_* / eval LLM settings) or pass a "
|
|
37
|
+
"stub in tests."
|
|
38
|
+
)
|
|
39
|
+
return llm_fn
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def extract_json_object(raw: str) -> dict[str, Any] | None:
|
|
43
|
+
"""
|
|
44
|
+
容错提取 LLM 输出中的 JSON 对象。
|
|
45
|
+
|
|
46
|
+
容忍 ```json 围栏、前后缀说明文本; 失败返回 None。
|
|
47
|
+
"""
|
|
48
|
+
text = (raw or "").strip()
|
|
49
|
+
fence = re.search(r"```(?:json)?\s*(.*?)\s*```", text, re.DOTALL)
|
|
50
|
+
if fence:
|
|
51
|
+
text = fence.group(1).strip()
|
|
52
|
+
|
|
53
|
+
start = text.find("{")
|
|
54
|
+
end = text.rfind("}")
|
|
55
|
+
if start < 0 or end <= start:
|
|
56
|
+
return None
|
|
57
|
+
try:
|
|
58
|
+
data = json.loads(text[start : end + 1])
|
|
59
|
+
except json.JSONDecodeError:
|
|
60
|
+
return None
|
|
61
|
+
return data if isinstance(data, dict) else None
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
async def judge_json(
|
|
65
|
+
llm_fn: LLMFn,
|
|
66
|
+
system_prompt: str,
|
|
67
|
+
user_prompt: str,
|
|
68
|
+
max_retries: int = DEFAULT_MAX_RETRIES,
|
|
69
|
+
) -> dict[str, Any]:
|
|
70
|
+
"""
|
|
71
|
+
调用 LLM 并解析 JSON 响应; 解析失败带完整上下文重试。
|
|
72
|
+
|
|
73
|
+
Raises:
|
|
74
|
+
LLMNotConfiguredError: llm_fn 为 None
|
|
75
|
+
LLMJudgeError: 重试用尽仍无法解析 (附最后一次原始输出)
|
|
76
|
+
"""
|
|
77
|
+
require_llm_fn(llm_fn)
|
|
78
|
+
|
|
79
|
+
last_raw = ""
|
|
80
|
+
last_error: Exception | None = None
|
|
81
|
+
for attempt in range(max_retries + 1):
|
|
82
|
+
try:
|
|
83
|
+
last_raw = await llm_fn(system_prompt, user_prompt)
|
|
84
|
+
except Exception as e:
|
|
85
|
+
# 传输层错误同样重试 (judge 输出不稳定/网络抖动, design §Risks)
|
|
86
|
+
last_error = e
|
|
87
|
+
if attempt < max_retries:
|
|
88
|
+
await asyncio.sleep(0)
|
|
89
|
+
continue
|
|
90
|
+
raise LLMJudgeError(f"LLM call failed after {max_retries + 1} attempts: {e}") from e
|
|
91
|
+
|
|
92
|
+
parsed = extract_json_object(last_raw)
|
|
93
|
+
if parsed is not None:
|
|
94
|
+
return parsed
|
|
95
|
+
last_error = LLMJudgeError("response is not a parsable JSON object")
|
|
96
|
+
|
|
97
|
+
raise LLMJudgeError(
|
|
98
|
+
f"LLM judge failed after {max_retries + 1} attempts ({last_error}); "
|
|
99
|
+
f"last response: {last_raw[:200]!r}"
|
|
100
|
+
)
|
|
@@ -0,0 +1,150 @@
|
|
|
1
|
+
"""Prompt variant A/B — PromptMetric (纯 LLM 层, 不经 AgentRunner).
|
|
2
|
+
|
|
3
|
+
对多个 Prompt 变体: 渲染模板 (str.format(**context)) → llm_fn 生成回答 →
|
|
4
|
+
指标打分 → n_trials 取平均 → 声明胜者。不经过 AgentRunner — Agent 行为层的
|
|
5
|
+
A/B 由 Dashboard run compare 覆盖, 本模块补 Prompt 层的快速对比通道。
|
|
6
|
+
|
|
7
|
+
胜者判定 (v1 简化语义): 各指标平均分求和最大者; 结果对象保留逐 trial 明细,
|
|
8
|
+
为 Phase 3 显著性检验预留。
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
from dataclasses import dataclass, field
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
from agent_eval.metrics.base import Metric, MetricResult
|
|
17
|
+
from agent_eval.metrics.llm_judge import LLMFn, require_llm_fn
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class PromptTemplateError(Exception):
|
|
21
|
+
"""变体模板渲染失败 (str.format 缺 key / 格式错误) — 校验错误。"""
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@dataclass
|
|
25
|
+
class PromptVariant:
|
|
26
|
+
"""Prompt 变体 (template 经 str.format(**context) 渲染)"""
|
|
27
|
+
|
|
28
|
+
name: str
|
|
29
|
+
template: str
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@dataclass
|
|
33
|
+
class PromptTrialDetail:
|
|
34
|
+
"""单次试验明细 (Phase 3 显著性检验预留)"""
|
|
35
|
+
|
|
36
|
+
output: str
|
|
37
|
+
scores: dict[str, float] = field(default_factory=dict) # metric name → score
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
@dataclass
|
|
41
|
+
class PromptComparisonResult:
|
|
42
|
+
"""变体对比结果 (含逐 trial 明细与 winner 标注)"""
|
|
43
|
+
|
|
44
|
+
variant_name: str
|
|
45
|
+
metric_scores: dict[str, float] = field(default_factory=dict) # n_trials 平均
|
|
46
|
+
trials: list[PromptTrialDetail] = field(default_factory=list)
|
|
47
|
+
winner: bool = False
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class PromptMetric:
|
|
51
|
+
"""
|
|
52
|
+
Prompt A/B 对比 — 渲染 → 生成 → 打分 → n_trials 平均 → 声明胜者。
|
|
53
|
+
|
|
54
|
+
用法:
|
|
55
|
+
pm = PromptMetric(variants=[...], metrics=[...], llm_fn=llm_fn)
|
|
56
|
+
results = await pm.compare(context={"question": "..."}, n_trials=3)
|
|
57
|
+
winner = pm.declare_winner(results)
|
|
58
|
+
|
|
59
|
+
胜者判定 (v1): 指标平均分求和最大 (平分取声明序首个)。
|
|
60
|
+
变体模板作为 system prompt 传给 llm_fn (user 为空串, 与蓝图一致)。
|
|
61
|
+
"""
|
|
62
|
+
|
|
63
|
+
def __init__(
|
|
64
|
+
self,
|
|
65
|
+
variants: list[PromptVariant],
|
|
66
|
+
metrics: list[Metric],
|
|
67
|
+
llm_fn: LLMFn | None = None,
|
|
68
|
+
):
|
|
69
|
+
if not variants:
|
|
70
|
+
raise ValueError("PromptMetric requires at least one variant")
|
|
71
|
+
if not metrics:
|
|
72
|
+
raise ValueError("PromptMetric requires at least one metric")
|
|
73
|
+
self.variants = list(variants)
|
|
74
|
+
self.metrics = list(metrics)
|
|
75
|
+
self.llm_fn = llm_fn
|
|
76
|
+
|
|
77
|
+
async def compare(
|
|
78
|
+
self,
|
|
79
|
+
context: dict[str, Any],
|
|
80
|
+
n_trials: int = 3,
|
|
81
|
+
) -> list[PromptComparisonResult]:
|
|
82
|
+
"""
|
|
83
|
+
对比所有变体 (每个变体 n_trials 次试验取平均)。
|
|
84
|
+
|
|
85
|
+
Raises:
|
|
86
|
+
LLMNotConfiguredError: 未注入 llm_fn (明确配置错误)
|
|
87
|
+
PromptTemplateError: 模板渲染缺 key / 格式错误 (校验错误)
|
|
88
|
+
ValueError: n_trials < 1
|
|
89
|
+
"""
|
|
90
|
+
llm_fn = require_llm_fn(self.llm_fn)
|
|
91
|
+
if n_trials < 1:
|
|
92
|
+
raise ValueError(f"n_trials must be >= 1, got {n_trials}")
|
|
93
|
+
|
|
94
|
+
results: list[PromptComparisonResult] = []
|
|
95
|
+
for variant in self.variants:
|
|
96
|
+
prompt = self._render(variant, context)
|
|
97
|
+
trials: list[PromptTrialDetail] = []
|
|
98
|
+
for _ in range(n_trials):
|
|
99
|
+
output = await llm_fn(prompt, "")
|
|
100
|
+
detail = PromptTrialDetail(output=output)
|
|
101
|
+
for metric in self.metrics:
|
|
102
|
+
result: MetricResult = await metric.measure(
|
|
103
|
+
input=prompt,
|
|
104
|
+
actual_output=output,
|
|
105
|
+
)
|
|
106
|
+
detail.scores[metric.name] = result.score
|
|
107
|
+
trials.append(detail)
|
|
108
|
+
|
|
109
|
+
metric_scores = {
|
|
110
|
+
metric.name: sum(t.scores[metric.name] for t in trials) / len(trials)
|
|
111
|
+
for metric in self.metrics
|
|
112
|
+
}
|
|
113
|
+
results.append(PromptComparisonResult(
|
|
114
|
+
variant_name=variant.name,
|
|
115
|
+
metric_scores=metric_scores,
|
|
116
|
+
trials=trials,
|
|
117
|
+
))
|
|
118
|
+
|
|
119
|
+
self.declare_winner(results)
|
|
120
|
+
return results
|
|
121
|
+
|
|
122
|
+
@staticmethod
|
|
123
|
+
def declare_winner(
|
|
124
|
+
results: list[PromptComparisonResult],
|
|
125
|
+
) -> PromptComparisonResult:
|
|
126
|
+
"""声明胜者 (v1 求和语义: 指标平均分求和最大; 平分取首个)。
|
|
127
|
+
|
|
128
|
+
原地标注: 恰有一个变体 winner=True (清除既有标注)。返回胜者。
|
|
129
|
+
"""
|
|
130
|
+
if not results:
|
|
131
|
+
raise ValueError("no comparison results to declare a winner from")
|
|
132
|
+
winner = max(results, key=lambda r: sum(r.metric_scores.values()))
|
|
133
|
+
for r in results:
|
|
134
|
+
r.winner = r is winner
|
|
135
|
+
return winner
|
|
136
|
+
|
|
137
|
+
@staticmethod
|
|
138
|
+
def _render(variant: PromptVariant, context: dict[str, Any]) -> str:
|
|
139
|
+
"""渲染变体模板; 缺 key / 格式错误转 PromptTemplateError (校验错误)。"""
|
|
140
|
+
try:
|
|
141
|
+
return variant.template.format(**context)
|
|
142
|
+
except KeyError as e:
|
|
143
|
+
raise PromptTemplateError(
|
|
144
|
+
f"variant '{variant.name}' template missing context key: {e} "
|
|
145
|
+
f"(template: {variant.template!r})"
|
|
146
|
+
) from e
|
|
147
|
+
except (IndexError, ValueError) as e:
|
|
148
|
+
raise PromptTemplateError(
|
|
149
|
+
f"variant '{variant.name}' template render failed: {e}"
|
|
150
|
+
) from e
|