Files

271 lines
10 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""专家模型池:统一 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 CompletionsDeepSeek / 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 Keyenv: {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