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)
This commit is contained in:
+33
-15
@@ -214,6 +214,8 @@ class HuggingFaceClassifier(BaseClassifier):
|
||||
"""可选:基于 transformers 的序列分类模型。
|
||||
|
||||
仅当安装 torch+transformers 且模型可加载时可用;否则抛错提示。
|
||||
失败安全(T-R3,采纳 llmrouter 分类失败静默降级思想):模型加载成功但
|
||||
推理期异常时,自动回落内置规则分类器,不让单次推理异常打垮路由。
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str, num_labels: int = 5, confidence_floor: float = 0.55):
|
||||
@@ -229,23 +231,27 @@ class HuggingFaceClassifier(BaseClassifier):
|
||||
)
|
||||
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:
|
||||
import torch # type: ignore
|
||||
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))},
|
||||
)
|
||||
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:
|
||||
@@ -265,6 +271,18 @@ def build_classifier(cfg: Dict) -> BaseClassifier:
|
||||
clf.rules = {**DOMAIN_RULES, **external}
|
||||
return clf
|
||||
if ctype == "hf":
|
||||
return HuggingFaceClassifier(cfg.get("model", "Qwen/Qwen3-0.6B"), confidence_floor=floor)
|
||||
# 失败安全(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)")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user