- 新增 router_system/explain.py:CandidateProfile(capability/cost/latency/ reliability/difficulty 画像)+ explain_routing(硬过滤先行、逐条人话拒绝理由、 五维归一评分、同分按领域字典序、confidence=0.8+分差推导封顶 0.99) - Router._explain 纯解释层装配(解释层任何异常不影响主链路,D-G4 同款纪律); 路由胜负与既有决策完全等价,/chat 响应新增 route_explanation 字段(缓存命中为 None) - experts:Expert.nominal_latency_ms 标称延迟元数据(mock 1/hf 300/api 800) - stats:upgraded_by_domain 计数 + domain_reliability()(无数据给 0.9 中性先验) - .gitignore 登记 extra/(参考项目目录不入库) - 新增 tests/test_explain.py 8 项;全量 33 passed(基线 25)
107 lines
4.9 KiB
Python
107 lines
4.9 KiB
Python
"""可解释路由评分层(T-R1,采纳 ai-model-router 五维评分 + 硬过滤带拒绝理由设计)。
|
||
|
||
对分类产生的候选领域专家做透明评分与解释。路由胜负保持既有逻辑不变
|
||
(本模块是纯解释层,不参与决策):
|
||
|
||
- 硬过滤先行:无规则命中/低置信度/缺专家的候选先淘汰,逐条记录人话拒绝理由
|
||
- 五维加权评分(各维归一到 0-1,权重可配):
|
||
capability 分类器该领域原始分归一(语义匹配强度)
|
||
cost_efficiency 1/(1+cost*100) 平滑映射,零成本本地模型恒 1.0
|
||
speed 1/(1+标称延迟/200)(200ms 参考延迟)
|
||
reliability Judge 通过率(无数据时 0.9 中性先验)
|
||
quality 难度 × 领域纵深启发式
|
||
- 置信度可推导:confidence = 0.8 + (最高分 - 次高分),分差越大越自信(封顶 0.99)
|
||
- 同分决胜按领域名字典序(与仓库既有确定性约定一致)
|
||
|
||
输出 RoutingExplanation.to_dict() 直接进 /chat 响应的 route_explanation 字段。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
from dataclasses import dataclass, field
|
||
from typing import Any, Dict, List, Optional
|
||
|
||
# 领域纵深启发式(该领域知识可承载的任务深度,0-1)
|
||
_DOMAIN_DEPTH = {
|
||
"code": 0.90, "math": 0.90, "legal": 0.85, "medical": 0.85, "general": 0.60,
|
||
}
|
||
|
||
DEFAULT_WEIGHTS: Dict[str, float] = {
|
||
"capability": 0.35,
|
||
"cost_efficiency": 0.20,
|
||
"speed": 0.15,
|
||
"reliability": 0.20,
|
||
"quality": 0.10,
|
||
}
|
||
|
||
|
||
@dataclass
|
||
class CandidateProfile:
|
||
"""一个候选领域专家的评分输入(由 Router 从分类/专家池/统计装配)。"""
|
||
domain: str
|
||
model: str # 专家模型名
|
||
capability: float # 分类器该领域原始分归一 0-1
|
||
cost_per_mtok: float # $ / 1M output tokens 估计
|
||
nominal_latency_ms: float # 标称延迟(专家后端元数据)
|
||
reliability: float # Judge 通过率(无数据给中性先验)
|
||
difficulty_score: float # 查询难度分 0-1(质量维输入)
|
||
hard_fail_reason: Optional[str] = None # 非 None 即硬过滤淘汰(人话理由)
|
||
|
||
|
||
@dataclass
|
||
class RoutingExplanation:
|
||
"""一次路由的完整解释:胜者 + 排名 + 拒绝名单 + 权重快照。"""
|
||
winner: str # 胜出领域;无存活候选时为 "fallback"
|
||
confidence: float # 0.8 + 分差推导(封顶 0.99)
|
||
ranked: List[Dict[str, Any]] = field(default_factory=list)
|
||
rejected: List[Dict[str, Any]] = field(default_factory=list)
|
||
weights: Dict[str, float] = field(default_factory=dict)
|
||
|
||
def to_dict(self) -> Dict[str, Any]:
|
||
return {
|
||
"winner": self.winner,
|
||
"confidence": round(self.confidence, 4),
|
||
"ranked": self.ranked,
|
||
"rejected": self.rejected,
|
||
"weights": dict(self.weights),
|
||
}
|
||
|
||
|
||
def score_candidate(c: CandidateProfile, weights: Dict[str, float]) -> Dict[str, Any]:
|
||
"""单候选五维评分(全部绝对映射 0-1,跨候选可比、可单测)。"""
|
||
dims = {
|
||
"capability": max(0.0, min(1.0, c.capability)),
|
||
"cost_efficiency": 1.0 / (1.0 + max(0.0, c.cost_per_mtok) * 100.0),
|
||
"speed": 1.0 / (1.0 + max(0.0, c.nominal_latency_ms) / 200.0),
|
||
"reliability": max(0.0, min(1.0, c.reliability)),
|
||
"quality": max(0.0, min(1.0, c.difficulty_score)) * _DOMAIN_DEPTH.get(c.domain, 0.60),
|
||
}
|
||
total = sum(dims[k] * weights.get(k, 0.0) for k in dims)
|
||
return {"domain": c.domain, "model": c.model,
|
||
"total": round(total, 4), "dims": {k: round(v, 4) for k, v in dims.items()}}
|
||
|
||
|
||
def explain_routing(candidates: List[CandidateProfile],
|
||
weights: Optional[Dict[str, float]] = None) -> RoutingExplanation:
|
||
"""硬过滤 -> 五维加权 -> 排名与置信度推导(确定性:同分按领域名字典序)。"""
|
||
w = dict(DEFAULT_WEIGHTS if weights is None else weights)
|
||
rejected: List[Dict[str, Any]] = []
|
||
alive: List[CandidateProfile] = []
|
||
for c in candidates:
|
||
if c.hard_fail_reason is not None:
|
||
rejected.append({"domain": c.domain, "reason": c.hard_fail_reason})
|
||
else:
|
||
alive.append(c)
|
||
|
||
scored = [score_candidate(c, w) for c in alive]
|
||
scored.sort(key=lambda s: (-s["total"], s["domain"]))
|
||
|
||
if not scored:
|
||
# 全部被硬过滤(含低置信度直连回退路径):决策即走大模型回退
|
||
return RoutingExplanation(winner="fallback", confidence=0.80,
|
||
ranked=[], rejected=rejected, weights=w)
|
||
|
||
gap = scored[0]["total"] - (scored[1]["total"] if len(scored) > 1 else 0.0)
|
||
confidence = min(0.99, 0.80 + max(0.0, gap))
|
||
return RoutingExplanation(winner=scored[0]["domain"], confidence=confidence,
|
||
ranked=scored, rejected=rejected, weights=w)
|