Files
projectAIpopular/gateway/sense/grader.py
T
tzt 27a092a3f8 feat(sense): T-X3 路由决策缓存——live 模式 embed+线性头去重(采纳 cortiq 决策哈希缓存)
- gateway/sense/decision_cache.py:DecisionCache(sha256(scope+文本) 键、
  LRU + TTL 60s、4096 条上限、时钟可注入),对齐 auth._AuthCache 进程内模式
- Grader:仅 live 模式缓存纯决策负载(probs/head_version/vec);
  命中跳过 embed + 线性头预测;tier 仍按当前 conformal 阈值即时重算;
  观察落盘(observer.log)不因缓存命中跳过——校准数据完整性不受影响
- collect/shadow 模式不走缓存(校准必须全量);invalidate() 同步清空缓存
  (阈值/工件切换即时生效,TTL 仅兜底陈旧)

pytest 455 passed(T-X2 后 447 + 8)
2026-09-18 22:41:15 +08:00

174 lines
6.8 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.
"""分级决策组合(T-G5,§8 主时序):
embedder.embed(挂 -> fallback=T2+规则门,D-G4
-> features.gatet1_hard_ok
-> LinearHead.predict -> {p1,p2,p3}
-> conformal 阈值:p1>=τ1 且 t1_hard_ok -> T1p3>=τ3 或 repo_signals -> T3;其余 T2
-> mode 裁剪:collect/shadow 只写观察(decided≠executed),live 返回决策
-> observer.log(全模式必写)
"""
from __future__ import annotations
import time
import uuid
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional
from gateway.sense.calibrate import load_active
from gateway.sense.classifier import LinearHead
from gateway.sense.config import SenseConfig
from gateway.sense.decision_cache import DecisionCache
from gateway.sense.embedder import embed
from gateway.sense.errors import EmbedderDown
from gateway.sense.features import gate
@dataclass
class TierDecision:
tier: str # 决策档(live 时即执行档的候选)
probs: Dict[str, float]
hard_gates: Dict[str, Any]
thresholds_version: str = ""
head_version: str = ""
mode: str = "collect"
fallback: bool = False # True = embedder/工件故障,规则门退化
executed_tier: str = "" # 现行为执行的档位(collect/shadow 记录用)
def _rule_tier(feats) -> str:
"""规则门兜底(无模型时的决策):仓库级 -> T3;t1 硬门过 -> T1;否则 T2。"""
if feats.repo_signals:
return "T3"
if feats.t1_hard_ok:
return "T1"
return "T2"
class Grader:
"""决策器(持有 active 工件缓存;工件/阈值切换后调 invalidate)。"""
def __init__(self, cfg: SenseConfig, store, observer=None):
self.cfg = cfg
self.store = store
self.observer = observer
self._head: Optional[LinearHead] = None
self._head_loaded = False
self._thresholds: Optional[Dict[str, Any]] = None
self._dcache = DecisionCache() # T-X3live 模式决策缓存
def _load_head(self):
if not self._head_loaded:
art = self.store.active_artifact("head")
if art:
try:
self._head = LinearHead.load(art["path"], version=art["version"])
except Exception: # noqa: BLE001
self._head = None
self._head_loaded = True
return self._head
def invalidate(self) -> None:
"""工件 promote 后调用(重载 active 工件与阈值;同步清空决策缓存)。"""
self._head = None
self._head_loaded = False
self._thresholds = None
self._dcache.clear()
def _thresholds_cached(self) -> Dict[str, Any]:
if self._thresholds is None:
self._thresholds = load_active(self.store, self.cfg.models_dir)
return self._thresholds
async def decide(self, query_or_messages, consumer: str,
domain: str = "", request_id: str = "",
executed_tier: str = "") -> TierDecision:
"""分级决策(§8 时序;全模式写观察)。
T-X3live 模式下对 (consumer, 文本) 的纯决策负载(probs/head_version/vec
做 60s LRU 缓存,命中时跳过 embed + 线性头;观察/审计照常落盘。
collect/shadow 模式不走缓存(校准数据必须全量产出)。
"""
ts = time.time()
rid = request_id or ("rt" + uuid.uuid4().hex[:10])
text = (query_or_messages if isinstance(query_or_messages, str)
else "\n".join(str(m.get("content") or "")
for m in query_or_messages))
feats = gate(query_or_messages, consumer, self.cfg)
probs: Dict[str, float] = {}
head_version = ""
fallback = False
vec: Optional[List[int]] = None
cached = self._dcache.get(consumer, text) \
if self.cfg.mode == "live" else None
if cached is not None:
probs = dict(cached["probs"])
head_version = str(cached["head_version"])
vec = cached["vec"]
else:
try:
vec = await embed(text, self.cfg.embedder)
except EmbedderDown:
fallback = True
head = None if fallback else self._load_head()
if head is None:
fallback = True
if not fallback and vec is not None:
probs = head.predict([float(v) for v in vec])
head_version = head.version
if self.cfg.mode == "live":
self._dcache.put(consumer, text,
{"probs": dict(probs),
"head_version": head_version,
"vec": vec})
th = self._thresholds_cached()
th_version = str(th.get("version") or "")
# ---- 决策 ----
if fallback:
tier = _rule_tier(feats) # D-G4 规则门退化
else:
t1_ok = (probs.get("t1", 0.0) >= float(th.get("t1", 0.9))
and feats.t1_hard_ok)
t3_ok = (probs.get("t3", 0.0) >= float(th.get("t3", 0.9))
or feats.repo_signals)
if t3_ok and not t1_ok:
tier = "T3"
elif t1_ok:
tier = "T1"
else:
tier = "T2"
# ---- mode 裁剪(D-G7----
if self.cfg.mode == "live":
executed = tier # live:决策即执行
else:
# collect/shadow:现行为——规则门等价(T1 门/T3 信号)近似 v2 现状
executed = executed_tier or _rule_tier(feats)
decision = TierDecision(
tier=tier, probs=probs,
hard_gates={"turns": feats.turns, "est_tokens": feats.est_tokens,
"single_turn": feats.single_turn,
"intent_blocked": feats.intent_blocked,
"repo_signals": feats.repo_signals,
"over_length": feats.over_length,
"t1_hard_ok": feats.t1_hard_ok},
thresholds_version=th_version, head_version=head_version,
mode=self.cfg.mode, fallback=fallback,
executed_tier=executed)
# ---- 观察必写(全模式)----
if self.observer is not None:
self.observer.log(__import__("gateway.sense.observer",
fromlist=["Observation"]).Observation(
request_id=rid, consumer=consumer, decided_tier=tier,
executed_tier=executed, probs=probs,
policy_version=f"head:{head_version}|th:{th_version}",
features=decision.hard_gates, embedding=vec,
bucket="default", domain=domain, ts=ts))
return decision