feat: 多专业小模型+路由模型系统 MVP(mock 全链路 + FastAPI 网关 + 论文调研)
This commit is contained in:
@@ -0,0 +1,141 @@
|
||||
"""两阶段路由缓存(对齐实现方案):
|
||||
- L1 精确缓存:完全相同的查询 -> 直接命中
|
||||
- L2 语义缓存:字符 n-gram 余弦相似度(零依赖)-> 相似查询命中
|
||||
- 命中 N 次(promote_frequency)后提升为精确缓存
|
||||
|
||||
说明:语义缓存中的"完全相同查询"(相似度=1.0)直接计为 exact 命中;
|
||||
高频语义命中会提升为 O(1) 的精确缓存条目。
|
||||
|
||||
只缓存"未升级"的结果(升级路径每次都走大模型,不缓存,避免陈旧)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
|
||||
@dataclass
|
||||
class CacheEntry:
|
||||
result: Dict[str, Any]
|
||||
hits: int = 1
|
||||
|
||||
|
||||
def _ngrams(text: str, n: int = 3) -> List[str]:
|
||||
"""字符 n-gram(去空白、小写),用于轻量语义相似度。"""
|
||||
cleaned = re.sub(r"\s+", "", text.lower())
|
||||
if len(cleaned) < n:
|
||||
return [cleaned]
|
||||
return [cleaned[i:i + n] for i in range(len(cleaned) - n + 1)]
|
||||
|
||||
|
||||
def _cosine(vec_a: Dict[str, float], vec_b: Dict[str, float]) -> float:
|
||||
if not vec_a or not vec_b:
|
||||
return 0.0
|
||||
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
|
||||
nb = sum(v * v for v in vec_b.values()) ** 0.5
|
||||
if na == 0 or nb == 0:
|
||||
return 0.0
|
||||
return dot / (na * nb)
|
||||
|
||||
|
||||
def _tf_vector(grams: List[str]) -> Dict[str, float]:
|
||||
vec: Dict[str, float] = {}
|
||||
for g in grams:
|
||||
vec[g] = vec.get(g, 0.0) + 1.0
|
||||
return vec
|
||||
|
||||
|
||||
class RouterCache:
|
||||
"""L1 精确缓存 + L2 语义缓存。"""
|
||||
|
||||
def __init__(self, semantic_enabled: bool = True, similarity_threshold: float = 0.88,
|
||||
promote_frequency: int = 5, max_exact: int = 10000, max_semantic: int = 5000):
|
||||
self.semantic_enabled = semantic_enabled
|
||||
self.similarity_threshold = similarity_threshold
|
||||
self.promote_frequency = promote_frequency
|
||||
self.max_exact = max_exact
|
||||
self.max_semantic = max_semantic
|
||||
self._exact: Dict[str, CacheEntry] = {}
|
||||
self._semantic: List[Tuple[str, CacheEntry]] = [] # (query, entry)
|
||||
self._sem_vecs: Dict[str, Dict[str, float]] = {}
|
||||
self.hits = {"exact": 0, "semantic": 0}
|
||||
self.misses = 0
|
||||
|
||||
# ---- 查询 ----
|
||||
def get(self, query: str) -> Optional[Tuple[Optional[str], Dict[str, Any]]]:
|
||||
"""返回 (level, result);未命中返回 None。level: 'exact' | 'semantic'"""
|
||||
entry = self._exact.get(query)
|
||||
if entry is not None:
|
||||
self.hits["exact"] += 1
|
||||
return ("exact", entry.result)
|
||||
|
||||
if self.semantic_enabled:
|
||||
q_vec = _tf_vector(_ngrams(query))
|
||||
best_sim = 0.0
|
||||
best_query: Optional[str] = None
|
||||
best_result: Optional[Dict[str, Any]] = None
|
||||
for q, e in self._semantic:
|
||||
sim = _cosine(q_vec, self._sem_vecs.get(q, {}))
|
||||
if sim > best_sim:
|
||||
best_sim = sim
|
||||
best_query = q
|
||||
best_result = e.result
|
||||
if best_query is not None and best_sim >= self.similarity_threshold:
|
||||
# 完全相同查询(相似度=1.0)计为 exact 命中
|
||||
is_exact = best_sim >= 0.999
|
||||
level = "exact" if is_exact else "semantic"
|
||||
self.hits[level] += 1
|
||||
self._semantic_hit(best_query)
|
||||
return (level, best_result)
|
||||
|
||||
self.misses += 1
|
||||
return None
|
||||
|
||||
def _semantic_hit(self, query: str):
|
||||
"""语义命中:累计命中次数,达到阈值提升为精确缓存。"""
|
||||
for i, (q, e) in enumerate(self._semantic):
|
||||
if q == query:
|
||||
e.hits += 1
|
||||
if e.hits >= self.promote_frequency:
|
||||
self._exact[query] = e
|
||||
self._semantic.pop(i)
|
||||
self._sem_vecs.pop(query, None)
|
||||
break
|
||||
|
||||
# ---- 写入 ----
|
||||
def put(self, query: str, result: Dict[str, Any]):
|
||||
if query in self._exact:
|
||||
return
|
||||
entry = CacheEntry(result=result)
|
||||
if self.semantic_enabled:
|
||||
if len(self._semantic) >= self.max_semantic:
|
||||
old_q, _ = self._semantic.pop(0)
|
||||
self._sem_vecs.pop(old_q, None)
|
||||
self._semantic.append((query, entry))
|
||||
self._sem_vecs[query] = _tf_vector(_ngrams(query))
|
||||
else:
|
||||
self._exact[query] = entry
|
||||
if len(self._exact) > self.max_exact:
|
||||
self._exact.pop(next(iter(self._exact)))
|
||||
|
||||
# ---- 统计 ----
|
||||
def stats(self) -> Dict[str, Any]:
|
||||
total = self.hits["exact"] + self.hits["semantic"] + self.misses
|
||||
return {
|
||||
"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.hits = {"exact": 0, "semantic": 0}
|
||||
self.misses = 0␍
|
||||
Reference in New Issue
Block a user