Files
projectAIpopular/router_system/classifier.py
T
tzt 0dad895aac feat(v1): T-R2 采纳 llmrouter「规则文档即配置」——分类规则外置 config/routes.json
- classifier.load_domain_rules():支持 .json(必有)/.yaml(pyyaml 可选,与
  config.py 同一可选依赖纪律);文件按 domain 整域覆盖内置 DOMAIN_RULES
  (与 v2 知识库'文件按 id 覆盖'同一惯例);缺失/格式非法/条目非法整体
  安全回退内置(llmrouter 失败安全思想,规则文档编辑错误不打垮路由)
- build_classifier:cfg.rules_file 显式指定,未指定时约定路径 config/routes.json
  存在即自动加载(约定优于配置,ROUTES.md 精髓:改文档即改行为,可 review 可版本化)
- 新增 tests/test_rules_external.py 5 项;全量 38 passed(33+5)
2026-09-19 09:24:56 +08:00

271 lines
12 KiB
Python
Raw 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.
"""意图分类器:识别查询领域(code/math/legal/medical/general)与难度。
- RuleClassifier:关键词/正则规则打分,纯标准库,零依赖,可离线运行。
- HuggingFaceClassifier:可选,基于 transformers 的分类模型(需安装 ML 依赖)。
置信度设计:每个领域有一组 (关键词, 权重)。命中权重求和得原始分 s,
confidence = 1 - exp(-s),保证 s=1 -> 0.63s=2 -> 0.86s=3 -> 0.95。
无领域命中(或最高分领域为 general)时置信度低,触发 should_fallback。
规则外置(T-R2,采纳 llmrouter「规则文档即配置」设计):领域规则支持从
config/routes.json(或 .yamlpyyaml 可选)加载,文件按 domain 覆盖内置
DOMAIN_RULES(与 v2 知识库"文件按 id 覆盖"同一惯例);文件缺失/格式非法
整体安全回退内置(llmrouter 失败安全思想)。
"""
from __future__ import annotations
import json
import math
from pathlib import Path
from typing import Dict, List, Optional, Tuple
from .difficulty import estimate_difficulty
from .models import Classification
# ---------------------------------------------------------------
# 领域关键词规则: (关键词, 权重)
# ---------------------------------------------------------------
DOMAIN_RULES: Dict[str, List[Tuple[str, float]]] = {
"code": [
# 中文
("python", 1.2), ("java", 1.2), ("javascript", 1.2), ("typescript", 1.2),
("代码", 1.2), ("编程", 1.2), ("函数", 0.9), ("接口", 0.8), ("报错", 0.9),
("调试", 0.9), ("部署", 0.8), ("算法", 0.8), ("数组", 0.8), ("排序", 0.9),
("正则", 0.8), ("数据库", 0.7), ("sql", 0.8), ("git", 0.7), ("api", 0.7),
("变量", 0.7), ("循环", 0.7), ("递归", 0.8), ("重构", 0.8), ("编译", 0.9),
("测试", 0.6), ("前端", 0.8), ("后端", 0.8), ("爬虫", 0.8), ("脚本", 0.7),
# 英文
("function", 0.9), ("class", 0.8), ("bug", 0.9), ("debug", 0.9),
("compile", 0.9), ("error", 0.6), ("code", 0.7), ("script", 0.7),
("algorithm", 0.8), ("sort", 0.7), ("array", 0.7), ("regex", 0.8),
("import", 0.7), ("loop", 0.7), ("recursion", 0.8), ("refactor", 0.8),
("deploy", 0.8), ("docker", 0.8), ("kubernetes", 0.8),
("async", 0.7), ("flask", 0.7), ("django", 0.7),
("索引", 0.8), ("优化", 0.7), ("查询", 0.6),
],
"math": [
("数学", 1.2), ("方程", 1.0), ("求解", 0.8), ("导数", 1.0), ("积分", 1.0),
("矩阵", 0.9), ("概率", 0.9), ("统计", 0.8), ("证明", 0.8), ("定理", 0.9),
("微积分", 1.1), ("代数", 0.9), ("几何", 0.9), ("不等式", 0.9),
("equation", 1.0), ("derivative", 1.0), ("integral", 1.0), ("calculus", 1.1),
("matrix", 0.9), ("probability", 0.9), ("statistics", 0.8), ("proof", 0.8),
("theorem", 0.9), ("algebra", 0.9), ("geometry", 0.9), ("sqrt", 0.8),
("gcd", 0.8), ("lim", 0.8), ("polynomial", 0.9), ("summation", 0.7),
("math", 0.7), ("解", 0.6), ("计算", 0.8), ("等于", 0.6), ("求值", 0.7), ("函数", 0.6),
],
"legal": [
("法律", 1.2), ("合同", 1.0), ("法条", 1.0), ("合规", 1.0), ("诉讼", 1.0),
("知识产权", 1.1), ("版权", 0.9), ("专利", 0.9), ("违约", 0.9), ("赔偿", 0.8),
("仲裁", 0.9), ("劳动法", 1.0), ("刑法", 1.0), ("民法典", 1.0),
("法规", 0.8), ("条款", 0.7), ("律师", 0.8), ("起诉", 0.9), ("判决", 0.9),
("law", 1.0), ("legal", 1.1), ("contract", 1.0), ("compliance", 1.0),
("litigation", 1.0), ("copyright", 0.9), ("patent", 0.9), ("trademark", 0.9),
("liability", 0.9), ("regulatory", 0.8), ("jurisdiction", 0.9),
("clause", 0.8), ("agreement", 0.7), ("申请", 0.6),
],
"medical": [
("医疗", 1.2), ("药物", 1.0), ("症状", 1.0), ("诊断", 1.0), ("治疗", 0.9),
("医生", 0.9), ("血压", 0.9), ("高血压", 1.0), ("糖尿病", 1.0), ("感冒", 0.9),
("剂量", 0.9), ("副作用", 0.9), ("手术", 0.9), ("患者", 0.9),
("吃药", 0.9), ("发烧", 1.0), ("疫苗", 0.9), ("感染", 0.9), ("体检", 0.7),
("medical", 1.0), ("patient", 0.9), ("symptom", 1.0), ("disease", 0.9),
("diagnosis", 1.0), ("treatment", 0.8), ("prescription", 1.0),
("dosage", 0.9), ("side effect", 0.9), ("hypertension", 1.0),
("diabetes", 1.0), ("surgery", 0.8), ("clinic", 0.7), ("vaccine", 0.9),
("infection", 0.9),
],
"general": [
("总结", 0.4), ("翻译", 0.4), ("介绍", 0.4), ("解释", 0.3),
("summarize", 0.4), ("translate", 0.4), ("explain", 0.3),
("introduce", 0.3), ("what is", 0.3), ("tell me", 0.3),
("write an essay", 0.4), ("邮件", 0.4), ("email", 0.3),
("推荐", 0.3), ("评价", 0.3),
],
}
_STOPWORDS = {
"的", "了", "吗", "呢", "啊", "是", "在", "有", "和", "与", "或", "及", "一个", "如何",
"the", "a", "an", "is", "are", "to", "of", "in", "on", "for", "with", "and",
"or", "do", "does", "can", "could", "would", "should", "please", "me", "my",
}
# 外置规则默认路径(约定优于配置:文件存在即自动加载,T-R2)
DEFAULT_RULES_FILE = Path(__file__).resolve().parent.parent / "config" / "routes.json"
def load_domain_rules(path: Optional[str | Path] = None
) -> Optional[Dict[str, List[Tuple[str, float]]]]:
"""加载外置分类规则文件,返回 {domain: [(关键词, 权重), ...]}。
文件缺失 / 格式非法 / 条目非法时返回 None(调用方安全回退内置规则,
不抛异常——规则文档可被人工编辑,编辑错误不应打垮路由)。
支持 .json(必有)与 .yamlpyyaml 可选依赖)。
"""
p = Path(path) if path else DEFAULT_RULES_FILE
if not p.exists():
return None
try:
if p.suffix.lower() in (".yaml", ".yml"):
try:
import yaml # type: ignore
except ImportError:
return None
with open(p, "r", encoding="utf-8") as f:
data = yaml.safe_load(f)
else:
with open(p, "r", encoding="utf-8") as f:
data = json.load(f)
except (OSError, ValueError, Exception): # noqa: BLE001 解析失败一律回退
return None
if not isinstance(data, dict):
return None
rules: Dict[str, List[Tuple[str, float]]] = {}
for domain, entries in data.items():
if not isinstance(domain, str) or not isinstance(entries, list):
return None
pairs: List[Tuple[str, float]] = []
for entry in entries:
if (not isinstance(entry, (list, tuple)) or len(entry) != 2
or not isinstance(entry[0], str)
or not isinstance(entry[1], (int, float))):
return None
pairs.append((entry[0], float(entry[1])))
rules[domain] = pairs
return rules
class BaseClassifier:
def classify(self, query: str) -> Classification:
raise NotImplementedError
def should_fallback(self, classification: Classification, threshold: float) -> bool:
return classification.confidence < threshold
class RuleClassifier(BaseClassifier):
"""基于关键词规则的分类器(零依赖)。"""
def __init__(self, confidence_floor: float = 0.55):
self.confidence_floor = confidence_floor
self.rules = DOMAIN_RULES
def _score(self, q: str) -> Dict[str, float]:
"""对已 lowercase 的查询按领域规则打分,返回有命中的领域分数。
命中词列表只对最终胜出领域有意义(classify 的唯一消费点),
故不在各领域上白建 list——胜出后由 _matched_rules 单独收集。
"""
scores: Dict[str, float] = {}
for domain, rules in self.rules.items():
s = 0.0
for kw, w in rules:
if kw in q:
s += w
if s > 0:
scores[domain] = s
return scores
def _matched_rules(self, q: str, domain: str) -> List[str]:
"""收集指定领域命中的关键词(保持登记顺序)。"""
return [kw for kw, _w in self.rules.get(domain, ()) if kw in q]
def classify(self, query: str) -> Classification:
q = query.lower()
raw = self._score(q)
if not raw:
# 完全无命中 -> general,低置信度
diff, ds = estimate_difficulty(query)
return Classification(
domain="general",
confidence=0.50,
difficulty=diff,
difficulty_score=ds,
raw_scores={},
matched_rules=[],
)
# 同分决胜:按领域名字典序,保证与规则表排列顺序无关的确定性
best_domain = max(sorted(raw), key=lambda d: raw[d])
best_score = raw[best_domain]
confidence = 1.0 - math.exp(-best_score)
# general 领域天然置信度压低
if best_domain == "general":
confidence = min(confidence, self.confidence_floor + 0.05)
# 与次高分的差距影响置信度(区分度)
if len(raw) > 1:
second = max(v for d, v in raw.items() if d != best_domain)
if second > 0.7 * best_score:
confidence *= 0.85
diff, ds = estimate_difficulty(query)
return Classification(
domain=best_domain,
confidence=round(min(0.99, confidence), 4),
difficulty=diff,
difficulty_score=ds,
raw_scores={k: round(v, 3) for k, v in raw.items()},
matched_rules=self._matched_rules(q, best_domain),
)
class HuggingFaceClassifier(BaseClassifier):
"""可选:基于 transformers 的序列分类模型。
仅当安装 torch+transformers 且模型可加载时可用;否则抛错提示。
"""
def __init__(self, model_name: str, num_labels: int = 5, confidence_floor: float = 0.55):
try:
from transformers import AutoModelForSequenceClassification, AutoTokenizer
except ImportError as e:
raise RuntimeError(
"HuggingFaceClassifier 需要安装 ML 依赖:pip install -r requirements-ml.txt"
) from e
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
self.model = AutoModelForSequenceClassification.from_pretrained(
model_name, num_labels=num_labels
)
self.labels = ["code", "math", "legal", "medical", "general"]
self.confidence_floor = confidence_floor
def classify(self, query: str) -> Classification:
import torch # type: ignore
inputs = self.tokenizer(query, return_tensors="pt", truncation=True, max_length=256)
with torch.no_grad():
logits = self.model(**inputs).logits
probs = torch.softmax(logits, dim=-1)[0]
idx = int(probs.argmax())
diff, ds = estimate_difficulty(query)
return Classification(
domain=self.labels[idx],
confidence=round(float(probs[idx]), 4),
difficulty=diff,
difficulty_score=ds,
raw_scores={self.labels[i]: round(float(probs[i]), 3) for i in range(len(self.labels))},
)
def build_classifier(cfg: Dict) -> BaseClassifier:
"""根据配置构建分类器。cfg 为 classifier 段配置。
T-R2cfg.rules_file 指定外置规则文件;未指定时若约定路径
config/routes.json 存在则自动加载。文件按 domain 覆盖内置规则。
"""
ctype = cfg.get("type", "rule")
floor = cfg.get("confidence_floor", 0.55)
if ctype == "rule":
clf = RuleClassifier(confidence_floor=floor)
rules_file = cfg.get("rules_file")
external = load_domain_rules(rules_file) if rules_file \
else (load_domain_rules() if DEFAULT_RULES_FILE.exists() else None)
if external:
clf.rules = {**DOMAIN_RULES, **external}
return clf
if ctype == "hf":
return HuggingFaceClassifier(cfg.get("model", "Qwen/Qwen3-0.6B"), confidence_floor=floor)
raise ValueError(f"未知分类器类型: {ctype}(支持 rule | hf")