"""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, }