Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9fdf91c2fb | ||
|
|
66b6fd8c3e |
@@ -25,3 +25,6 @@ cached_results/
|
|||||||
Thumbs.db
|
Thumbs.db
|
||||||
.idea/
|
.idea/
|
||||||
.vscode/
|
.vscode/
|
||||||
|
|
||||||
|
# 安全扫描器工作目录(不入库)
|
||||||
|
.mimosa/
|
||||||
|
|||||||
@@ -0,0 +1,29 @@
|
|||||||
|
# 分支:v1-model-routing — 本地多智能体协作模型路由(第一代)
|
||||||
|
|
||||||
|
> **快照点**:`1e51167`(v1 MVP 完成时点,仓库首个提交)。
|
||||||
|
> 此分支为**历史路标**,冻结不再演进;集成主线见 `master`。
|
||||||
|
|
||||||
|
## 这一代是什么
|
||||||
|
|
||||||
|
- **核心命题**:用"规则知识库 + 任务拆解 Planner + 专业执行器池 + 质量控制器(Judge) + 最后处理者"
|
||||||
|
在限定条件下替代单一通用大模型——**本地多智能体协作路由**
|
||||||
|
- 架构:8 领域规则知识库 → 规则分类器(置信度 1−e^−s)→ Planner 任务拆解(DAG)
|
||||||
|
→ 黑板/前向链 → 规则执行器(L0 零参数、确定性模板)→ RuleJudge 五维评分 → mock 兜底
|
||||||
|
- 两级路由:`domain_groups`(tech/professional/lifestyle/general)→ 组内 RuleClassifier;三级子领域识别
|
||||||
|
- 接口:FastAPI `/chat` `/health` `/metrics`;核心包 `router_system/` 零第三方依赖
|
||||||
|
|
||||||
|
## 如何运行
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
.venv\Scripts\python.exe scripts/demo.py --trace # L0 演示(离线可跑)
|
||||||
|
.venv\Scripts\python.exe scripts/eval.py # 迷你评估
|
||||||
|
.venv\Scripts\python.exe -m pytest tests -q # 测试(初版 20 项)
|
||||||
|
.venv\Scripts\python.exe scripts/serve.py --port 8000 # 网关
|
||||||
|
```
|
||||||
|
|
||||||
|
## 定位与结论
|
||||||
|
|
||||||
|
- **方向结论**:规则路由在 9 域平衡集实测 74.4% 准确率、68.9% 升级率(`research/routerarena/`)——
|
||||||
|
假阳性过高,**被 v2 端云协作风取代**;本代代码在主线保留为 legacy(`POST /chat/legacy`)
|
||||||
|
- 设计文档:`实现方案_多专业小模型+路由模型.md`、`可行性调研与落地实现路线报告.md`、
|
||||||
|
`research/2026_papers_survey.md`、`research/routerarena/01_results_and_gap_analysis.md`
|
||||||
+155
-141
@@ -1,141 +1,155 @@
|
|||||||
"""两阶段路由缓存(对齐实现方案):
|
"""两阶段路由缓存(对齐实现方案):
|
||||||
- L1 精确缓存:完全相同的查询 -> 直接命中
|
- L1 精确缓存:完全相同的查询 -> 直接命中
|
||||||
- L2 语义缓存:字符 n-gram 余弦相似度(零依赖)-> 相似查询命中
|
- L2 语义缓存:字符 n-gram 余弦相似度(零依赖)-> 相似查询命中
|
||||||
- 命中 N 次(promote_frequency)后提升为精确缓存
|
- 命中 N 次(promote_frequency)后提升为精确缓存
|
||||||
|
|
||||||
说明:语义缓存中的"完全相同查询"(相似度=1.0)直接计为 exact 命中;
|
说明:语义缓存中的"完全相同查询"(相似度=1.0)直接计为 exact 命中;
|
||||||
高频语义命中会提升为 O(1) 的精确缓存条目。
|
高频语义命中会提升为 O(1) 的精确缓存条目。
|
||||||
|
|
||||||
只缓存"未升级"的结果(升级路径每次都走大模型,不缓存,避免陈旧)。
|
只缓存"未升级"的结果(升级路径每次都走大模型,不缓存,避免陈旧)。
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
性能设计(2026-09 优化):
|
||||||
|
- 每条语义缓存条目在写入时预计算并缓存向量范数,查询时免重复计算(原来每对比较都重算)
|
||||||
import re
|
- 语义查找单遍完成:扫描即跟踪最优条目与命中计数,命中后不再二次线性查找
|
||||||
from dataclasses import dataclass
|
- 相似度达到 1.0(完全相同查询)时提前终止扫描(余弦相似度上界,不可能更优)
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
@dataclass
|
import re
|
||||||
class CacheEntry:
|
from dataclasses import dataclass
|
||||||
result: Dict[str, Any]
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
hits: int = 1
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
def _ngrams(text: str, n: int = 3) -> List[str]:
|
class CacheEntry:
|
||||||
"""字符 n-gram(去空白、小写),用于轻量语义相似度。"""
|
result: Dict[str, Any]
|
||||||
cleaned = re.sub(r"\s+", "", text.lower())
|
hits: int = 1
|
||||||
if len(cleaned) < n:
|
|
||||||
return [cleaned]
|
|
||||||
return [cleaned[i:i + n] for i in range(len(cleaned) - n + 1)]
|
def _ngrams(text: str, n: int = 3) -> List[str]:
|
||||||
|
"""字符 n-gram(去空白、小写),用于轻量语义相似度。"""
|
||||||
|
cleaned = re.sub(r"\s+", "", text.lower())
|
||||||
def _cosine(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float:
|
if len(cleaned) < n:
|
||||||
if not vec_a or not vec_b:
|
return [cleaned]
|
||||||
return 0.0
|
return [cleaned[i:i + n] for i in range(len(cleaned) - n + 1)]
|
||||||
common = set(vec_a) & set(vec_b)
|
|
||||||
dot = sum(vec_a[k] * vec_b[k] for k in common)
|
|
||||||
na = sum(v * v for v in vec_a.values()) ** 0.5
|
def _tf_vector(grams: List[str]) -> Dict[str, float]:
|
||||||
nb = sum(v * v for v in vec_b.values()) ** 0.5
|
vec: Dict[str, float] = {}
|
||||||
if na == 0 or nb == 0:
|
for g in grams:
|
||||||
return 0.0
|
vec[g] = vec.get(g, 0.0) + 1.0
|
||||||
return dot / (na * nb)
|
return vec
|
||||||
|
|
||||||
|
|
||||||
def _tf_vector(grams: List[str]) -> Dict[str, float]:
|
def _norm(vec: Dict[str, float]) -> float:
|
||||||
vec: Dict[str, float] = {}
|
return sum(v * v for v in vec.values()) ** 0.5
|
||||||
for g in grams:
|
|
||||||
vec[g] = vec.get(g, 0.0) + 1.0
|
|
||||||
return vec
|
def _dot(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float:
|
||||||
|
"""点积:遍历较小的一方,另一侧用 get 兜底。"""
|
||||||
|
if len(vec_a) > len(vec_b):
|
||||||
class RouterCache:
|
vec_a, vec_b = vec_b, vec_a
|
||||||
"""L1 精确缓存 + L2 语义缓存。"""
|
return sum(v * vec_b.get(k, 0.0) for k, v in vec_a.items())
|
||||||
|
|
||||||
def __init__(self, semantic_enabled: bool = True, similarity_threshold: float = 0.88,
|
|
||||||
promote_frequency: int = 5, max_exact: int = 10000, max_semantic: int = 5000):
|
class RouterCache:
|
||||||
self.semantic_enabled = semantic_enabled
|
"""L1 精确缓存 + L2 语义缓存。"""
|
||||||
self.similarity_threshold = similarity_threshold
|
|
||||||
self.promote_frequency = promote_frequency
|
def __init__(self, semantic_enabled: bool = True, similarity_threshold: float = 0.88,
|
||||||
self.max_exact = max_exact
|
promote_frequency: int = 5, max_exact: int = 10000, max_semantic: int = 5000):
|
||||||
self.max_semantic = max_semantic
|
self.semantic_enabled = semantic_enabled
|
||||||
self._exact: Dict[str, CacheEntry] = {}
|
self.similarity_threshold = similarity_threshold
|
||||||
self._semantic: List[Tuple[str, CacheEntry]] = [] # (query, entry)
|
self.promote_frequency = promote_frequency
|
||||||
self._sem_vecs: Dict[str, Dict[str, float]] = {}
|
self.max_exact = max_exact
|
||||||
self.hits = {"exact": 0, "semantic": 0}
|
self.max_semantic = max_semantic
|
||||||
self.misses = 0
|
self._exact: Dict[str, CacheEntry] = {}
|
||||||
|
self._semantic: List[Tuple[str, CacheEntry]] = [] # (query, entry)
|
||||||
# ---- 查询 ----
|
self._sem_vecs: Dict[str, Dict[str, float]] = {}
|
||||||
def get(self, query: str) -> Optional[Tuple[Optional[str], Dict[str, Any]]]:
|
self._sem_norms: Dict[str, float] = {} # 预计算范数,避免查询期重算
|
||||||
"""返回 (level, result);未命中返回 None。level: 'exact' | 'semantic'"""
|
self.hits = {"exact": 0, "semantic": 0}
|
||||||
entry = self._exact.get(query)
|
self.misses = 0
|
||||||
if entry is not None:
|
|
||||||
self.hits["exact"] += 1
|
# ---- 查询 ----
|
||||||
return ("exact", entry.result)
|
def get(self, query: str) -> Optional[Tuple[Optional[str], Dict[str, Any]]]:
|
||||||
|
"""返回 (level, result);未命中返回 None。level: 'exact' | 'semantic'"""
|
||||||
if self.semantic_enabled:
|
entry = self._exact.get(query)
|
||||||
q_vec = _tf_vector(_ngrams(query))
|
if entry is not None:
|
||||||
best_sim = 0.0
|
self.hits["exact"] += 1
|
||||||
best_query: Optional[str] = None
|
return ("exact", entry.result)
|
||||||
best_result: Optional[Dict[str, Any]] = None
|
|
||||||
for q, e in self._semantic:
|
if self.semantic_enabled:
|
||||||
sim = _cosine(q_vec, self._sem_vecs.get(q, {}))
|
q_vec = _tf_vector(_ngrams(query))
|
||||||
if sim > best_sim:
|
q_norm = _norm(q_vec)
|
||||||
best_sim = sim
|
best_sim = 0.0
|
||||||
best_query = q
|
best_idx = -1
|
||||||
best_result = e.result
|
if q_norm > 0.0:
|
||||||
if best_query is not None and best_sim >= self.similarity_threshold:
|
# 单遍扫描:同时跟踪最优相似度与条目位置
|
||||||
# 完全相同查询(相似度=1.0)计为 exact 命中
|
for i, (q, _e) in enumerate(self._semantic):
|
||||||
is_exact = best_sim >= 0.999
|
n_q = self._sem_norms.get(q, 0.0)
|
||||||
level = "exact" if is_exact else "semantic"
|
if n_q <= 0.0:
|
||||||
self.hits[level] += 1
|
continue
|
||||||
self._semantic_hit(best_query)
|
sim = _dot(q_vec, self._sem_vecs.get(q, {})) / (q_norm * n_q)
|
||||||
return (level, best_result)
|
if sim > best_sim:
|
||||||
|
best_sim = sim
|
||||||
self.misses += 1
|
best_idx = i
|
||||||
return None
|
if sim >= 1.0:
|
||||||
|
break # 余弦相似度上界:完全相同查询,提前终止
|
||||||
def _semantic_hit(self, query: str):
|
if best_idx >= 0 and best_sim >= self.similarity_threshold:
|
||||||
"""语义命中:累计命中次数,达到阈值提升为精确缓存。"""
|
best_q, best_entry = self._semantic[best_idx]
|
||||||
for i, (q, e) in enumerate(self._semantic):
|
# 完全相同查询(相似度=1.0)计为 exact 命中
|
||||||
if q == query:
|
is_exact = best_sim >= 0.999
|
||||||
e.hits += 1
|
level = "exact" if is_exact else "semantic"
|
||||||
if e.hits >= self.promote_frequency:
|
self.hits[level] += 1
|
||||||
self._exact[query] = e
|
self._bump_semantic(best_idx, best_q, best_entry)
|
||||||
self._semantic.pop(i)
|
return (level, best_entry.result)
|
||||||
self._sem_vecs.pop(query, None)
|
|
||||||
break
|
self.misses += 1
|
||||||
|
return None
|
||||||
# ---- 写入 ----
|
|
||||||
def put(self, query: str, result: Dict[str, Any]):
|
def _bump_semantic(self, idx: int, query: str, entry: CacheEntry):
|
||||||
if query in self._exact:
|
"""语义命中:累计命中次数,达到阈值提升为精确缓存(O(1),无需二次查找)。"""
|
||||||
return
|
entry.hits += 1
|
||||||
entry = CacheEntry(result=result)
|
if entry.hits >= self.promote_frequency:
|
||||||
if self.semantic_enabled:
|
self._exact[query] = entry
|
||||||
if len(self._semantic) >= self.max_semantic:
|
self._semantic.pop(idx)
|
||||||
old_q, _ = self._semantic.pop(0)
|
self._sem_vecs.pop(query, None)
|
||||||
self._sem_vecs.pop(old_q, None)
|
self._sem_norms.pop(query, None)
|
||||||
self._semantic.append((query, entry))
|
|
||||||
self._sem_vecs[query] = _tf_vector(_ngrams(query))
|
# ---- 写入 ----
|
||||||
else:
|
def put(self, query: str, result: Dict[str, Any]):
|
||||||
self._exact[query] = entry
|
if query in self._exact:
|
||||||
if len(self._exact) > self.max_exact:
|
return
|
||||||
self._exact.pop(next(iter(self._exact)))
|
entry = CacheEntry(result=result)
|
||||||
|
if self.semantic_enabled:
|
||||||
# ---- 统计 ----
|
if len(self._semantic) >= self.max_semantic:
|
||||||
def stats(self) -> Dict[str, Any]:
|
old_q, _ = self._semantic.pop(0)
|
||||||
total = self.hits["exact"] + self.hits["semantic"] + self.misses
|
self._sem_vecs.pop(old_q, None)
|
||||||
return {
|
self._sem_norms.pop(old_q, None)
|
||||||
"exact_hits": self.hits["exact"],
|
self._semantic.append((query, entry))
|
||||||
"semantic_hits": self.hits["semantic"],
|
vec = _tf_vector(_ngrams(query))
|
||||||
"misses": self.misses,
|
self._sem_vecs[query] = vec
|
||||||
"hit_rate": round((self.hits["exact"] + self.hits["semantic"]) / total, 4) if total else 0.0,
|
self._sem_norms[query] = _norm(vec)
|
||||||
"exact_size": len(self._exact),
|
else:
|
||||||
"semantic_size": len(self._semantic),
|
self._exact[query] = entry
|
||||||
}
|
if len(self._exact) > self.max_exact:
|
||||||
|
self._exact.pop(next(iter(self._exact)))
|
||||||
def clear(self):
|
|
||||||
self._exact.clear()
|
# ---- 统计 ----
|
||||||
self._semantic.clear()
|
def stats(self) -> Dict[str, Any]:
|
||||||
self._sem_vecs.clear()
|
total = self.hits["exact"] + self.hits["semantic"] + self.misses
|
||||||
self.hits = {"exact": 0, "semantic": 0}
|
return {
|
||||||
self.misses = 0␍
|
"exact_hits": self.hits["exact"],
|
||||||
|
"semantic_hits": self.hits["semantic"],
|
||||||
|
"misses": self.misses,
|
||||||
|
"hit_rate": round((self.hits["exact"] + self.hits["semantic"]) / total, 4) if total else 0.0,
|
||||||
|
"exact_size": len(self._exact),
|
||||||
|
"semantic_size": len(self._semantic),
|
||||||
|
}
|
||||||
|
|
||||||
|
def clear(self):
|
||||||
|
self._exact.clear()
|
||||||
|
self._semantic.clear()
|
||||||
|
self._sem_vecs.clear()
|
||||||
|
self._sem_norms.clear()
|
||||||
|
self.hits = {"exact": 0, "semantic": 0}
|
||||||
|
self.misses = 0
|
||||||
|
|||||||
@@ -128,7 +128,8 @@ class RuleClassifier(BaseClassifier):
|
|||||||
matched_rules=[],
|
matched_rules=[],
|
||||||
)
|
)
|
||||||
|
|
||||||
best_domain = max(raw, key=raw.get)
|
# 同分决胜:按领域名字典序,保证与规则表排列顺序无关的确定性
|
||||||
|
best_domain = max(sorted(raw), key=lambda d: raw[d])
|
||||||
best_score = raw[best_domain]
|
best_score = raw[best_domain]
|
||||||
confidence = 1.0 - math.exp(-best_score)
|
confidence = 1.0 - math.exp(-best_score)
|
||||||
|
|
||||||
@@ -138,7 +139,7 @@ class RuleClassifier(BaseClassifier):
|
|||||||
|
|
||||||
# 与次高分的差距影响置信度(区分度)
|
# 与次高分的差距影响置信度(区分度)
|
||||||
if len(raw) > 1:
|
if len(raw) > 1:
|
||||||
second = sorted(raw.values(), reverse=True)[1]
|
second = max(v for d, v in raw.items() if d != best_domain)
|
||||||
if second > 0.7 * best_score:
|
if second > 0.7 * best_score:
|
||||||
confidence *= 0.85
|
confidence *= 0.85
|
||||||
|
|
||||||
|
|||||||
+221
-217
@@ -1,217 +1,221 @@
|
|||||||
"""主路由器:协调 缓存 -> 分类 -> 专家 -> Judge -> 大模型回退 的完整链路。
|
"""主路由器:协调 缓存 -> 分类 -> 专家 -> Judge -> 大模型回退 的完整链路。
|
||||||
|
|
||||||
流程(对齐实现方案):
|
流程(对齐实现方案):
|
||||||
1. 检查缓存(L1 精确 / L2 语义)
|
1. 检查缓存(L1 精确 / L2 语义)
|
||||||
2. 低置信度查询直接走大模型(should_fallback)
|
2. 低置信度查询直接走大模型(should_fallback)
|
||||||
3. 分类器输出领域 + 难度
|
3. 分类器输出领域 + 难度
|
||||||
4. 选择专家模型生成
|
4. 选择专家模型生成
|
||||||
5. Judge 评估质量
|
5. Judge 评估质量
|
||||||
6. 质量不达标 -> 升级大模型
|
6. 质量不达标 -> 升级大模型
|
||||||
7. 记录指标、写缓存、返回结果
|
7. 记录指标、写缓存、返回结果
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
from .cache import RouterCache
|
from .cache import RouterCache
|
||||||
from .classifier import BaseClassifier, build_classifier
|
from .classifier import BaseClassifier, build_classifier
|
||||||
from .config import load_config
|
from .config import load_config
|
||||||
from .experts import Expert, build_expert_pool
|
from .experts import Expert, build_expert_pool
|
||||||
from .fallback import FallbackProvider, build_fallback
|
from .fallback import FallbackProvider, build_fallback
|
||||||
from .judge import BaseJudge, build_judge
|
from .judge import BaseJudge, build_judge
|
||||||
from .models import Classification, ExpertResponse, RouterResult, now_ms
|
from .models import Classification, ExpertResponse, RouterResult, now_ms
|
||||||
from .stats import Stats
|
from .stats import Stats
|
||||||
|
|
||||||
|
|
||||||
class Router:
|
class Router:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
classifier: BaseClassifier,
|
classifier: BaseClassifier,
|
||||||
experts: Dict[str, Expert],
|
experts: Dict[str, Expert],
|
||||||
judge: BaseJudge,
|
judge: BaseJudge,
|
||||||
fallback: FallbackProvider,
|
fallback: FallbackProvider,
|
||||||
cache: Optional[RouterCache] = None,
|
cache: Optional[RouterCache] = None,
|
||||||
stats: Optional[Stats] = None,
|
stats: Optional[Stats] = None,
|
||||||
config: Optional[Dict[str, Any]] = None,
|
config: Optional[Dict[str, Any]] = None,
|
||||||
):
|
):
|
||||||
self.classifier = classifier
|
self.classifier = classifier
|
||||||
self.experts = experts
|
self.experts = experts
|
||||||
self.judge = judge
|
self.judge = judge
|
||||||
self.fallback = fallback
|
self.fallback = fallback
|
||||||
self.cache = cache or RouterCache()
|
self.cache = cache or RouterCache()
|
||||||
self.stats = stats or Stats()
|
self.stats = stats or Stats()
|
||||||
cfg = config or {}
|
cfg = config or {}
|
||||||
rcfg = cfg.get("router", {})
|
rcfg = cfg.get("router", {})
|
||||||
self.low_confidence_threshold = rcfg.get("low_confidence_threshold", 0.60)
|
self.low_confidence_threshold = rcfg.get("low_confidence_threshold", 0.60)
|
||||||
self.judge_fallback_threshold = rcfg.get("judge_fallback_threshold", 0.70)
|
self.judge_fallback_threshold = rcfg.get("judge_fallback_threshold", 0.70)
|
||||||
self.cache_enabled = cfg.get("cache", {}).get("enabled", True)
|
self.cache_enabled = cfg.get("cache", {}).get("enabled", True)
|
||||||
|
|
||||||
# ---------------------------------------------------------------
|
# ---------------------------------------------------------------
|
||||||
async def route(self, query: str) -> RouterResult:
|
async def route(self, query: str) -> RouterResult:
|
||||||
start = now_ms()
|
start = now_ms()
|
||||||
route: list = []
|
route: list = []
|
||||||
|
|
||||||
# ---- Step 1: 缓存 ----
|
# ---- Step 1: 缓存 ----
|
||||||
if self.cache_enabled:
|
if self.cache_enabled:
|
||||||
hit = self.cache.get(query)
|
hit = self.cache.get(query)
|
||||||
if hit is not None:
|
if hit is not None:
|
||||||
level, cached = hit
|
level, cached = hit
|
||||||
latency = now_ms() - start
|
latency = now_ms() - start
|
||||||
result = RouterResult(
|
result = RouterResult(
|
||||||
query=query,
|
query=query,
|
||||||
response=cached.get("response", ""),
|
response=cached.get("response", ""),
|
||||||
domain=cached.get("domain", "general"),
|
domain=cached.get("domain", "general"),
|
||||||
difficulty=cached.get("difficulty", "medium"),
|
difficulty=cached.get("difficulty", "medium"),
|
||||||
confidence=cached.get("confidence", 0.0),
|
confidence=cached.get("confidence", 0.0),
|
||||||
upgraded=False,
|
upgraded=False,
|
||||||
quality_score=cached.get("quality_score", 0.0),
|
quality_score=cached.get("quality_score", 0.0),
|
||||||
model_used=cached.get("model_used", ""),
|
model_used=cached.get("model_used", ""),
|
||||||
route=["cache:" + level],
|
route=["cache:" + level],
|
||||||
latency_ms=latency,
|
latency_ms=latency,
|
||||||
cache_hit=True,
|
cache_hit=True,
|
||||||
cache_level=level,
|
cache_level=level,
|
||||||
cost_est=0.0,
|
cost_est=0.0,
|
||||||
)
|
)
|
||||||
self.stats.record(latency, result.domain, result.difficulty, False, True, level, 0.0, result.model_used)
|
self.stats.record(latency, result.domain, result.difficulty, False, True, level, 0.0, result.model_used)
|
||||||
return result
|
return result
|
||||||
route.append("cache:miss")
|
route.append("cache:miss")
|
||||||
|
|
||||||
# ---- Step 2: 分类 ----
|
# ---- Step 2: 分类 ----
|
||||||
classification = self.classifier.classify(query)
|
classification = self.classifier.classify(query)
|
||||||
route.append(f"classify:{classification.domain}@{classification.confidence:.2f}/{classification.difficulty}")
|
route.append(f"classify:{classification.domain}@{classification.confidence:.2f}/{classification.difficulty}")
|
||||||
|
|
||||||
# 低置信度 -> 直接走大模型
|
# 低置信度 -> 直接走大模型
|
||||||
if self.classifier.should_fallback(classification, self.low_confidence_threshold):
|
if self.classifier.should_fallback(classification, self.low_confidence_threshold):
|
||||||
route.append("direct_fallback")
|
route.append("direct_fallback")
|
||||||
fb = await self._call_fallback(query)
|
return await self._fallback_result(query, classification, route, start)
|
||||||
latency = now_ms() - start
|
|
||||||
result = self._finalize(query, classification, fb, quality_score=0.0,
|
# ---- Step 3: 选择专家 ----
|
||||||
upgraded=True, route=route, latency_ms=latency,
|
domain = classification.domain
|
||||||
model_used=fb.model_used, cost_est=fb.cost_est)
|
expert = self.experts.get(domain)
|
||||||
self._record(result, latency)
|
if expert is None:
|
||||||
return result
|
expert = self.experts.get("general")
|
||||||
|
route.append("expert:fallback-to-general")
|
||||||
# ---- Step 3: 选择专家 ----
|
else:
|
||||||
domain = classification.domain
|
route.append(f"expert:{expert.name}")
|
||||||
expert = self.experts.get(domain)
|
|
||||||
if expert is None:
|
# ---- Step 4: 生成 ----
|
||||||
expert = self.experts.get("general")
|
try:
|
||||||
route.append("expert:fallback-to-general")
|
expert_resp = await expert.generate(query, classification.difficulty)
|
||||||
else:
|
except Exception as e:
|
||||||
route.append(f"expert:{expert.name}")
|
self.stats.record_error()
|
||||||
|
route.append(f"expert_error:{type(e).__name__}")
|
||||||
# ---- Step 4: 生成 ----
|
return await self._fallback_result(query, classification, route, start,
|
||||||
try:
|
error=str(e))
|
||||||
expert_resp = await expert.generate(query, classification.difficulty)
|
|
||||||
except Exception as e:
|
# ---- Step 5: Judge 评估 ----
|
||||||
self.stats.record_error()
|
try:
|
||||||
route.append(f"expert_error:{type(e).__name__}")
|
evaluation = await self.judge.evaluate(query, expert_resp.text, domain)
|
||||||
fb = await self._call_fallback(query)
|
except Exception:
|
||||||
latency = now_ms() - start
|
evaluation = None
|
||||||
result = self._finalize(query, classification, fb, quality_score=0.0,
|
route.append("judge_error")
|
||||||
upgraded=True, route=route, latency_ms=latency,
|
|
||||||
model_used=fb.model_used, cost_est=fb.cost_est,
|
quality_score = evaluation.overall_score if evaluation else 0.0
|
||||||
error=str(e))
|
route.append(f"judge:{quality_score:.2f}")
|
||||||
self._record(result, latency)
|
|
||||||
return result
|
upgraded = False
|
||||||
|
final_resp = expert_resp
|
||||||
# ---- Step 5: Judge 评估 ----
|
if evaluation is not None and evaluation.needs_fallback:
|
||||||
try:
|
route.append("upgrade")
|
||||||
evaluation = await self.judge.evaluate(query, expert_resp.text, domain)
|
final_resp = await self._call_fallback(query)
|
||||||
except Exception:
|
upgraded = True
|
||||||
evaluation = None
|
|
||||||
route.append("judge_error")
|
latency = now_ms() - start
|
||||||
|
result = self._finalize(query, classification, final_resp, quality_score=quality_score,
|
||||||
quality_score = evaluation.overall_score if evaluation else 0.0
|
upgraded=upgraded, route=route, latency_ms=latency,
|
||||||
route.append(f"judge:{quality_score:.2f}")
|
model_used=final_resp.model_used, cost_est=final_resp.cost_est)
|
||||||
|
self._record(result, latency)
|
||||||
upgraded = False
|
|
||||||
final_resp = expert_resp
|
# 未升级的结果写缓存
|
||||||
if evaluation is not None and evaluation.needs_fallback:
|
if self.cache_enabled and not upgraded and result.response:
|
||||||
route.append("upgrade")
|
self.cache.put(query, result.to_dict())
|
||||||
final_resp = await self._call_fallback(query)
|
|
||||||
upgraded = True
|
return result
|
||||||
|
|
||||||
latency = now_ms() - start
|
# ---------------------------------------------------------------
|
||||||
result = self._finalize(query, classification, final_resp, quality_score=quality_score,
|
async def _fallback_result(self, query: str, classification: Classification,
|
||||||
upgraded=upgraded, route=route, latency_ms=latency,
|
route: list, start: float,
|
||||||
model_used=final_resp.model_used, cost_est=final_resp.cost_est)
|
error: Optional[str] = None) -> RouterResult:
|
||||||
self._record(result, latency)
|
"""兜底路径的公共收尾:调用大模型回退 -> finalize -> 记录指标。
|
||||||
|
|
||||||
# 未升级的结果写缓存
|
低置信度直连、专家异常两条路径共用,避免收尾逻辑三处重复。
|
||||||
if self.cache_enabled and not upgraded and result.response:
|
"""
|
||||||
self.cache.put(query, result.to_dict())
|
fb = await self._call_fallback(query)
|
||||||
|
latency = now_ms() - start
|
||||||
return result
|
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,
|
||||||
async def _call_fallback(self, query: str) -> ExpertResponse:
|
error=error)
|
||||||
try:
|
self._record(result, latency)
|
||||||
return await self.fallback.generate(query)
|
return result
|
||||||
except Exception as e:
|
|
||||||
# 回退也失败:返回错误占位响应
|
async def _call_fallback(self, query: str) -> ExpertResponse:
|
||||||
return ExpertResponse(
|
try:
|
||||||
text=f"[系统错误] 专家与大模型回退均失败:{type(e).__name__}: {e}",
|
return await self.fallback.generate(query)
|
||||||
model_used=f"error:{self.fallback.name}",
|
except Exception as e:
|
||||||
cost_est=0.0,
|
# 回退也失败:返回错误占位响应
|
||||||
)
|
return ExpertResponse(
|
||||||
|
text=f"[系统错误] 专家与大模型回退均失败:{type(e).__name__}: {e}",
|
||||||
@staticmethod
|
model_used=f"error:{self.fallback.name}",
|
||||||
def _finalize(query: str, classification: Classification, resp: ExpertResponse,
|
cost_est=0.0,
|
||||||
quality_score: float, upgraded: bool, route: list,
|
)
|
||||||
latency_ms: float, model_used: str, cost_est: float,
|
|
||||||
error: Optional[str] = None) -> RouterResult:
|
@staticmethod
|
||||||
return RouterResult(
|
def _finalize(query: str, classification: Classification, resp: ExpertResponse,
|
||||||
query=query,
|
quality_score: float, upgraded: bool, route: list,
|
||||||
response=resp.text,
|
latency_ms: float, model_used: str, cost_est: float,
|
||||||
domain=classification.domain,
|
error: Optional[str] = None) -> RouterResult:
|
||||||
difficulty=classification.difficulty,
|
return RouterResult(
|
||||||
confidence=classification.confidence,
|
query=query,
|
||||||
upgraded=upgraded,
|
response=resp.text,
|
||||||
quality_score=quality_score,
|
domain=classification.domain,
|
||||||
model_used=model_used,
|
difficulty=classification.difficulty,
|
||||||
route=route,
|
confidence=classification.confidence,
|
||||||
latency_ms=latency_ms,
|
upgraded=upgraded,
|
||||||
cache_hit=False,
|
quality_score=quality_score,
|
||||||
cost_est=cost_est,
|
model_used=model_used,
|
||||||
error=error,
|
route=route,
|
||||||
)
|
latency_ms=latency_ms,
|
||||||
|
cache_hit=False,
|
||||||
def _record(self, result: RouterResult, latency_ms: float):
|
cost_est=cost_est,
|
||||||
self.stats.record(
|
error=error,
|
||||||
latency_ms,
|
)
|
||||||
result.domain,
|
|
||||||
result.difficulty,
|
def _record(self, result: RouterResult, latency_ms: float):
|
||||||
result.upgraded,
|
self.stats.record(
|
||||||
result.cache_hit,
|
latency_ms,
|
||||||
result.cache_level,
|
result.domain,
|
||||||
result.cost_est,
|
result.difficulty,
|
||||||
result.model_used,
|
result.upgraded,
|
||||||
)
|
result.cache_hit,
|
||||||
|
result.cache_level,
|
||||||
# ---------------------------------------------------------------
|
result.cost_est,
|
||||||
def health(self) -> Dict[str, Any]:
|
result.model_used,
|
||||||
return {
|
)
|
||||||
"status": "ok",
|
|
||||||
"domains": list(self.experts.keys()),
|
# ---------------------------------------------------------------
|
||||||
"classifier": type(self.classifier).__name__,
|
def health(self) -> Dict[str, Any]:
|
||||||
"judge": type(self.judge).__name__,
|
return {
|
||||||
"fallback": type(self.fallback).__name__,
|
"status": "ok",
|
||||||
}
|
"domains": list(self.experts.keys()),
|
||||||
|
"classifier": type(self.classifier).__name__,
|
||||||
|
"judge": type(self.judge).__name__,
|
||||||
def build_router(config_path: Optional[str] = None) -> Router:
|
"fallback": type(self.fallback).__name__,
|
||||||
"""从配置构建完整 Router(默认 mock 全链路,零依赖可跑)。"""
|
}
|
||||||
config = load_config(config_path)
|
|
||||||
classifier = build_classifier(config.get("classifier", {}))
|
|
||||||
experts = build_expert_pool(config.get("experts", {}), config.get("domains", []))
|
def build_router(config_path: Optional[str] = None) -> Router:
|
||||||
judge = build_judge(config.get("judge", {}), config.get("router", {}).get("judge_fallback_threshold", 0.70))
|
"""从配置构建完整 Router(默认 mock 全链路,零依赖可跑)。"""
|
||||||
fallback = build_fallback(config.get("fallback", {}))
|
config = load_config(config_path)
|
||||||
cache_cfg = config.get("cache", {})
|
classifier = build_classifier(config.get("classifier", {}))
|
||||||
cache = RouterCache(
|
experts = build_expert_pool(config.get("experts", {}), config.get("domains", []))
|
||||||
semantic_enabled=cache_cfg.get("semantic_enabled", True),
|
judge = build_judge(config.get("judge", {}), config.get("router", {}).get("judge_fallback_threshold", 0.70))
|
||||||
similarity_threshold=cache_cfg.get("similarity_threshold", 0.88),
|
fallback = build_fallback(config.get("fallback", {}))
|
||||||
promote_frequency=cache_cfg.get("promote_frequency", 5),
|
cache_cfg = config.get("cache", {})
|
||||||
)
|
cache = RouterCache(
|
||||||
stats = Stats()
|
semantic_enabled=cache_cfg.get("semantic_enabled", True),
|
||||||
return Router(classifier, experts, judge, fallback, cache, stats, config)␍
|
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)
|
||||||
|
|||||||
@@ -1,6 +1,42 @@
|
|||||||
from router_system.cache import RouterCache
|
from router_system.cache import RouterCache
|
||||||
|
|
||||||
|
|
||||||
|
def test_semantic_lookup_after_many_entries():
|
||||||
|
"""多条目下语义命中正确(范数预计算 + 单遍扫描的回归)。"""
|
||||||
|
c = RouterCache(similarity_threshold=0.5)
|
||||||
|
for i in range(50):
|
||||||
|
c.put(f"完全不相关的查询主题编号{i}关于烹饪的意见", {"response": f"r{i}"})
|
||||||
|
c.put("用 Python 实现快速排序函数", {"response": "code-answer"})
|
||||||
|
level, got = c.get("用 Python 实现快速排序的函数写法") # 相似但不完全相同
|
||||||
|
assert level in ("semantic", "exact")
|
||||||
|
assert got["response"] == "code-answer"
|
||||||
|
|
||||||
|
|
||||||
|
def test_promotion_clears_semantic_state():
|
||||||
|
"""提升为精确缓存后,语义列表与范数索引无残留。"""
|
||||||
|
c = RouterCache(promote_frequency=2)
|
||||||
|
c.put("查询甲", {"response": "a"})
|
||||||
|
first = c.get("查询甲") # 相似度=1.0 计 exact,hits 达阈值即提升
|
||||||
|
assert first is not None and first[0] == "exact"
|
||||||
|
second = c.get("查询甲")
|
||||||
|
assert second is not None and second[0] == "exact"
|
||||||
|
assert c.stats()["exact_size"] == 1
|
||||||
|
assert c.stats()["semantic_size"] == 0
|
||||||
|
assert len(c._sem_norms) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_semantic_eviction_clears_norms():
|
||||||
|
"""语义缓存满员淘汰最旧条目时,向量与范数索引同步清理。"""
|
||||||
|
c = RouterCache(max_semantic=2)
|
||||||
|
c.put("查询一", {"response": "1"})
|
||||||
|
c.put("查询二", {"response": "2"})
|
||||||
|
c.put("查询三", {"response": "3"}) # 淘汰查询一
|
||||||
|
assert len(c._semantic) == 2
|
||||||
|
assert len(c._sem_vecs) == 2
|
||||||
|
assert len(c._sem_norms) == 2
|
||||||
|
assert c.get("查询一") is None
|
||||||
|
|
||||||
|
|
||||||
def test_exact_hit():
|
def test_exact_hit():
|
||||||
c = RouterCache()
|
c = RouterCache()
|
||||||
result = {"response": "hello", "domain": "general"}
|
result = {"response": "hello", "domain": "general"}
|
||||||
|
|||||||
+63
-44
@@ -1,44 +1,63 @@
|
|||||||
"""分类器单元测试。"""
|
"""分类器单元测试。"""
|
||||||
from router_system.classifier import RuleClassifier
|
from router_system.classifier import RuleClassifier
|
||||||
|
|
||||||
|
|
||||||
def test_code_classification():
|
def test_tie_break_is_deterministic():
|
||||||
clf = RuleClassifier()
|
"""同分决胜:按领域名字典序,与规则表排列顺序无关。"""
|
||||||
r = clf.classify("用 Python 写一个快速排序函数")
|
clf = RuleClassifier()
|
||||||
assert r.domain == "code"
|
clf.rules = {"zeta": [("x", 1.0)], "alpha": [("x", 1.0)]}
|
||||||
assert r.confidence > 0.7
|
r = clf.classify("x")
|
||||||
|
assert r.domain == "alpha"
|
||||||
|
|
||||||
def test_math_classification():
|
|
||||||
clf = RuleClassifier()
|
def test_distinctiveness_penalty():
|
||||||
r = clf.classify("求解方程 x^2 - 5x + 6 = 0")
|
"""次高分占比高(语义含混)时置信度被压低;单一领域命中不受影响。"""
|
||||||
assert r.domain == "math"
|
clf = RuleClassifier()
|
||||||
assert r.confidence > 0.7
|
clf.rules = {"a": [("kw", 1.0)], "b": [("kw", 0.9)]}
|
||||||
|
r_ambiguous = clf.classify("kw")
|
||||||
|
clf_clear = RuleClassifier()
|
||||||
def test_legal_classification():
|
clf_clear.rules = {"a": [("kw", 1.0)], "b": [("other", 0.1)]}
|
||||||
clf = RuleClassifier()
|
r_clear = clf_clear.classify("kw")
|
||||||
r = clf.classify("劳动合同到期不续签需要支付经济补偿吗")
|
assert r_clear.confidence > r_ambiguous.confidence
|
||||||
assert r.domain == "legal"
|
|
||||||
|
|
||||||
|
def test_code_classification():
|
||||||
def test_medical_classification():
|
clf = RuleClassifier()
|
||||||
clf = RuleClassifier()
|
r = clf.classify("用 Python 写一个快速排序函数")
|
||||||
r = clf.classify("高血压患者日常饮食需要注意什么")
|
assert r.domain == "code"
|
||||||
assert r.domain == "medical"
|
assert r.confidence > 0.7
|
||||||
|
|
||||||
|
|
||||||
def test_general_low_confidence():
|
def test_math_classification():
|
||||||
clf = RuleClassifier()
|
clf = RuleClassifier()
|
||||||
r = clf.classify("今天天气怎么样")
|
r = clf.classify("求解方程 x^2 - 5x + 6 = 0")
|
||||||
# 未命中任何领域 -> 低置信度,触发 should_fallback
|
assert r.domain == "math"
|
||||||
assert r.domain == "general"
|
assert r.confidence > 0.7
|
||||||
assert clf.should_fallback(r, 0.6) is True
|
|
||||||
|
|
||||||
|
def test_legal_classification():
|
||||||
def test_difficulty_estimation():
|
clf = RuleClassifier()
|
||||||
clf = RuleClassifier()
|
r = clf.classify("劳动合同到期不续签需要支付经济补偿吗")
|
||||||
easy = clf.classify("1 + 1 = ?")
|
assert r.domain == "legal"
|
||||||
hard = clf.classify("证明费马大定理并推导其推论,给出详细步骤")
|
|
||||||
assert hard.difficulty in ("medium", "hard")
|
|
||||||
assert easy.difficulty == "easy"␍
|
def test_medical_classification():
|
||||||
|
clf = RuleClassifier()
|
||||||
|
r = clf.classify("高血压患者日常饮食需要注意什么")
|
||||||
|
assert r.domain == "medical"
|
||||||
|
|
||||||
|
|
||||||
|
def test_general_low_confidence():
|
||||||
|
clf = RuleClassifier()
|
||||||
|
r = clf.classify("今天天气怎么样")
|
||||||
|
# 未命中任何领域 -> 低置信度,触发 should_fallback
|
||||||
|
assert r.domain == "general"
|
||||||
|
assert clf.should_fallback(r, 0.6) is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_difficulty_estimation():
|
||||||
|
clf = RuleClassifier()
|
||||||
|
easy = clf.classify("1 + 1 = ?")
|
||||||
|
hard = clf.classify("证明费马大定理并推导其推论,给出详细步骤")
|
||||||
|
assert hard.difficulty in ("medium", "hard")
|
||||||
|
assert easy.difficulty == "easy"
|
||||||
|
|||||||
Reference in New Issue
Block a user