chore: T-P-1 工作区收敛——并行会话成果与历史未入库文件整理入库
- 入库历史遗漏源码/测试:router_system 9 模块(agent/executors/inference/knowledge/ memory/planner/skills/trace)、tests 11 个测试文件、config/knowledge 领域知识 - 入库根目录方案文档(v2/v3/可行性×2)、references 文献(arxiv 14-18/cnki_open/ 参考文献清单)、research 论文素材(routerarena/paper/中文文献 PDF) - 前端构建产物刷新(新 hash);webapp 误写文档删除 - gitignore 增补:deepseek-harness、research/_refs、.mimosa/.zcode、网关日志/pid、 临时调试脚本、tests/e2e/node_modules、AI代理功能开发/prefix - 基线确认:318 passed
This commit is contained in:
@@ -0,0 +1,126 @@
|
||||
"""本系统 L0 路由器 → RouterArena BaseRouter 适配器。
|
||||
|
||||
核心职责:
|
||||
- 接收 RouterArena 的 prompt(不含 ground truth)
|
||||
- 调用本系统 Router.route() 得到 (domain, confidence, quality_score)
|
||||
- 按 `DOMAIN_TO_MODEL_SLOT` 映射到 config.models 中的目标模型名
|
||||
- 兜底:confidence<0.60 或 quality<0.70 → 升到 reasoning-mid
|
||||
|
||||
合规说明(依据 RouterArena README L34-37 + base_router.py L139-167):
|
||||
- 路由阶段不接触 ground truth answer
|
||||
- 不在 RouterArena 数据上训练/微调
|
||||
- 8 领域映射表是手工设计,不基于 RouterArena 标签学习
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict
|
||||
|
||||
# 把项目根加入 path,以便 import router_system
|
||||
_PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
if str(_PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(_PROJECT_ROOT))
|
||||
|
||||
from router_system.router import Router, build_router # noqa: E402
|
||||
|
||||
from .base_router import BaseRouter # noqa: E402
|
||||
|
||||
|
||||
# L0 域 → RouterArena 候选模型槽位 的映射
|
||||
# 依据 00_integration_plan.md §3.2:
|
||||
# - code → gpt-4o-mini(代码强项)
|
||||
# - math/legal/medical/finance → claude-3-haiku-20240307(推理严谨性)
|
||||
# - life/general → gemini-2.0-flash-001(实用+快速)
|
||||
# - education → deepseek-chat(中文/教育)
|
||||
DOMAIN_TO_MODEL_SLOT: Dict[str, str] = {
|
||||
"code": "gpt-4o-mini",
|
||||
"math": "claude-3-haiku-20240307",
|
||||
"legal": "claude-3-haiku-20240307",
|
||||
"medical": "claude-3-haiku-20240307",
|
||||
"finance": "claude-3-haiku-20240307",
|
||||
"life": "gemini-2.0-flash-001",
|
||||
"education": "deepseek-chat",
|
||||
"general": "gemini-2.0-flash-001",
|
||||
}
|
||||
|
||||
# 兜底:低置信度或低质 → 升到更稳的模型
|
||||
ESCALATION_MODEL_SLOT = "mistral-medium"
|
||||
|
||||
# 阈值:与本系统 config.yaml 默认对齐
|
||||
LOW_CONFIDENCE_THRESHOLD = 0.60
|
||||
JUDGE_FALLBACK_THRESHOLD = 0.70
|
||||
|
||||
|
||||
class ESExpertRouter(BaseRouter):
|
||||
"""把本系统 L0 Router 适配为 RouterArena BaseRouter。
|
||||
|
||||
重要:_get_prediction 必须只基于 query(无 ground truth),
|
||||
返回值必须在 config.models 中。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
router_name: str,
|
||||
config_path: str = "",
|
||||
underlying_router: Router = None,
|
||||
):
|
||||
super().__init__(router_name, config_path=config_path or None)
|
||||
# 复用本系统 L0 Router(默认配置;如已 build 过可注入)
|
||||
self._router = underlying_router or build_router()
|
||||
|
||||
def _get_prediction(self, query: str) -> str:
|
||||
"""依据 L0 路由决策返回目标模型名。"""
|
||||
# 注意:这里我们用同步方式跑异步 router.route()
|
||||
# BaseRouter 的 _get_prediction 是同步签名;RouterArena 的 generate_prediction_file
|
||||
# 默认单线程顺序调用 8400 条,async 包装完全兼容
|
||||
import asyncio
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
if loop.is_running():
|
||||
# 已经在 async 上下文(不太可能,但兜底)
|
||||
return self._get_prediction_sync(query)
|
||||
return loop.run_until_complete(self._route_one(query))
|
||||
except RuntimeError:
|
||||
# 没有 event loop,临时建一个
|
||||
return asyncio.run(self._route_one(query))
|
||||
|
||||
async def _route_one(self, query: str) -> str:
|
||||
result = await self._router.route(query)
|
||||
return self._decide(result.domain, result.confidence, result.quality_score)
|
||||
|
||||
def _get_prediction_sync(self, query: str) -> str:
|
||||
"""event loop 已运行时的兜底(按当前 router 的同步视图返回默认值)。"""
|
||||
# 我们没有同步入口,但 generate_prediction_file 是顺序同步调用 _get_prediction,
|
||||
# 不会与 async 上下文冲突,所以此分支极少触发;保守返回 generalist-fast
|
||||
return DOMAIN_TO_MODEL_SLOT["general"]
|
||||
|
||||
def _decide(self, domain: str, confidence: float, quality_score: float) -> str:
|
||||
"""路由决策:按阈值升档。"""
|
||||
if domain not in DOMAIN_TO_MODEL_SLOT:
|
||||
domain = "general"
|
||||
# 兜底升级
|
||||
if confidence < LOW_CONFIDENCE_THRESHOLD or quality_score < JUDGE_FALLBACK_THRESHOLD:
|
||||
return ESCALATION_MODEL_SLOT
|
||||
return DOMAIN_TO_MODEL_SLOT[domain]
|
||||
|
||||
def diagnostics(self, query: str) -> Dict[str, Any]:
|
||||
"""返回完整路由诊断(用于科研报告,不影响 RouterArena 协议)。"""
|
||||
import asyncio
|
||||
return asyncio.run(self._diagnose_one(query))
|
||||
|
||||
async def _diagnose_one(self, query: str) -> Dict[str, Any]:
|
||||
result = await self._router.route(query)
|
||||
return {
|
||||
"query": query,
|
||||
"domain": result.domain,
|
||||
"subdomain": result.subdomain,
|
||||
"subdomain2": result.subdomain2,
|
||||
"difficulty": result.difficulty,
|
||||
"confidence": result.confidence,
|
||||
"quality_score": result.quality_score,
|
||||
"upgraded": result.upgraded,
|
||||
"model_used": result.model_used,
|
||||
"route": result.route,
|
||||
"selected_slot": self._decide(result.domain, result.confidence, result.quality_score),
|
||||
}
|
||||
Reference in New Issue
Block a user