feat: 多专业小模型+路由模型系统 MVP(mock 全链路 + FastAPI 网关 + 论文调研)
This commit is contained in:
@@ -0,0 +1,65 @@
|
||||
"""CLI 演示:构建 mock 全链路路由系统,跑一组样例查询并打印结果。
|
||||
|
||||
用法:
|
||||
python scripts/demo.py [--query "自定义查询"] [--batch]
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
from router_system.router import build_router
|
||||
|
||||
SAMPLE_QUERIES = [
|
||||
"用 Python 写一个快速排序函数,并解释时间复杂度",
|
||||
"求解方程 x^2 - 5x + 6 = 0",
|
||||
"劳动合同里约定离职后两年内不得从事同行业,是否有效?",
|
||||
"高血压患者日常饮食需要注意什么?",
|
||||
"给我总结一下深度学习中注意力机制的优缺点",
|
||||
"为什么天空是蓝色的?",
|
||||
"帮我调试这段代码:def f(x): return x + 1 报 TypeError",
|
||||
"求 ∫ x^2 dx 从 0 到 1 的定积分是多少?",
|
||||
]
|
||||
|
||||
|
||||
async def run_demo(router, queries, verbose: bool = False):
|
||||
for q in queries:
|
||||
r = await router.route(q)
|
||||
print("=" * 72)
|
||||
print(f"Q: {q}")
|
||||
print(f" domain={r.domain} difficulty={r.difficulty} conf={r.confidence:.2f} "
|
||||
f"upgraded={r.upgraded} quality={r.quality_score:.2f} model={r.model_used} "
|
||||
f"latency={r.latency_ms:.1f}ms cache={r.cache_hit}({r.cache_level}) cost=${r.cost_est:.6f}")
|
||||
print(f" route: {' -> '.join(r.route)}")
|
||||
if verbose:
|
||||
print(f" --- response ---\n{r.response[:400]}")
|
||||
|
||||
|
||||
async def main():
|
||||
parser = argparse.ArgumentParser(description="多专家路由系统 CLI 演示")
|
||||
parser.add_argument("--query", type=str, default=None, help="单条查询(覆盖默认样例)")
|
||||
parser.add_argument("--batch", action="store_true", help="批量模式(打印全部响应)")
|
||||
parser.add_argument("--verbose", action="store_true", help="打印响应正文")
|
||||
args = parser.parse_args()
|
||||
|
||||
router = build_router()
|
||||
print("系统组件:", router.health())
|
||||
print()
|
||||
|
||||
if args.query:
|
||||
await run_demo(router, [args.query], verbose=True)
|
||||
else:
|
||||
await run_demo(router, SAMPLE_QUERIES, verbose=args.verbose)
|
||||
|
||||
print()
|
||||
print("=" * 72)
|
||||
print("运行指标:", router.stats.summary())
|
||||
print("缓存统计:", router.cache.stats())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())␍
|
||||
@@ -0,0 +1,80 @@
|
||||
"""迷你评估:对带标注的样例查询评估分类准确率、升级率、成本。
|
||||
|
||||
用法:
|
||||
python scripts/eval.py [--repeat 2] [--config path]
|
||||
--repeat 用于把样例跑 N 遍,验证语义缓存命中与降本效果。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import sys
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
from router_system.router import build_router
|
||||
|
||||
# (query, 期望领域)
|
||||
BENCH = [
|
||||
("用 Python 实现二分查找", "code"),
|
||||
("这段 JavaScript 为什么报错:undefined is not a function", "code"),
|
||||
("帮我优化这个 SQL 查询的索引", "code"),
|
||||
("求解一元二次方程 ax^2+bx+c=0 的求根公式", "math"),
|
||||
("证明勾股定理", "math"),
|
||||
("计算 3x + 5 = 20,x 等于多少", "math"),
|
||||
("劳动合同到期不续签,公司需要支付经济补偿吗", "legal"),
|
||||
("在合同中约定违约金上限 30%,是否合规", "legal"),
|
||||
("专利申请的流程和费用大概是多少", "legal"),
|
||||
("高血压患者可以吃哪些降压药,副作用是什么", "medical"),
|
||||
("感冒发烧 38.5 度,需要吃退烧药吗", "medical"),
|
||||
("糖尿病患者的日常饮食建议", "medical"),
|
||||
("介绍一下 Transformer 架构", "general"),
|
||||
("写一封请假邮件", "general"),
|
||||
("为什么天空是蓝色的", "general"),
|
||||
]
|
||||
|
||||
|
||||
async def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--repeat", type=int, default=2, help="重复轮数(验证缓存)")
|
||||
parser.add_argument("--config", type=str, default=None)
|
||||
args = parser.parse_args()
|
||||
|
||||
router = build_router(args.config)
|
||||
correct = Counter()
|
||||
total = 0
|
||||
upgraded = 0
|
||||
cache_hits = 0
|
||||
|
||||
for round_i in range(args.repeat):
|
||||
for q, expected in BENCH:
|
||||
r = await router.route(q)
|
||||
total += 1
|
||||
if r.domain == expected:
|
||||
correct["total"] += 1
|
||||
else:
|
||||
correct[f"misclass->{r.domain}"] += 1
|
||||
if r.upgraded:
|
||||
upgraded += 1
|
||||
if r.cache_hit:
|
||||
cache_hits += 1
|
||||
|
||||
acc = correct["total"] / total
|
||||
print(f"样例数: {len(BENCH)} x {args.repeat} 轮 = {total} 次请求")
|
||||
print(f"分类准确率: {acc:.1%} ({correct['total']}/{total})")
|
||||
print(f"升级率: {upgraded/total:.1%} ({upgraded}/{total})")
|
||||
print(f"缓存命中率: {cache_hits/total:.1%} ({cache_hits}/{total})")
|
||||
print()
|
||||
print("运行指标:", router.stats.summary())
|
||||
print("缓存统计:", router.cache.stats())
|
||||
print()
|
||||
if acc < 0.8:
|
||||
print("⚠️ 准确率低于 80%,请检查分类规则。")
|
||||
else:
|
||||
print("✅ 分类准确率达标(≥80%)。")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())␍
|
||||
@@ -0,0 +1,60 @@
|
||||
"""启动路由网关服务(后台、无窗口)。
|
||||
|
||||
用法:
|
||||
python scripts/serve.py [--port 8000] [--stop]
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
ROOT = Path(__file__).resolve().parent.parent
|
||||
PID_FILE = ROOT / "_gateway.pid"
|
||||
|
||||
|
||||
def start(port: int):
|
||||
if PID_FILE.exists():
|
||||
old = PID_FILE.read_text().strip()
|
||||
if old:
|
||||
print(f"已有服务运行 (pid={old}),先执行 --stop 再启动。")
|
||||
return
|
||||
out = open(ROOT / "_gateway.out.log", "ab", buffering=0)
|
||||
err = open(ROOT / "_gateway.err.log", "ab", buffering=0)
|
||||
flags = 0x00000008 | 0x08000000 # DETACHED_PROCESS | CREATE_NO_WINDOW
|
||||
proc = subprocess.Popen(
|
||||
[sys.executable, "-m", "uvicorn", "gateway.api:app",
|
||||
"--host", "127.0.0.1", "--port", str(port)],
|
||||
cwd=str(ROOT),
|
||||
stdout=out,
|
||||
stderr=err,
|
||||
creationflags=flags,
|
||||
close_fds=True,
|
||||
)
|
||||
PID_FILE.write_text(str(proc.pid))
|
||||
print(f"gateway 已启动 pid={proc.pid} port={port},日志: _gateway.out.log")
|
||||
|
||||
|
||||
def stop():
|
||||
if not PID_FILE.exists():
|
||||
print("没有运行中的服务。")
|
||||
return
|
||||
pid = int(PID_FILE.read_text().strip())
|
||||
try:
|
||||
import signal
|
||||
os.kill(pid, signal.SIGTERM)
|
||||
print(f"已发送终止信号 pid={pid}")
|
||||
except ProcessLookupError:
|
||||
print(f"进程 {pid} 不存在,清理 pid 文件。")
|
||||
PID_FILE.unlink(missing_ok=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--port", type=int, default=8000)
|
||||
parser.add_argument("--stop", action="store_true", help="停止服务")
|
||||
args = parser.parse_args()
|
||||
if args.stop:
|
||||
stop()
|
||||
else:
|
||||
start(args.port)
|
||||
@@ -0,0 +1,68 @@
|
||||
"""可选:用 transformers 训练 BERT 级意图分类器(需 requirements-ml.txt)。
|
||||
|
||||
这是"正式环境"路线(对照实现方案 3.2 路由器选型表):
|
||||
- 规则分类器(当前默认): 90-93% 准确率,零成本,适合 MVP
|
||||
- 小 LLM / BERT 分类器: 94-97% 准确率,适合正式环境
|
||||
|
||||
用法(示例):
|
||||
python scripts/train_classifier.py --data data/train.jsonl --output models/classifier
|
||||
|
||||
数据格式(每行一个 JSON):
|
||||
{"query": "...", "domain": "code|math|legal|medical|general"}
|
||||
|
||||
注意:本脚本是流水线骨架,需自行准备数据。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data", required=True, help="训练数据 jsonl 路径")
|
||||
parser.add_argument("--output", default="models/classifier", help="输出目录")
|
||||
parser.add_argument("--base", default="Qwen/Qwen3-0.6B", help="基础模型")
|
||||
parser.add_argument("--epochs", type=int, default=3)
|
||||
args = parser.parse_args()
|
||||
|
||||
try:
|
||||
from transformers import (AutoModelForSequenceClassification, AutoTokenizer,
|
||||
TrainingArguments)
|
||||
except ImportError as e:
|
||||
print("需要 ML 依赖:pip install -r requirements-ml.txt")
|
||||
raise SystemExit(1) from e
|
||||
|
||||
# 数据格式转换
|
||||
samples = []
|
||||
with open(args.data, encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line:
|
||||
obj = json.loads(line)
|
||||
samples.append(obj)
|
||||
print(f"加载 {len(samples)} 条训练样本")
|
||||
|
||||
labels = ["code", "math", "legal", "medical", "general"]
|
||||
label2id = {l: i for i, l in enumerate(labels)}
|
||||
|
||||
# 这里仅演示训练流水线;正式训练请使用 datasets 库构建 Dataset。
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.base)
|
||||
model = AutoModelForSequenceClassification.from_pretrained(args.base, num_labels=5)
|
||||
|
||||
print("训练参数示例:")
|
||||
print(TrainingArguments(
|
||||
output_dir=args.output,
|
||||
num_train_epochs=args.epochs,
|
||||
per_device_train_batch_size=8,
|
||||
learning_rate=2e-5,
|
||||
))
|
||||
print("请参考实现方案 3.2 路由器选型表,构建带标注数据集后执行正式训练。")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()␍
|
||||
Reference in New Issue
Block a user