docfield_extract / eval /score.py
kenzychew's picture
eval: report documents that never reached the model, grouped by cause
3907a18
Raw
History Blame Contribute Delete
13.3 kB
"""Score phase: compute all metrics and the threshold sweep from the cache.
Purely offline -- it loads cached predictions (``eval.cache``), runs the pure
aggregations (``eval.metrics``), and formats human-readable tables. Re-running
with a different threshold grid or comparison never re-runs inference, which is
the whole point of the two-phase split.
The formatted output covers the four things the methodology asks for (data spec
section 6): per-field precision/recall/F1, per-critical-field metrics, document
routing stats, and the threshold sweep trade-off curve. It also prints the
confidence distribution (so a flat sweep is explained by the backend exposing no
per-field confidence) and, as analysis only, the lowest threshold that reaches a
target auto-accept precision -- the operator still chooses the value.
Scores can come from one of two places, and every report says which:
- **as cached** (default) -- the ``confidence`` frozen in at predict time.
- **current rules** (``revalidate=True``) -- recomputed offline from the cached
predicted documents via ``eval.revalidate``.
The two diverge whenever a validation or scoring rule has changed since the
predict run. That divergence is always detected and reported, in both modes, so
a stale cache can never be mistaken for agreement (see ``eval.revalidate``).
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from eval.cache import DEFAULT_CACHE_BASE, read_entries
from eval.metrics import (
CRITICAL_FIELDS,
THRESHOLDS,
FieldMetric,
SweepRow,
compute_field_metrics,
confidence_histogram,
smallest_threshold_meeting,
summarize_errors,
sweep_thresholds,
)
from eval.normalize import is_present
from eval.revalidate import Drift, revalidate_entries
from eval.splits import DEFAULT_SPLITS_DIR, describe, select
# Auto-accept precision target on critical fields (data spec section 6).
TARGET_CRITICAL_PRECISION: float = 0.98
# Where a report's confidence scores came from. Rendered verbatim in the report
# header so a reader never has to guess which rule set produced the numbers.
SOURCE_CACHED: str = "cached"
SOURCE_CURRENT: str = "current"
@dataclass(frozen=True)
class ScoreReport:
"""Everything the score phase computed for one dataset slice."""
dataset: str
n: int
labeled_fields: tuple[str, ...]
critical_labeled: tuple[str, ...]
field_metrics: list[FieldMetric]
sweep: list[SweepRow]
confidence_hist: dict[float, int]
n_error: int
score_source: str
drift: list[Drift]
split: str
n_cached: int
error_kinds: list[tuple[str, int]]
@property
def n_reached_model(self) -> int:
"""Documents that actually produced an extraction."""
return self.n - self.n_error
@property
def revalidated(self) -> bool:
"""Whether the metrics were computed under current rules."""
return self.score_source == SOURCE_CURRENT
def _labeled_fields(entries: list[dict[str, Any]]) -> tuple[str, ...]:
"""Union of the ``labeled_fields`` recorded across cached entries."""
seen: list[str] = []
for entry in entries:
for field in entry.get("labeled_fields", []):
if field not in seen:
seen.append(field)
return tuple(seen)
def _critical_labeled(labeled: tuple[str, ...], entries: list[dict[str, Any]]) -> tuple[str, ...]:
"""Critical fields the dataset actually labels *and* has gold present for.
A critical field with no gold anywhere in the slice cannot be scored, so it
is excluded from the critical-precision denominators.
"""
result: list[str] = []
for field in CRITICAL_FIELDS:
if field not in labeled:
continue
if any(is_present(field, entry.get("gold", {}).get(field)) for entry in entries):
result.append(field)
return tuple(result)
def build_report(
dataset: str,
*,
cache_base: Path = DEFAULT_CACHE_BASE,
thresholds: tuple[float, ...] = THRESHOLDS,
revalidate: bool = False,
split: str = "all",
splits_dir: Path = DEFAULT_SPLITS_DIR,
) -> ScoreReport:
"""Load the cache for a dataset and compute the full score report.
Drift between the cached scores and current rules is detected on every call,
regardless of ``revalidate``, and recorded on the report. Substituting the
recomputed scores is opt-in; surfacing the disagreement is not.
Args:
dataset: Dataset name whose cache to score.
cache_base: Root cache directory. Defaults to ``eval/cache``.
thresholds: The threshold grid to sweep.
revalidate: Score under current rules by recomputing validation and
confidence from the cached predicted documents, instead of using the
scalars frozen in at predict time. Offline either way.
split: Which cached documents to score -- "all", "tuning", or "heldout"
(see ``eval.splits``). Selection is by id, so it does not depend on
cache read order.
splits_dir: Directory holding the split manifests.
Returns:
A :class:`ScoreReport`.
Raises:
FileNotFoundError: If no cached entries exist for the dataset.
SplitError: If ``split`` cannot be resolved for this dataset.
"""
entries = read_entries(cache_base, dataset)
if not entries:
raise FileNotFoundError(
f"No cached predictions for dataset {dataset!r} under {cache_base}. "
"Run the predict phase first."
)
n_cached = len(entries)
entries = select(entries, split, dataset=dataset, splits_dir=splits_dir)
# Always recompute, so drift is visible even when scoring from the cache.
recomputed, drift = revalidate_entries(entries)
scored = recomputed if revalidate else entries
labeled = _labeled_fields(scored)
critical_labeled = _critical_labeled(labeled, scored)
field_metrics = compute_field_metrics(scored, labeled)
sweep = sweep_thresholds(scored, critical_labeled, thresholds)
hist = confidence_histogram(scored)
n_error = sum(1 for entry in scored if entry.get("error"))
error_kinds = summarize_errors(scored)
return ScoreReport(
dataset=dataset,
n=len(entries),
labeled_fields=labeled,
critical_labeled=critical_labeled,
field_metrics=field_metrics,
sweep=sweep,
confidence_hist=hist,
n_error=n_error,
score_source=SOURCE_CURRENT if revalidate else SOURCE_CACHED,
drift=drift,
split=split,
n_cached=n_cached,
error_kinds=error_kinds,
)
# --- Formatting ----------------------------------------------------------------
def _pct(value: float | None) -> str:
"""Format an optional ratio as a percentage, or 'n/a' when undefined."""
return " n/a" if value is None else f"{value * 100:5.1f}%"
def _format_field_table(report: ScoreReport) -> list[str]:
lines = [
"Per-field metrics (whole slice):",
f" {'field':<16} {'P':>7} {'R':>7} {'F1':>7} {'pred':>4} {'gold':>4} {'ok':>4}",
]
for metric in report.field_metrics:
marker = " *" if metric.field in report.critical_labeled else " "
lines.append(
f"{marker}{metric.field:<16} {_pct(metric.precision)} {_pct(metric.recall)} "
f"{_pct(metric.f1)} {metric.n_pred:>4} {metric.n_gold:>4} {metric.n_match:>4}"
)
lines.append(" (* = critical field)")
return lines
def _format_sweep_table(report: ScoreReport, coarse_step: int = 5) -> list[str]:
lines = [
"Threshold sweep (critical fields on the auto-accepted subset):",
f" {'thr':>5} {'accept':>7} {'accept%':>8} {'crit P':>8} {'crit R':>8}",
]
for index, row in enumerate(report.sweep):
# Print a coarse grid plus the final threshold to keep it readable.
is_grid = index % coarse_step == 0 or index == len(report.sweep) - 1
if not is_grid:
continue
lines.append(
f" {row.threshold:>5.2f} {row.n_accepted:>7} {row.accept_rate * 100:>7.1f}% "
f"{_pct(row.crit_precision)} {_pct(row.crit_recall)}"
)
return lines
def _format_confidence(report: ScoreReport) -> list[str]:
origin = (
"recomputed under CURRENT rules"
if report.revalidated
else "AS CACHED at predict time"
)
lines = [f"Confidence distribution ({origin}):"]
for value, count in report.confidence_hist.items():
bar = "#" * count
lines.append(f" {value:>5.2f} {count:>3} {bar}")
return lines
def _format_outcomes(report: ScoreReport) -> list[str]:
"""Render how many documents reached the model at all, and why not.
A document that never reached the model is not a model error, but it lowers
recall exactly as a genuine miss does (nothing predicted against a gold
value present) while leaving precision untouched. Stating the split here
keeps a reader from reading an infrastructure outage as extraction quality.
"""
lines = ["Document outcomes:"]
if not report.n_error:
lines.append(f" all {report.n} documents reached the model.")
return lines
share = report.n_error / report.n * 100 if report.n else 0.0
lines += [
f" reached the model : {report.n_reached_model:>4} of {report.n}",
f" never reached the model: {report.n_error:>4} of {report.n} ({share:.1f}%)",
]
for cause, count in report.error_kinds:
lines.append(f" {count:>4} {cause}")
lines += [
" A document that never reached the model produced no extraction. It is",
" counted as a miss in recall (gold present, nothing predicted) but not in",
" precision, so recall below is depressed by the outage, not by the model.",
]
return lines
def _format_drift(report: ScoreReport, max_shown: int = 8) -> list[str]:
"""Render the cache-drift warning (or a one-line all-clear)."""
if not report.drift:
return [
"Cache drift check:",
f" OK -- all {report.n} cached scores agree with current rules.",
]
n_drift = len(report.drift)
n_routing = sum(1 for d in report.drift if d.routing_changed)
lines = [
"Cache drift check:",
f" WARNING -- {n_drift} of {report.n} cached scores disagree with current rules"
+ (f" ({n_routing} change the hard-failure verdict)." if n_routing else "."),
" The cached scalars were frozen by the rule set in force at predict time;",
" a validation or scoring rule has changed since.",
]
lines.append(
" These metrics USE the recomputed scores."
if report.revalidated
else " These metrics USE the stale cached scores -- pass --revalidate for current rules."
)
for record in report.drift[:max_shown]:
lines.append(f" {record.describe()}")
if n_drift > max_shown:
lines.append(f" ... and {n_drift - max_shown} more.")
return lines
def _format_routing(report: ScoreReport) -> list[str]:
target = smallest_threshold_meeting(report.sweep, TARGET_CRITICAL_PRECISION)
lines = [
"Operating point analysis (you choose the threshold):",
f" Target: auto-accept precision on critical fields "
f"{report.critical_labeled or '(none labeled)'} >= "
f"{TARGET_CRITICAL_PRECISION:.2f}",
]
if not report.critical_labeled:
lines.append(
" This dataset labels none of total/tax/invoice_number, so critical "
"auto-accept precision cannot be measured here."
)
elif target is None:
lines.append(
" No threshold in the sweep reaches the target with any "
"auto-accepted document (see the confidence distribution above)."
)
else:
lines.append(
f" Lowest qualifying threshold: {target.threshold:.2f} "
f"(accept {target.n_accepted}/{target.n_total} = "
f"{target.accept_rate * 100:.1f}%, crit P {_pct(target.crit_precision)}, "
f"crit R {_pct(target.crit_recall)})."
)
return lines
def format_report(report: ScoreReport) -> str:
"""Render a :class:`ScoreReport` as a plain-text report.
Args:
report: The computed score report.
Returns:
A multi-line string ready to print.
"""
source = (
"CURRENT RULES (recomputed offline from cached predictions)"
if report.revalidated
else "AS CACHED at predict time (may be stale vs current rules)"
)
header = [
"=" * 68,
f"Evaluation: {report.dataset} (n={report.n}, errors={report.n_error})",
f"Split: {describe(report.split, report.n, report.n_cached)}",
f"Scores: {source}",
f"Labeled fields: {', '.join(report.labeled_fields) or '(none)'}",
"=" * 68,
]
sections = [
header,
_format_field_table(report),
_format_outcomes(report),
_format_confidence(report),
_format_drift(report),
_format_sweep_table(report),
_format_routing(report),
]
return "\n\n".join("\n".join(section) for section in sections)