limoXD's picture
Release v0.11 finite-selection hardening
5ee4f7e verified
Raw
History Blame Contribute Delete
7.94 kB
from __future__ import annotations
import math
import re
import statistics
from collections import Counter
from collections.abc import Sequence
from dataclasses import dataclass
from typing import Literal
EvaluationLabel = Literal["clean", "typo", "unspecified"]
PredictionProvenance = Literal["provider", "deberta", "lfm", "rule", "unspecified"]
_REASON_PATTERN = re.compile(r"^[a-z][a-z0-9_]{0,63}$")
@dataclass(frozen=True, slots=True)
class EvaluationItem:
item_id: str
input_text: str
references: tuple[str, ...]
label: EvaluationLabel = "unspecified"
def __post_init__(self) -> None:
if not self.item_id:
raise ValueError("item_id must not be empty")
if not self.references or any(not reference for reference in self.references):
raise ValueError("references must contain non-empty strings")
if self.label not in ("clean", "typo", "unspecified"):
raise ValueError(f"unsupported evaluation label: {self.label}")
@dataclass(frozen=True, slots=True)
class Prediction:
item_id: str
candidates: tuple[str, ...]
provenance: PredictionProvenance = "unspecified"
reason: str | None = None
margin: float | None = None
def __post_init__(self) -> None:
if not self.item_id:
raise ValueError("prediction item_id must not be empty")
if any(not candidate for candidate in self.candidates):
raise ValueError("prediction candidates must not contain empty strings")
if len(set(self.candidates)) != len(self.candidates):
raise ValueError("prediction candidates must be unique")
if self.provenance not in ("provider", "deberta", "lfm", "rule", "unspecified"):
raise ValueError(f"unsupported prediction provenance: {self.provenance}")
if self.reason is not None and not _REASON_PATTERN.fullmatch(self.reason):
raise ValueError("prediction reason must be a lowercase machine-readable label")
if self.margin is not None and not math.isfinite(self.margin):
raise ValueError("prediction margin must be finite")
def character_edit_distance(left: str, right: str) -> int:
if len(left) < len(right):
left, right = right, left
previous = list(range(len(right) + 1))
for left_index, left_character in enumerate(left, start=1):
current = [left_index]
for right_index, right_character in enumerate(right, start=1):
current.append(
min(
current[-1] + 1,
previous[right_index] + 1,
previous[right_index - 1] + (left_character != right_character),
)
)
previous = current
return previous[-1]
def minimum_cer(value: str, references: Sequence[str]) -> float:
return min(
character_edit_distance(value, reference) / len(reference) for reference in references
)
def evaluate_predictions(
items: Sequence[EvaluationItem],
predictions: Sequence[Prediction],
*,
candidate_limit: int,
) -> dict[str, object]:
if candidate_limit < 1:
raise ValueError("candidate_limit must be positive")
if not items:
raise ValueError("evaluation items must not be empty")
item_ids = [item.item_id for item in items]
if len(set(item_ids)) != len(item_ids):
raise ValueError("duplicate item_id in evaluation items")
prediction_ids = [prediction.item_id for prediction in predictions]
if len(set(prediction_ids)) != len(prediction_ids):
raise ValueError("duplicate prediction item_id")
unknown_ids = set(prediction_ids).difference(item_ids)
if unknown_ids:
raise ValueError(f"prediction has unknown item_id: {min(unknown_ids)}")
for item in items:
if item.label == "clean" and item.input_text not in item.references:
raise ValueError(f"clean item input must be a reference: {item.item_id}")
prediction_by_id = {prediction.item_id: prediction for prediction in predictions}
cers: list[float] = []
baseline_cers: list[float] = []
correct = 0
abstentions = 0
candidate_recall = 0
accepted_rows = 0
accepted_correct = 0
clean_rows = 0
overcorrections = 0
typo_rows = 0
typo_correct = 0
improved_rows = 0
worsened_rows = 0
unchanged_rows = 0
accepted_candidate_miss_rows = 0
selection_error_rows = 0
provenance_counts: Counter[str] = Counter()
selection_reason_counts: Counter[str] = Counter()
reported_margins: list[float] = []
for item in items:
prediction = prediction_by_id.get(item.item_id)
candidates = prediction.candidates[:candidate_limit] if prediction else ()
if prediction is not None:
provenance_counts[prediction.provenance] += 1
if prediction.reason is not None:
selection_reason_counts[prediction.reason] += 1
if prediction.margin is not None:
reported_margins.append(prediction.margin)
abstentions += not candidates
accepted_rows += bool(candidates)
has_reference_candidate = any(candidate in item.references for candidate in candidates)
candidate_recall += has_reference_candidate
effective_output = candidates[0] if candidates else item.input_text
is_correct = effective_output in item.references
accepted_candidate_miss_rows += bool(candidates) and not has_reference_candidate
selection_error_rows += bool(candidates) and has_reference_candidate and not is_correct
correct += is_correct
accepted_correct += bool(candidates) and is_correct
baseline_cer = minimum_cer(item.input_text, item.references)
effective_cer = minimum_cer(effective_output, item.references)
baseline_cers.append(baseline_cer)
cers.append(effective_cer)
if effective_cer < baseline_cer:
improved_rows += 1
elif effective_cer > baseline_cer:
worsened_rows += 1
else:
unchanged_rows += 1
if item.label == "clean":
clean_rows += 1
overcorrections += effective_output != item.input_text
elif item.label == "typo":
typo_rows += 1
typo_correct += is_correct
rows = len(items)
return {
"rows": rows,
"candidate_limit": candidate_limit,
"metrics": {
"effective_acc_at_1": correct / rows if rows else 0.0,
"candidate_recall_at_k": candidate_recall / rows if rows else 0.0,
"baseline_mean_min_cer": (
statistics.fmean(baseline_cers) if baseline_cers else 0.0
),
"mean_min_cer": statistics.fmean(cers) if cers else 0.0,
"abstention_rate": abstentions / rows if rows else 0.0,
"accepted_rows": accepted_rows,
"accepted_accuracy": (accepted_correct / accepted_rows if accepted_rows else None),
"clean_rows": clean_rows,
"overcorrection_rate": (overcorrections / clean_rows if clean_rows else None),
"typo_rows": typo_rows,
"typo_accuracy": typo_correct / typo_rows if typo_rows else None,
"improved_rows": improved_rows,
"worsened_rows": worsened_rows,
"unchanged_rows": unchanged_rows,
"candidate_miss_rows": rows - candidate_recall,
"accepted_candidate_miss_rows": accepted_candidate_miss_rows,
"selection_error_rows": selection_error_rows,
"declared_provenance_counts": dict(sorted(provenance_counts.items())),
"selection_reason_counts": dict(sorted(selection_reason_counts.items())),
"reported_margin_rows": len(reported_margins),
"mean_reported_margin": (
statistics.fmean(reported_margins) if reported_margins else None
),
},
}