feat: 多专业小模型+路由模型系统 MVP(mock 全链路 + FastAPI 网关 + 论文调研)
This commit is contained in:
@@ -0,0 +1,217 @@
|
||||
"""主路由器:协调 缓存 -> 分类 -> 专家 -> Judge -> 大模型回退 的完整链路。
|
||||
|
||||
流程(对齐实现方案):
|
||||
1. 检查缓存(L1 精确 / L2 语义)
|
||||
2. 低置信度查询直接走大模型(should_fallback)
|
||||
3. 分类器输出领域 + 难度
|
||||
4. 选择专家模型生成
|
||||
5. Judge 评估质量
|
||||
6. 质量不达标 -> 升级大模型
|
||||
7. 记录指标、写缓存、返回结果
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from .cache import RouterCache
|
||||
from .classifier import BaseClassifier, build_classifier
|
||||
from .config import load_config
|
||||
from .experts import Expert, build_expert_pool
|
||||
from .fallback import FallbackProvider, build_fallback
|
||||
from .judge import BaseJudge, build_judge
|
||||
from .models import Classification, ExpertResponse, RouterResult, now_ms
|
||||
from .stats import Stats
|
||||
|
||||
|
||||
class Router:
|
||||
def __init__(
|
||||
self,
|
||||
classifier: BaseClassifier,
|
||||
experts: Dict[str, Expert],
|
||||
judge: BaseJudge,
|
||||
fallback: FallbackProvider,
|
||||
cache: Optional[RouterCache] = None,
|
||||
stats: Optional[Stats] = None,
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.classifier = classifier
|
||||
self.experts = experts
|
||||
self.judge = judge
|
||||
self.fallback = fallback
|
||||
self.cache = cache or RouterCache()
|
||||
self.stats = stats or Stats()
|
||||
cfg = config or {}
|
||||
rcfg = cfg.get("router", {})
|
||||
self.low_confidence_threshold = rcfg.get("low_confidence_threshold", 0.60)
|
||||
self.judge_fallback_threshold = rcfg.get("judge_fallback_threshold", 0.70)
|
||||
self.cache_enabled = cfg.get("cache", {}).get("enabled", True)
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
async def route(self, query: str) -> RouterResult:
|
||||
start = now_ms()
|
||||
route: list = []
|
||||
|
||||
# ---- Step 1: 缓存 ----
|
||||
if self.cache_enabled:
|
||||
hit = self.cache.get(query)
|
||||
if hit is not None:
|
||||
level, cached = hit
|
||||
latency = now_ms() - start
|
||||
result = RouterResult(
|
||||
query=query,
|
||||
response=cached.get("response", ""),
|
||||
domain=cached.get("domain", "general"),
|
||||
difficulty=cached.get("difficulty", "medium"),
|
||||
confidence=cached.get("confidence", 0.0),
|
||||
upgraded=False,
|
||||
quality_score=cached.get("quality_score", 0.0),
|
||||
model_used=cached.get("model_used", ""),
|
||||
route=["cache:" + level],
|
||||
latency_ms=latency,
|
||||
cache_hit=True,
|
||||
cache_level=level,
|
||||
cost_est=0.0,
|
||||
)
|
||||
self.stats.record(latency, result.domain, result.difficulty, False, True, level, 0.0, result.model_used)
|
||||
return result
|
||||
route.append("cache:miss")
|
||||
|
||||
# ---- Step 2: 分类 ----
|
||||
classification = self.classifier.classify(query)
|
||||
route.append(f"classify:{classification.domain}@{classification.confidence:.2f}/{classification.difficulty}")
|
||||
|
||||
# 低置信度 -> 直接走大模型
|
||||
if self.classifier.should_fallback(classification, self.low_confidence_threshold):
|
||||
route.append("direct_fallback")
|
||||
fb = await self._call_fallback(query)
|
||||
latency = now_ms() - start
|
||||
result = self._finalize(query, classification, fb, quality_score=0.0,
|
||||
upgraded=True, route=route, latency_ms=latency,
|
||||
model_used=fb.model_used, cost_est=fb.cost_est)
|
||||
self._record(result, latency)
|
||||
return result
|
||||
|
||||
# ---- Step 3: 选择专家 ----
|
||||
domain = classification.domain
|
||||
expert = self.experts.get(domain)
|
||||
if expert is None:
|
||||
expert = self.experts.get("general")
|
||||
route.append("expert:fallback-to-general")
|
||||
else:
|
||||
route.append(f"expert:{expert.name}")
|
||||
|
||||
# ---- Step 4: 生成 ----
|
||||
try:
|
||||
expert_resp = await expert.generate(query, classification.difficulty)
|
||||
except Exception as e:
|
||||
self.stats.record_error()
|
||||
route.append(f"expert_error:{type(e).__name__}")
|
||||
fb = await self._call_fallback(query)
|
||||
latency = now_ms() - start
|
||||
result = self._finalize(query, classification, fb, quality_score=0.0,
|
||||
upgraded=True, route=route, latency_ms=latency,
|
||||
model_used=fb.model_used, cost_est=fb.cost_est,
|
||||
error=str(e))
|
||||
self._record(result, latency)
|
||||
return result
|
||||
|
||||
# ---- Step 5: Judge 评估 ----
|
||||
try:
|
||||
evaluation = await self.judge.evaluate(query, expert_resp.text, domain)
|
||||
except Exception:
|
||||
evaluation = None
|
||||
route.append("judge_error")
|
||||
|
||||
quality_score = evaluation.overall_score if evaluation else 0.0
|
||||
route.append(f"judge:{quality_score:.2f}")
|
||||
|
||||
upgraded = False
|
||||
final_resp = expert_resp
|
||||
if evaluation is not None and evaluation.needs_fallback:
|
||||
route.append("upgrade")
|
||||
final_resp = await self._call_fallback(query)
|
||||
upgraded = True
|
||||
|
||||
latency = now_ms() - start
|
||||
result = self._finalize(query, classification, final_resp, quality_score=quality_score,
|
||||
upgraded=upgraded, route=route, latency_ms=latency,
|
||||
model_used=final_resp.model_used, cost_est=final_resp.cost_est)
|
||||
self._record(result, latency)
|
||||
|
||||
# 未升级的结果写缓存
|
||||
if self.cache_enabled and not upgraded and result.response:
|
||||
self.cache.put(query, result.to_dict())
|
||||
|
||||
return result
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
async def _call_fallback(self, query: str) -> ExpertResponse:
|
||||
try:
|
||||
return await self.fallback.generate(query)
|
||||
except Exception as e:
|
||||
# 回退也失败:返回错误占位响应
|
||||
return ExpertResponse(
|
||||
text=f"[系统错误] 专家与大模型回退均失败:{type(e).__name__}: {e}",
|
||||
model_used=f"error:{self.fallback.name}",
|
||||
cost_est=0.0,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _finalize(query: str, classification: Classification, resp: ExpertResponse,
|
||||
quality_score: float, upgraded: bool, route: list,
|
||||
latency_ms: float, model_used: str, cost_est: float,
|
||||
error: Optional[str] = None) -> RouterResult:
|
||||
return RouterResult(
|
||||
query=query,
|
||||
response=resp.text,
|
||||
domain=classification.domain,
|
||||
difficulty=classification.difficulty,
|
||||
confidence=classification.confidence,
|
||||
upgraded=upgraded,
|
||||
quality_score=quality_score,
|
||||
model_used=model_used,
|
||||
route=route,
|
||||
latency_ms=latency_ms,
|
||||
cache_hit=False,
|
||||
cost_est=cost_est,
|
||||
error=error,
|
||||
)
|
||||
|
||||
def _record(self, result: RouterResult, latency_ms: float):
|
||||
self.stats.record(
|
||||
latency_ms,
|
||||
result.domain,
|
||||
result.difficulty,
|
||||
result.upgraded,
|
||||
result.cache_hit,
|
||||
result.cache_level,
|
||||
result.cost_est,
|
||||
result.model_used,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
def health(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"status": "ok",
|
||||
"domains": list(self.experts.keys()),
|
||||
"classifier": type(self.classifier).__name__,
|
||||
"judge": type(self.judge).__name__,
|
||||
"fallback": type(self.fallback).__name__,
|
||||
}
|
||||
|
||||
|
||||
def build_router(config_path: Optional[str] = None) -> Router:
|
||||
"""从配置构建完整 Router(默认 mock 全链路,零依赖可跑)。"""
|
||||
config = load_config(config_path)
|
||||
classifier = build_classifier(config.get("classifier", {}))
|
||||
experts = build_expert_pool(config.get("experts", {}), config.get("domains", []))
|
||||
judge = build_judge(config.get("judge", {}), config.get("router", {}).get("judge_fallback_threshold", 0.70))
|
||||
fallback = build_fallback(config.get("fallback", {}))
|
||||
cache_cfg = config.get("cache", {})
|
||||
cache = RouterCache(
|
||||
semantic_enabled=cache_cfg.get("semantic_enabled", True),
|
||||
similarity_threshold=cache_cfg.get("similarity_threshold", 0.88),
|
||||
promote_frequency=cache_cfg.get("promote_frequency", 5),
|
||||
)
|
||||
stats = Stats()
|
||||
return Router(classifier, experts, judge, fallback, cache, stats, config)␍
|
||||
Reference in New Issue
Block a user