dev-strender's picture
Replace v24-era demo with v34 pipeline demo (engine-vendored bundle)
9c84f9d verified
Raw History Blame
8.9 kB
"""Scoring utilities ported from chosun-projects."""
def issue_count_to_score(count: int) -> float:
"""Convert issue count to 0-1 score. Used by title and image_caption eval."""
mapping = {0: 1.0, 1: 0.75, 2: 0.5, 3: 0.25}
return mapping.get(count, 0.0)
def error_count_to_severity(count: int) -> str:
"""Convert error count to 7-tier severity. Used by body eval."""
mapping = {0: "perfect", 1: "excellent", 2: "good", 3: "okay", 4: "fair", 5: "poor"}
return mapping.get(count, "severe")
def severity_to_score(severity: str) -> float:
"""Convert severity label to 0-1 score."""
mapping = {
"perfect": 1.0,
"excellent": 0.9,
"good": 0.7,
"okay": 0.5,
"fair": 0.3,
"poor": 0.1,
"severe": 0.0,
}
return mapping.get(severity, 0.0)
def score_to_severity(score: float, tiers: int = 5) -> str:
"""Convert 0-1 score to severity label."""
if tiers == 7:
if score >= 0.95:
return "perfect"
elif score >= 0.8:
return "excellent"
elif score >= 0.6:
return "good"
elif score >= 0.4:
return "okay"
elif score >= 0.2:
return "fair"
elif score >= 0.05:
return "poor"
else:
return "severe"
else: # tiers == 5
if score >= 0.9:
return "perfect"
elif score >= 0.7:
return "excellent"
elif score >= 0.4:
return "good"
elif score >= 0.2:
return "okay"
else:
return "poor"
def weighted_average(scores: dict[str, float], weights: dict[str, float]) -> float:
"""Compute weighted average of scores."""
total = sum(scores.get(k, 0.0) * w for k, w in weights.items())
return total
def compute_lcs(seq1: list[str], seq2: list[str]) -> list[tuple[int, int]]:
"""Compute Longest Common Subsequence. Returns list of (i, j) index pairs.
Ported from chosun-projects proofread/eval/scoring.py.
"""
m, n = len(seq1), len(seq2)
dp = [[0] * (n + 1) for _ in range(m + 1)]
for i in range(1, m + 1):
for j in range(1, n + 1):
if seq1[i - 1] == seq2[j - 1]:
dp[i][j] = dp[i - 1][j - 1] + 1
else:
dp[i][j] = max(dp[i - 1][j], dp[i][j - 1])
# Backtrack to find LCS indices
result = []
i, j = m, n
while i > 0 and j > 0:
if seq1[i - 1] == seq2[j - 1]:
result.append((i - 1, j - 1))
i -= 1
j -= 1
elif dp[i - 1][j] > dp[i][j - 1]:
i -= 1
else:
j -= 1
return list(reversed(result))
def find_diffs(
original_tokens: list[str],
corrected_tokens: list[str],
lcs_pairs: list[tuple[int, int]],
) -> list[dict]:
"""Find differences between original and corrected using LCS alignment.
Returns list of diffs with keys: orig_tokens, corrected_tokens,
orig_start, orig_end, corr_start, corr_end.
"""
diffs = []
prev_i, prev_j = -1, -1
for curr_i, curr_j in lcs_pairs:
if curr_i > prev_i + 1 or curr_j > prev_j + 1:
orig_slice = original_tokens[prev_i + 1 : curr_i]
corr_slice = corrected_tokens[prev_j + 1 : curr_j]
if orig_slice or corr_slice:
diffs.append({
"orig_tokens": orig_slice,
"corrected_tokens": corr_slice,
"orig_start": prev_i + 1,
"orig_end": curr_i,
"corr_start": prev_j + 1,
"corr_end": curr_j,
})
prev_i, prev_j = curr_i, curr_j
# Handle trailing differences
orig_tail = original_tokens[prev_i + 1 :]
corr_tail = corrected_tokens[prev_j + 1 :]
if orig_tail or corr_tail:
diffs.append({
"orig_tokens": orig_tail,
"corrected_tokens": corr_tail,
"orig_start": prev_i + 1,
"orig_end": len(original_tokens),
"corr_start": prev_j + 1,
"corr_end": len(corrected_tokens),
})
return diffs
def merge_adjacent_diffs(diffs: list[dict], max_gap: int = 2) -> list[dict]:
"""Merge adjacent diffs that are within max_gap tokens of each other."""
if not diffs:
return []
merged = [diffs[0]]
for d in diffs[1:]:
prev = merged[-1]
if d["orig_start"] - prev["orig_end"] <= max_gap:
prev["orig_tokens"].extend(d["orig_tokens"])
prev["corrected_tokens"].extend(d["corrected_tokens"])
prev["orig_end"] = d["orig_end"]
prev["corr_end"] = d["corr_end"]
else:
merged.append(d)
return merged
def compute_proofread_metrics(
original: str, golden: str, predicted: str
) -> dict[str, int | float]:
"""Compute proofread evaluation metrics using token-level LCS.
Returns: {tp, fp, missing, redundant, precision, recall}
"""
orig_tokens = original.split()
gold_tokens = golden.split()
pred_tokens = predicted.split()
# Find golden diffs (what should have been corrected)
gold_lcs = compute_lcs(orig_tokens, gold_tokens)
gold_diffs = find_diffs(orig_tokens, gold_tokens, gold_lcs)
gold_diffs = merge_adjacent_diffs(gold_diffs)
# Find predicted diffs (what was actually corrected)
pred_lcs = compute_lcs(orig_tokens, pred_tokens)
pred_diffs = find_diffs(orig_tokens, pred_tokens, pred_lcs)
pred_diffs = merge_adjacent_diffs(pred_diffs)
# Compare diffs using 2-pass matching
tp = 0
fp = 0
gold_by_start = {d["orig_start"]: d for d in gold_diffs}
pred_by_start = {d["orig_start"]: d for d in pred_diffs}
# 1st pass: exact orig_start position matching
matched_gold = set()
matched_pred = set()
for start in set(gold_by_start.keys()) & set(pred_by_start.keys()):
gold_corr = " ".join(gold_by_start[start]["corrected_tokens"])
pred_corr = " ".join(pred_by_start[start]["corrected_tokens"])
if gold_corr == pred_corr:
tp += 1
else:
fp += 1
matched_gold.add(start)
matched_pred.add(start)
# 2nd pass: content-based matching for unmatched diffs
# Handles cases where merge_adjacent_diffs grouped differently due to
# extra corrections shifting token positions.
# Match if golden's orig_tokens are a subset of (or equal to) pred's orig_tokens
# and the golden correction appears within the pred correction.
unmatched_gold = [d for d in gold_diffs if d["orig_start"] not in matched_gold]
unmatched_pred = [d for d in pred_diffs if d["orig_start"] not in matched_pred]
for g in list(unmatched_gold):
g_orig = " ".join(g["orig_tokens"])
g_corr = " ".join(g["corrected_tokens"])
candidates = []
for p in unmatched_pred:
p_orig = " ".join(p["orig_tokens"])
p_corr = " ".join(p["corrected_tokens"])
# Check if golden's original text is contained in pred's original text
# (merge may have grouped more tokens together in pred)
if g_orig in p_orig or p_orig in g_orig:
# Check if the golden correction appears in pred correction
if g_corr in p_corr:
candidates.append((p, True)) # correction matches
else:
candidates.append((p, False)) # same spot, different correction
if not candidates:
continue
# Prefer matching corrections, then closest position
candidates.sort(key=lambda x: (not x[1], abs(x[0]["orig_start"] - g["orig_start"])))
best, corr_matches = candidates[0]
if corr_matches:
tp += 1
else:
fp += 1
unmatched_gold.remove(g)
unmatched_pred.remove(best)
missing = len(unmatched_gold)
redundant = len(unmatched_pred)
# Compute precision and recall
# Default 100 only when nothing to correct AND nothing was corrected (perfect no-op)
# Default 0 when there were things to correct but model did nothing, or vice versa
no_golden = (tp + fp + missing) == 0 # nothing to correct
no_pred = (tp + fp + redundant) == 0 # model corrected nothing
if no_golden and no_pred:
precision, recall = 100.0, 100.0 # perfect: nothing to fix, nothing changed
else:
precision = (tp / (tp + fp + redundant) * 100) if not no_pred else 0.0
recall = (tp / (tp + fp + missing) * 100) if not no_golden else 0.0
f1 = (2 * precision * recall / (precision + recall)) if (precision + recall) > 0 else 0.0
return {
"tp": tp,
"fp": fp,
"missing": missing,
"redundant": redundant,
"precision": precision,
"recall": recall,
"f1": f1,
}