- labeler.py:derive_true_tier 纯函数(规则0 人工优先 / plan_multi->T3 / 失败->下一档(T3保持) / 成功->executed / 空结果跳过)+ derive_true_tiers 批量回填 + 180d 留存清理 - calibrate.py:compute_thresholds(τ1/τ3 分组扫描——并列分数整组判定, 精度 >= 1-α 的最大覆盖阈值;<min_labels -> ok=False 沿用 last-good); save_thresholds 工件落盘+登记;load_active 回退链 active->同kind扫描->last-good->保守值 - store:all_observations/list_artifacts 支撑方法 - 测试 +9(四规则/批量清理/覆盖达标/标签不足/工件往返与回退),全量 395 passed
133 lines
5.1 KiB
Python
133 lines
5.1 KiB
Python
"""split-conformal 阈值校准(T-G3)+ last-good 回退 + 工件读写(D-G3)。
|
||
|
||
目标(§5 policy.alpha):P(true>T1 | 判 T1) ≤ α(τ3 同理对 T3)。
|
||
实现:calib 集按分数降序扫描,取**满足精度的最低阈值**(覆盖率最大);
|
||
无任何满足点 -> 阈值取最高分 + ε(全拒,保守)。
|
||
标签数 < min_labels -> 沿用 last-good(D-G3)-> 再无 -> 内置保守值。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import time
|
||
from pathlib import Path
|
||
from typing import Any, Dict, List, Optional, Tuple
|
||
|
||
from gateway.sense.store import json_dumps
|
||
|
||
CONSERVATIVE = {"t1": 0.90, "t3": 0.90}
|
||
THRESHOLD_KIND = "thresholds"
|
||
|
||
|
||
def _threshold_for(samples: List[Tuple[float, bool]], alpha: float) -> Optional[float]:
|
||
"""样本 [(score, true==tier)] -> 满足精度 >= 1-α 的最低阈值(并列分数整组判定)。
|
||
|
||
按唯一分数降序累积计数;每组边界处检查精度;取满足精度的最大覆盖(最低边界)。
|
||
无任何满足点 -> None(调用方全拒,保守)。
|
||
"""
|
||
if not samples:
|
||
return None
|
||
groups: Dict[float, List[int]] = {}
|
||
for score, ok in samples:
|
||
g = groups.setdefault(score, [0, 0])
|
||
g[0] += 1
|
||
g[1] += 1 if ok else 0
|
||
total = 0
|
||
correct = 0
|
||
best: Optional[float] = None
|
||
for score in sorted(groups, reverse=True):
|
||
cnt, ok_cnt = groups[score]
|
||
total += cnt
|
||
correct += ok_cnt
|
||
if correct / total >= 1.0 - alpha:
|
||
best = score
|
||
return best
|
||
|
||
|
||
def compute_thresholds(rows: List[Dict[str, Any]], alpha: float = 0.05,
|
||
min_labels: int = 500) -> Dict[str, Any]:
|
||
"""rows: [{probs(JSON 串或 dict), true_tier}] -> {t1, t3, coverage, n, ok}。
|
||
|
||
ok=False 表示标签不足(沿用 last-good 的信号)。
|
||
coverage = 通过阈值的 calib 样本占比(实测覆盖率,§6 晋升门比对 α+2%)。
|
||
"""
|
||
labeled = []
|
||
for r in rows:
|
||
probs = r.get("probs")
|
||
if isinstance(probs, str):
|
||
try:
|
||
probs = json.loads(probs)
|
||
except json.JSONDecodeError:
|
||
continue
|
||
if not isinstance(probs, dict):
|
||
continue
|
||
true_tier = r.get("true_tier")
|
||
if true_tier not in ("T1", "T2", "T3"):
|
||
continue
|
||
labeled.append((float(probs.get("t1", 0)), float(probs.get("t3", 0)),
|
||
true_tier))
|
||
n = len(labeled)
|
||
if n < min_labels:
|
||
return {"t1": None, "t3": None, "coverage": 0.0, "n": n, "ok": False}
|
||
|
||
s1 = [(p1, t == "T1") for p1, _p3, t in labeled]
|
||
s3 = [(p3, t == "T3") for _p1, p3, t in labeled]
|
||
t1 = _threshold_for(s1, alpha)
|
||
t3 = _threshold_for(s3, alpha)
|
||
if t1 is None or t3 is None:
|
||
return {"t1": t1, "t3": t3, "coverage": 0.0, "n": n, "ok": False}
|
||
coverage = (sum(1 for p1, _p3, _t in labeled if p1 >= t1)
|
||
+ sum(1 for _p1, p3, _t in labeled if p3 >= t3)) / (2 * n)
|
||
return {"t1": round(t1, 4), "t3": round(t3, 4),
|
||
"coverage": round(coverage, 4), "n": n, "ok": True}
|
||
|
||
|
||
def save_thresholds(store, models_dir: str | Path, version: str,
|
||
th: Dict[str, Any], activate: bool = False) -> str:
|
||
"""阈值工件落盘 thresholds.json + 登记工件表。返回文件路径。"""
|
||
d = Path(models_dir) / version
|
||
d.mkdir(parents=True, exist_ok=True)
|
||
path = d / "thresholds.json"
|
||
path.write_text(json_dumps(th), encoding="utf-8")
|
||
store.register_artifact(version, THRESHOLD_KIND, str(path),
|
||
{k: th.get(k) for k in ("t1", "t3", "coverage", "n")},
|
||
active=activate)
|
||
return str(path)
|
||
|
||
|
||
def load_active(store, models_dir: str | Path) -> Dict[str, Any]:
|
||
"""active 阈值工件 -> 旧工件扫描回退(last-good)-> 内置保守值(D-G3 回退链)。"""
|
||
art = store.active_artifact(THRESHOLD_KIND)
|
||
if art:
|
||
th = _read_threshold(art)
|
||
if th is not None:
|
||
th.setdefault("version", art["version"])
|
||
return th
|
||
# active 文件缺失/损坏:按时间倒序扫描同 kind 工件(last-good 语义)
|
||
for cand in store.list_artifacts(THRESHOLD_KIND):
|
||
if cand["version"] == art["version"]:
|
||
continue
|
||
th = _read_threshold(cand)
|
||
if th is not None:
|
||
th.setdefault("version", cand["version"])
|
||
return th
|
||
last_good = Path(models_dir) / "last-good.json"
|
||
if last_good.exists():
|
||
try:
|
||
th = json.loads(last_good.read_text(encoding="utf-8"))
|
||
if th.get("t1") is not None and th.get("t3") is not None:
|
||
th.setdefault("version", "last-good")
|
||
return th
|
||
except (json.JSONDecodeError, OSError):
|
||
pass
|
||
return {**CONSERVATIVE, "version": "conservative"}
|
||
|
||
|
||
def _read_threshold(art: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||
try:
|
||
th = json.loads(Path(art["path"]).read_text(encoding="utf-8"))
|
||
if isinstance(th, dict) and th.get("t1") is not None and th.get("t3") is not None:
|
||
return th
|
||
except (json.JSONDecodeError, OSError):
|
||
pass
|
||
return None
|