feat: 多专业小模型+路由模型系统 MVP(mock 全链路 + FastAPI 网关 + 论文调研)
This commit is contained in:
@@ -0,0 +1,270 @@
|
||||
"""专家模型池:统一 Expert 接口,支持三种后端。
|
||||
|
||||
- MockExpert :确定性模板输出(零依赖,离线可跑,便于测试与演示)
|
||||
- HFExpert :HuggingFace transformers 真实小模型(可选,需 ML 依赖)
|
||||
- APIExpert :OpenAI 兼容 API(可选,需 API Key,如 DeepSeek)
|
||||
|
||||
成本估计:cost_est 按参数量粗估(美元/百万 token 的近似比例)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from .models import ExpertResponse
|
||||
|
||||
# 按模型规模粗估的相对成本($ / 1M output tokens,近似)
|
||||
MODEL_COST_EST = {
|
||||
"mock": 0.0,
|
||||
"0.5b": 0.02,
|
||||
"1b": 0.05,
|
||||
"1.7b": 0.08,
|
||||
"3b": 0.12,
|
||||
"4b": 0.15,
|
||||
"7b": 0.25,
|
||||
"70b": 2.50,
|
||||
"api": 1.00,
|
||||
}
|
||||
|
||||
|
||||
def _cost_for(model_name: str, default: str = "1b") -> float:
|
||||
mn = model_name.lower()
|
||||
for key in ("0.5b", "1.7b", "3b", "4b", "7b", "70b"):
|
||||
if key in mn:
|
||||
return MODEL_COST_EST[key]
|
||||
if "api" in mn or mn in ("deepseek-chat", "gpt-4o-mini", "claude"):
|
||||
return MODEL_COST_EST["api"]
|
||||
return MODEL_COST_EST.get(default, 0.1)
|
||||
|
||||
|
||||
def extract_content_terms(query: str) -> List[str]:
|
||||
"""抽取查询中的"内容词"(中文词/英文单词),用于 Judge 覆盖度与 Mock 回显。"""
|
||||
q = query.lower()
|
||||
terms: List[str] = []
|
||||
# 英文单词(>=2 字符)
|
||||
for w in re.findall(r"[a-z][a-z0-9_]{1,}", q):
|
||||
if w not in _STOPWORDS_EN and w not in terms:
|
||||
terms.append(w)
|
||||
# 中文:按 2-4 字窗口切分,保留含中文字符的片段
|
||||
cn = re.findall(r"[\u4e00-\u9fff]{2,8}", q)
|
||||
for c in cn:
|
||||
terms.append(c)
|
||||
return terms
|
||||
|
||||
|
||||
_STOPWORDS_EN = {
|
||||
"the", "a", "an", "is", "are", "to", "of", "in", "on", "for", "with", "and",
|
||||
"or", "do", "does", "can", "could", "would", "should", "please", "me", "my",
|
||||
"this", "that", "it", "be", "was", "were", "have", "has", "had", "will",
|
||||
"not", "no", "yes", "i", "you", "he", "she", "we", "they",
|
||||
}
|
||||
|
||||
|
||||
class Expert:
|
||||
name: str = "expert"
|
||||
|
||||
async def generate(self, query: str, difficulty: str) -> ExpertResponse:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class MockExpert(Expert):
|
||||
"""确定性模板专家:零依赖,离线可跑。
|
||||
|
||||
输出会回显查询中的内容词以提高 Judge 覆盖度,并带领域结构,
|
||||
使端到端管线(分类 -> 专家 -> Judge -> 缓存)可被稳定测试与演示。
|
||||
"""
|
||||
|
||||
def __init__(self, name: str, domain: str, model: str = "mock"):
|
||||
self.name = name
|
||||
self.domain = domain
|
||||
self.model = model
|
||||
|
||||
async def generate(self, query: str, difficulty: str) -> ExpertResponse:
|
||||
await asyncio.sleep(0.001) # 模拟极短推理延迟
|
||||
terms = extract_content_terms(query)
|
||||
body = self._template(query, terms, difficulty)
|
||||
# 预估 token 数:中文约 1.5 字符/token,英文约 4 字符/token
|
||||
tokens = max(8, int(len(body) / 2.2))
|
||||
return ExpertResponse(
|
||||
text=body,
|
||||
model_used=self.model,
|
||||
latency_ms=1.0,
|
||||
tokens=tokens,
|
||||
cost_est=_cost_for(self.model) * tokens / 1_000_000,
|
||||
)
|
||||
|
||||
def _template(self, query: str, terms: List[str], difficulty: str) -> str:
|
||||
kw = "、".join(terms[:6]) if terms else "该主题"
|
||||
if self.domain == "code":
|
||||
return (
|
||||
f"(mock 代码专家)针对「{query}」的实现思路如下:\n\n"
|
||||
f"```python\n"
|
||||
f"def solve() -> None:\n"
|
||||
f" # 关键点:{kw}\n"
|
||||
f" # 1. 明确输入输出约束\n"
|
||||
f" # 2. 选择合适数据结构\n"
|
||||
f" # 3. 处理边界条件(空输入、极端值)\n"
|
||||
f" # 4. 补充单元测试\n"
|
||||
f" pass\n"
|
||||
f"```\n\n"
|
||||
f"复杂度:平均 O(n)。请按上述步骤补充具体实现。"
|
||||
)
|
||||
if self.domain == "math":
|
||||
return (
|
||||
f"(mock 数学专家)求解「{query}」的步骤:\n\n"
|
||||
f"1. 明确已知条件与目标:{kw}\n"
|
||||
f"2. 选择合适的方法(代数变形 / 积分 / 归纳等)\n"
|
||||
f"3. 逐步推导并验证中间结果\n"
|
||||
f"4. 检查边界与特殊情况\n\n"
|
||||
f"结论:在标准假设下,结果可化简为闭合形式。完整推导见正式解答。"
|
||||
)
|
||||
if self.domain == "legal":
|
||||
return (
|
||||
f"(mock 法律专家)关于「{query}」的初步法律分析:\n\n"
|
||||
f"相关要点:{kw}\n"
|
||||
f"1. 适用法规:请以现行有效法条为准(建议核对最新修订版)\n"
|
||||
f"2. 合同/合规风险点识别\n"
|
||||
f"3. 责任划分与救济途径\n\n"
|
||||
f"⚠️ 提示:以上为一般性分析,不构成正式法律意见,个案请咨询执业律师。"
|
||||
)
|
||||
if self.domain == "medical":
|
||||
return (
|
||||
f"(mock 医学专家)关于「{query}」的科普性说明:\n\n"
|
||||
f"相关关键词:{kw}\n"
|
||||
f"1. 常见表现与可能原因\n"
|
||||
f"2. 一般处理建议与注意事项\n"
|
||||
f"3. 何时需要就医(警示信号)\n\n"
|
||||
f"⚠️ 提示:内容仅供健康科普,不能替代医生诊断;如有不适请及时就医。"
|
||||
)
|
||||
return (
|
||||
f"(mock 通用专家)关于「{query}」的回答:\n\n"
|
||||
f"核心要点:{kw}\n"
|
||||
f"1. 背景与定义\n"
|
||||
f"2. 主要分类/维度\n"
|
||||
f"3. 实际应用与注意事项\n\n"
|
||||
f"如需更深入的分析,可以补充更多上下文。"
|
||||
)
|
||||
|
||||
|
||||
class HFExpert(Expert):
|
||||
"""可选:HuggingFace 真实小模型(需 requirements-ml.txt)。"""
|
||||
|
||||
def __init__(self, name: str, domain: str, model: str):
|
||||
self.name = name
|
||||
self.domain = domain
|
||||
self.model = model
|
||||
self._loaded = False
|
||||
self._model = None
|
||||
self._tokenizer = None
|
||||
|
||||
def _ensure_loaded(self):
|
||||
if self._loaded:
|
||||
return
|
||||
try:
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
except ImportError as e:
|
||||
raise RuntimeError("HFExpert 需要安装 ML 依赖:pip install -r requirements-ml.txt") from e
|
||||
self._tokenizer = AutoTokenizer.from_pretrained(self.model)
|
||||
self._model = AutoModelForCausalLM.from_pretrained(
|
||||
self.model, device_map="auto", torch_dtype="auto"
|
||||
)
|
||||
self._loaded = True
|
||||
|
||||
async def generate(self, query: str, difficulty: str) -> ExpertResponse:
|
||||
self._ensure_loaded()
|
||||
return await asyncio.to_thread(self._generate_sync, query)
|
||||
|
||||
def _generate_sync(self, query: str) -> ExpertResponse:
|
||||
messages = [{"role": "user", "content": query}]
|
||||
text = self._tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
inputs = self._tokenizer(text, return_tensors="pt").to(self._model.device)
|
||||
outputs = self._model.generate(
|
||||
**inputs,
|
||||
max_new_tokens=512,
|
||||
temperature=0.2,
|
||||
do_sample=True,
|
||||
)
|
||||
body = self._tokenizer.decode(outputs[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True)
|
||||
return ExpertResponse(
|
||||
text=body,
|
||||
model_used=self.model,
|
||||
latency_ms=0.0,
|
||||
tokens=512,
|
||||
cost_est=_cost_for(self.model) * 512 / 1_000_000,
|
||||
)
|
||||
|
||||
|
||||
class APIExpert(Expert):
|
||||
"""可选:OpenAI 兼容 Chat Completions(DeepSeek / OpenAI / 本地 vLLM)。"""
|
||||
|
||||
def __init__(self, name: str, domain: str, model: str, base_url: str, api_key: str):
|
||||
self.name = name
|
||||
self.domain = domain
|
||||
self.model = model
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self._client = None
|
||||
|
||||
def _get_client(self):
|
||||
if self._client is None:
|
||||
import httpx
|
||||
self._client = httpx.AsyncClient(timeout=60.0)
|
||||
return self._client
|
||||
|
||||
async def generate(self, query: str, difficulty: str) -> ExpertResponse:
|
||||
client = self._get_client()
|
||||
resp = await client.post(
|
||||
f"{self.base_url}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {self.api_key}"},
|
||||
json={
|
||||
"model": self.model,
|
||||
"messages": [{"role": "user", "content": query}],
|
||||
"temperature": 0.2,
|
||||
"max_tokens": 1024,
|
||||
},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
body = data["choices"][0]["message"]["content"]
|
||||
usage = data.get("usage", {})
|
||||
tokens = usage.get("completion_tokens", int(len(body) / 2.2))
|
||||
return ExpertResponse(
|
||||
text=body,
|
||||
model_used=self.model,
|
||||
latency_ms=0.0,
|
||||
tokens=tokens,
|
||||
cost_est=_cost_for(self.model) * tokens / 1_000_000,
|
||||
)
|
||||
|
||||
|
||||
def build_expert(domain: str, cfg: Dict) -> Expert:
|
||||
"""根据配置构建领域专家。cfg 为 experts.<domain> 段配置。"""
|
||||
etype = cfg.get("type", "mock")
|
||||
model = cfg.get("model", "mock")
|
||||
name = f"expert-{domain}"
|
||||
if etype == "mock":
|
||||
return MockExpert(name, domain, model)
|
||||
if etype == "hf":
|
||||
return HFExpert(name, domain, model)
|
||||
if etype == "api":
|
||||
base_url = cfg.get("base_url", "https://api.deepseek.com/v1")
|
||||
api_key = cfg.get("api_key") or _env(cfg.get("api_key_env", ""))
|
||||
if not api_key:
|
||||
raise RuntimeError(f"APIExpert({domain}) 缺少 API Key(env: {cfg.get('api_key_env')})")
|
||||
return APIExpert(name, domain, model, base_url, api_key)
|
||||
raise ValueError(f"未知专家后端类型: {etype}(支持 mock | hf | api)")
|
||||
|
||||
|
||||
def _env(name: str) -> Optional[str]:
|
||||
import os
|
||||
return os.environ.get(name) if name else None
|
||||
|
||||
|
||||
def build_expert_pool(experts_cfg: Dict[str, Dict], domains: List[str]) -> Dict[str, Expert]:
|
||||
"""构建完整专家池。"""
|
||||
pool: Dict[str, Expert] = {}
|
||||
for domain in domains:
|
||||
cfg = experts_cfg.get(domain, {"type": "mock", "model": "mock"})
|
||||
pool[domain] = build_expert(domain, cfg)
|
||||
return pool␍
|
||||
Reference in New Issue
Block a user