Files

194 lines
7.1 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""质量控制器(Judge):评估专家输出,决定是否升级大模型。
- RuleJudge:零依赖启发式(内容覆盖度 / 长度充分性 / 领域格式 / 安全提示),
稳定可测,适合 MVP 与离线演示。
- LLMJudge:可选,基于 transformers 小模型或 API 的 LLM-as-Judge。
设计对齐实现方案:overall_score < judge_fallback_threshold -> 升级大模型。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Dict, List
from .experts import extract_content_terms
# 各领域期望的响应长度范围(字符数)
_EXPECTED_LEN = {
"code": (60, 2000),
"math": (60, 2000),
"legal": (80, 3000),
"medical": (80, 3000),
"general": (40, 2000),
}
# 领域格式检查:响应应包含的标记
_DOMAIN_FORMAT_HINTS = {
"code": ["```", "def ", "function", "class "],
"math": ["步骤", "推导", "=", "解", "step"],
"legal": ["⚠", "法律", "意见", "合规", "contract", "law"],
"medical": ["⚠", "就医", "医生", "症状", "诊断", "symptom"],
"general": [],
}
@dataclass
class QualityEvaluation:
overall_score: float
scores: Dict[str, float] = field(default_factory=dict)
needs_fallback: bool = False
reasons: List[str] = field(default_factory=list)
def to_dict(self) -> Dict:
return {
"overall_score": round(self.overall_score, 4),
"scores": {k: round(v, 4) for k, v in self.scores.items()},
"needs_fallback": self.needs_fallback,
"reasons": self.reasons,
}
class BaseJudge:
async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation:
raise NotImplementedError
class RuleJudge(BaseJudge):
"""启发式质量评估(零依赖)。"""
def __init__(self, fallback_threshold: float = 0.70, kb=None):
self.fallback_threshold = fallback_threshold
self.kb = kb # 可选知识库:facts 维度校验用
async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation:
scores: Dict[str, float] = {}
reasons: List[str] = []
# 1) 内容覆盖度:查询中的内容词有多少出现在响应里
terms = extract_content_terms(query)
if terms:
hit = sum(1 for t in terms if t in response.lower())
coverage = hit / len(terms)
scores["coverage"] = coverage
if coverage < 0.4:
reasons.append(f"内容覆盖度低 ({coverage:.0%})")
else:
scores["coverage"] = 1.0
# 2) 长度充分性
lo, hi = _EXPECTED_LEN.get(domain, (40, 2000))
n = len(response)
if n < lo:
scores["length"] = max(0.0, n / lo)
reasons.append(f"响应过短 ({n} 字符)")
elif n > hi:
scores["length"] = 0.8
reasons.append(f"响应过长 ({n} 字符)")
else:
scores["length"] = 1.0
# 3) 领域格式检查
hints = _DOMAIN_FORMAT_HINTS.get(domain, [])
if hints:
hit_hints = sum(1 for h in hints if h in response)
scores["format"] = min(1.0, 0.4 + 0.2 * hit_hints)
if hit_hints == 0:
reasons.append("缺少领域格式特征")
else:
scores["format"] = 1.0
# 4) 安全/免责提示(法律、医疗领域应有警示语)
if domain in ("legal", "medical") and ("⚠" not in response and "提示" not in response):
scores["safety"] = 0.6
reasons.append("缺少免责提示")
else:
scores["safety"] = 1.0
# 5) 领域知识引用(facts 维度):法律/医疗响应应覆盖查询命中的知识条目
if domain in ("legal", "medical") and self.kb is not None:
facts = self.kb.facts(domain)
hit_facts = [f for f in facts if any(k in query for k in f.get("keywords", []))]
if hit_facts:
ok = sum(
1 for f in hit_facts
if any(k in response for k in f.get("keywords", []))
)
scores["facts"] = ok / len(hit_facts)
if scores["facts"] < 0.6:
reasons.append(f"领域知识引用不足 ({scores['facts']:.0%})")
else:
scores["facts"] = 1.0
else:
scores["facts"] = 1.0
weights = {"coverage": 0.35, "length": 0.15, "format": 0.15,
"safety": 0.15, "facts": 0.2}
overall = sum(scores.get(k, 0.0) * w for k, w in weights.items())
needs = overall < self.fallback_threshold
if needs:
reasons.append("质量分低于阈值,建议升级大模型")
return QualityEvaluation(
overall_score=round(overall, 4),
scores=scores,
needs_fallback=needs,
reasons=reasons,
)
class LLMJudge(BaseJudge):
"""可选:LLM-as-JudgeAPI 后端)。"""
def __init__(self, model: str, base_url: str, api_key: str, fallback_threshold: float = 0.70):
self.model = model
self.base_url = base_url.rstrip("/")
self.api_key = api_key
self.fallback_threshold = fallback_threshold
self._client = None
def _get_client(self):
if self._client is None:
import httpx
self._client = httpx.AsyncClient(timeout=60.0)
return self._client
async def evaluate(self, query: str, response: str, domain: str) -> QualityEvaluation:
client = self._get_client()
prompt = (
f"你是质量评审员。评估以下回答对查询的满足程度,输出 0-1 分(相关性/正确性/完整性)。\n"
f"查询: {query}\n领域: {domain}\n回答: {response[:2000]}\n"
f"只输出一个 0 到 1 之间的数字。"
)
resp = await client.post(
f"{self.base_url}/chat/completions",
headers={"Authorization": f"Bearer {self.api_key}"},
json={"model": self.model, "messages": [{"role": "user", "content": prompt}]},
)
resp.raise_for_status()
try:
score = float(resp.json()["choices"][0]["message"]["content"].strip())
score = max(0.0, min(1.0, score))
except Exception:
score = 0.5
return QualityEvaluation(
overall_score=score,
scores={"llm_judge": score},
needs_fallback=score < self.fallback_threshold,
)
def build_judge(cfg: Dict, fallback_threshold: float = 0.70, kb=None) -> BaseJudge:
"""cfg 为 judge 段配置。"""
jtype = cfg.get("type", "rule")
if jtype == "rule":
return RuleJudge(fallback_threshold=fallback_threshold, kb=kb)
if jtype == "llm":
import os
api_key = cfg.get("api_key") or os.environ.get(cfg.get("api_key_env", "API_KEY"))
return LLMJudge(
cfg.get("model", "deepseek-v4-flash"),
cfg.get("base_url", "https://api.deepseek.com"),
api_key or "",
fallback_threshold=fallback_threshold,
)
raise ValueError(f"未知 judge 类型: {jtype}(支持 rule | llm")