Files
projectAIpopular/router_system/classifier.py
T
tzt 06700a5aa3 feat(v1): T-R3 采纳 llmrouter 词边界与失败安全——难度英文标记整词命中 + HF 分类器两级回落
- difficulty:标记表编译拆分——英文标记改词边界正则(修真实误判:'int' 子串
  命中 'print'/'point'、'log' 命中 'logic'、'list' 命中 'listen'),中文标记
  保持子串语义;命中行为对合法用例不变(整词出现照常计数)
- classifier:HuggingFaceClassifier 推理期异常回落内置 RuleClassifier(单次
  推理异常不打垮路由);build_classifier 的 hf 分支构造失败(ML 依赖缺失/
  模型加载失败)打印提示并回落规则分类器(外置规则照常合并)
- 新增 tests/test_failsafe.py 5 项;全量 43 passed(38+5)
2026-09-19 09:27:15 +08:00

289 lines
13 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 且模型可加载时可用;否则抛错提示。
失败安全(T-R3,采纳 llmrouter 分类失败静默降级思想):模型加载成功但
推理期异常时,自动回落内置规则分类器,不让单次推理异常打垮路由。
"""
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
self._rule_fallback = RuleClassifier(confidence_floor=confidence_floor)
def classify(self, query: str) -> Classification:
try:
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))},
)
except Exception: # noqa: BLE001 推理失败回落规则分类器(失败安全)
return self._rule_fallback.classify(query)
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":
# 失败安全(T-R3):ML 依赖缺失 / 模型加载失败时回落规则分类器
try:
return HuggingFaceClassifier(
cfg.get("model", "Qwen/Qwen3-0.6B"), confidence_floor=floor)
except Exception as e: # noqa: BLE001
print(f"[classifier] HF 分类器不可用({type(e).__name__}),回落规则分类器")
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
raise ValueError(f"未知分类器类型: {ctype}(支持 rule | hf")