Codex_Extractor / smoke_signal /scripts /10_recalibrate.py
Pointf5ive
Training phase: scripts 09+10, --model flag in 03, gold schema patch
5ffc645
Raw
History Blame Contribute Delete
15.1 kB
#!/usr/bin/env python3
"""
Smoke Signal — Stage 12: Confidence Recalibration
==================================================
Recomputes confidence thresholds from the full gold set using a simple
precision/recall threshold analysis per region class.
Outputs:
- manifest/confidence_calibration.json
- training/runs/<RUN_ID>_recalibration.json
- manifest/run_log.csv entry (governance)
"""
import argparse
import csv
import json
from datetime import datetime, timezone
from pathlib import Path
from typing import Dict, List, Optional, Tuple
ROOT = Path(__file__).resolve().parents[1]
MANIFEST_CSV = ROOT / "manifest" / "source_manifest.csv"
CALIBRATION_JSON = ROOT / "manifest" / "confidence_calibration.json"
RUN_LOG_CSV = ROOT / "manifest" / "run_log.csv"
GOLD_FILE = ROOT / "gold" / "gold_corrections.jsonl"
RUNS_DIR = ROOT / "training" / "runs"
ELIGIBLE_RIGHTS = {"public-domain", "licensed-owned", "controlled-internal"}
BLOCKED_RIGHTS = {"unknown", "excluded"}
RUN_LOG_FIELDS = [
"run_id",
"date",
"operator",
"config_version",
"schema_version",
"source_batch",
"pages_processed",
"errors",
"cost_usd",
"output_path",
"notes",
]
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def utc_iso() -> str:
return utc_now().isoformat().replace("+00:00", "Z")
def ensure_run_dirs() -> None:
RUNS_DIR.mkdir(parents=True, exist_ok=True)
def ensure_run_log() -> None:
RUN_LOG_CSV.parent.mkdir(parents=True, exist_ok=True)
if RUN_LOG_CSV.exists():
return
with open(RUN_LOG_CSV, "w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=RUN_LOG_FIELDS)
writer.writeheader()
def append_run_log(row: Dict[str, str]) -> None:
ensure_run_log()
with open(RUN_LOG_CSV, "a", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=RUN_LOG_FIELDS)
writer.writerow({k: row.get(k, "") for k in RUN_LOG_FIELDS})
def load_manifest() -> Dict[str, Dict[str, str]]:
out: Dict[str, Dict[str, str]] = {}
if not MANIFEST_CSV.exists():
return out
with open(MANIFEST_CSV, newline="", encoding="utf-8") as f:
for row in csv.DictReader(f):
book_id = str(row.get("book_id", "")).strip()
if book_id:
out[book_id] = row
return out
def load_gold_records(path: Path) -> List[Dict]:
rows: List[Dict] = []
if not path.exists():
return rows
with open(path, encoding="utf-8") as f:
for idx, line in enumerate(f, start=1):
line = line.strip()
if not line:
continue
try:
rec = json.loads(line)
rec["_line"] = idx
rows.append(rec)
except json.JSONDecodeError:
continue
return rows
def parse_bool(value) -> Optional[bool]:
if isinstance(value, bool):
return value
if value is None:
return None
text = str(value).strip().lower()
if text in {"true", "1", "yes", "y"}:
return True
if text in {"false", "0", "no", "n"}:
return False
return None
def parse_confidence(value) -> Optional[float]:
try:
conf = float(value)
return max(0.0, min(1.0, conf))
except (TypeError, ValueError):
return None
def infer_was_correct(rec: Dict) -> Optional[bool]:
explicit = parse_bool(rec.get("was_correct"))
if explicit is not None:
return explicit
final_text = str(rec.get("final_text", "")).strip()
raw_text = str(rec.get("raw_text", rec.get("raw_ocr", ""))).strip()
if final_text and raw_text:
return final_text == raw_text
return None
def precision_recall_at_threshold(rows: List[Tuple[float, int]], threshold: float) -> Dict[str, float]:
tp = fp = fn = tn = 0
for conf, label in rows:
pred = 1 if conf >= threshold else 0
if pred == 1 and label == 1:
tp += 1
elif pred == 1 and label == 0:
fp += 1
elif pred == 0 and label == 1:
fn += 1
else:
tn += 1
precision = tp / (tp + fp) if (tp + fp) else 0.0
recall = tp / (tp + fn) if (tp + fn) else 0.0
return {
"tp": tp,
"fp": fp,
"fn": fn,
"tn": tn,
"precision": precision,
"recall": recall,
}
def f_beta(precision: float, recall: float, beta: float) -> float:
if precision <= 0 and recall <= 0:
return 0.0
beta2 = beta * beta
denom = (beta2 * precision) + recall
if denom <= 0:
return 0.0
return (1 + beta2) * (precision * recall) / denom
def select_thresholds(
rows: List[Tuple[float, int]],
auto_precision_target: float,
auto_recall_floor: float,
review_recall_target: float,
review_precision_floor: float,
quarantine_gap: float,
) -> Dict:
unique_thresholds = sorted({round(conf, 4) for conf, _ in rows})
if not unique_thresholds:
return {
"auto_accept": 0.85,
"review": 0.60,
"quarantine": 0.35,
"metrics": {},
}
# Add boundary values so we can always compute a fallback.
thresholds = sorted(set([0.0, 1.0] + unique_thresholds))
# Auto-accept: highest threshold meeting strict precision target.
auto_t = None
for t in thresholds:
m = precision_recall_at_threshold(rows, t)
if m["precision"] >= auto_precision_target and m["recall"] >= auto_recall_floor:
auto_t = t
if auto_t is None:
# Fallback: maximize F0.5 to prioritize precision.
auto_t = max(thresholds, key=lambda t: f_beta(
precision_recall_at_threshold(rows, t)["precision"],
precision_recall_at_threshold(rows, t)["recall"],
beta=0.5,
))
# Review threshold: below/at auto threshold, try to capture most true positives.
review_candidates = [t for t in thresholds if t <= auto_t]
review_t = None
for t in review_candidates:
m = precision_recall_at_threshold(rows, t)
if m["recall"] >= review_recall_target and m["precision"] >= review_precision_floor:
review_t = t
break
if review_t is None:
# Fallback: maximize F1 while respecting t <= auto_t.
review_t = max(review_candidates, key=lambda t: f_beta(
precision_recall_at_threshold(rows, t)["precision"],
precision_recall_at_threshold(rows, t)["recall"],
beta=1.0,
))
review_t = min(review_t, auto_t)
quarantine_t = max(0.0, review_t - quarantine_gap)
quarantine_t = min(quarantine_t, review_t)
# Round for readability and stable diffs.
auto_t = round(float(auto_t), 3)
review_t = round(float(review_t), 3)
quarantine_t = round(float(quarantine_t), 3)
# Guarantee monotonic order.
if review_t > auto_t:
review_t = auto_t
if quarantine_t > review_t:
quarantine_t = review_t
return {
"auto_accept": auto_t,
"review": review_t,
"quarantine": quarantine_t,
"metrics": {
"auto": precision_recall_at_threshold(rows, auto_t),
"review": precision_recall_at_threshold(rows, review_t),
"quarantine": precision_recall_at_threshold(rows, quarantine_t),
},
}
def build_region_rows(
gold_records: List[Dict],
manifest: Dict[str, Dict[str, str]],
rights_class_filter: Optional[str],
) -> Tuple[Dict[str, List[Tuple[float, int]]], Dict[str, int]]:
by_region: Dict[str, List[Tuple[float, int]]] = {}
counters = {
"input": len(gold_records),
"used": 0,
"skipped_missing_book": 0,
"skipped_missing_manifest": 0,
"skipped_blocked_rights": 0,
"skipped_rights_filter": 0,
"skipped_missing_confidence": 0,
"skipped_missing_label": 0,
}
for rec in gold_records:
book_id = str(rec.get("book_id", "")).strip()
if not book_id:
counters["skipped_missing_book"] += 1
continue
manifest_row = manifest.get(book_id)
if not manifest_row:
counters["skipped_missing_manifest"] += 1
continue
rights = str(manifest_row.get("rights_class", "unknown")).strip().lower()
if rights in BLOCKED_RIGHTS or rights not in ELIGIBLE_RIGHTS:
counters["skipped_blocked_rights"] += 1
continue
if rights_class_filter and rights != rights_class_filter:
counters["skipped_rights_filter"] += 1
continue
conf = parse_confidence(rec.get("confidence"))
if conf is None:
counters["skipped_missing_confidence"] += 1
continue
was_correct = infer_was_correct(rec)
if was_correct is None:
counters["skipped_missing_label"] += 1
continue
# Positive class = OCR output was correct.
label = 1 if was_correct else 0
region_class = str(rec.get("region_class", "narration")).strip() or "narration"
by_region.setdefault(region_class, []).append((conf, label))
by_region.setdefault("_default", []).append((conf, label))
counters["used"] += 1
return by_region, counters
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Smoke Signal — Stage 12: Recalibration")
parser.add_argument("--gold-file", default=str(GOLD_FILE), help="Path to gold corrections JSONL")
parser.add_argument(
"--rights-class",
default=None,
choices=sorted(ELIGIBLE_RIGHTS),
help="Optional rights class filter",
)
parser.add_argument("--operator", default="codex", help="Operator for governance log")
parser.add_argument("--auto-precision-target", type=float, default=0.98)
parser.add_argument("--auto-recall-floor", type=float, default=0.20)
parser.add_argument("--review-recall-target", type=float, default=0.90)
parser.add_argument("--review-precision-floor", type=float, default=0.60)
parser.add_argument("--quarantine-gap", type=float, default=0.20)
parser.add_argument("--min-samples-per-class", type=int, default=10)
parser.add_argument("--run-id", default=None, help="Optional explicit run id")
return parser.parse_args()
def main() -> None:
args = parse_args()
ensure_run_dirs()
run_id = args.run_id or f"SS-CAL-{utc_now().strftime('%Y%m%d-%H%M%S')}"
gold_path = Path(args.gold_file).expanduser().resolve()
if not gold_path.exists():
raise FileNotFoundError(f"Gold file not found: {gold_path}")
manifest = load_manifest()
if not manifest:
raise RuntimeError("Manifest is empty. Cannot enforce rights controls.")
gold_records = load_gold_records(gold_path)
if not gold_records:
raise RuntimeError("No valid gold records found.")
by_region, counters = build_region_rows(gold_records, manifest, args.rights_class)
if counters["used"] == 0:
raise RuntimeError(f"No usable records for recalibration after filtering. Counters: {counters}")
calibration: Dict[str, Dict] = {}
report_regions: Dict[str, Dict] = {}
for region_class, rows in by_region.items():
if len(rows) < args.min_samples_per_class and region_class != "_default":
# Too little data for a reliable per-class threshold; defer to default.
continue
selected = select_thresholds(
rows=rows,
auto_precision_target=args.auto_precision_target,
auto_recall_floor=args.auto_recall_floor,
review_recall_target=args.review_recall_target,
review_precision_floor=args.review_precision_floor,
quarantine_gap=args.quarantine_gap,
)
corrections = sum(1 for _, label in rows if label == 0)
calibration[region_class] = {
"auto_accept": selected["auto_accept"],
"review": selected["review"],
"quarantine": selected["quarantine"],
"corrections": corrections,
}
report_regions[region_class] = {
"samples": len(rows),
"correct": sum(1 for _, label in rows if label == 1),
"incorrect": sum(1 for _, label in rows if label == 0),
"thresholds": calibration[region_class],
"metrics": selected["metrics"],
}
# Ensure required defaults exist for runtime readers.
if "_default" not in calibration:
calibration["_default"] = {
"auto_accept": 0.85,
"review": 0.60,
"quarantine": 0.35,
"corrections": 0,
}
default_entry = calibration["_default"]
for cls in ["narration", "dialogue-speech-bubble", "caption", "title", "sign-label"]:
if cls not in calibration:
calibration[cls] = dict(default_entry)
CALIBRATION_JSON.parent.mkdir(parents=True, exist_ok=True)
with open(CALIBRATION_JSON, "w", encoding="utf-8") as f:
json.dump(calibration, f, indent=2, ensure_ascii=False)
report = {
"run_id": run_id,
"generated_at": utc_iso(),
"config_version": "ss_confidence_calibration_v0.1",
"schema_version": "ss_confidence_calibration_report_v1",
"gold_file": str(gold_path),
"rights_class_filter": args.rights_class,
"counters": counters,
"regions": report_regions,
"output_file": str(CALIBRATION_JSON),
}
report_path = RUNS_DIR / f"{run_id}_recalibration.json"
with open(report_path, "w", encoding="utf-8") as f:
json.dump(report, f, indent=2, ensure_ascii=False)
append_run_log(
{
"run_id": run_id,
"date": utc_now().date().isoformat(),
"operator": args.operator,
"config_version": "ss_confidence_calibration_v0.1",
"schema_version": "ss_confidence_calibration_report_v1",
"source_batch": args.rights_class or "auto",
"pages_processed": str(counters["used"]),
"errors": str(
counters["skipped_missing_book"]
+ counters["skipped_missing_manifest"]
+ counters["skipped_blocked_rights"]
+ counters["skipped_rights_filter"]
+ counters["skipped_missing_confidence"]
+ counters["skipped_missing_label"]
),
"cost_usd": "",
"output_path": str(CALIBRATION_JSON.relative_to(ROOT)),
"notes": json.dumps(
{
"auto_precision_target": args.auto_precision_target,
"review_recall_target": args.review_recall_target,
"regions_calibrated": sorted(report_regions.keys()),
},
ensure_ascii=False,
),
}
)
print(f"Recalibration complete: {CALIBRATION_JSON}")
print(f"Report: {report_path}")
print(f"Governance log updated: {RUN_LOG_CSV}")
if __name__ == "__main__":
main()