| |
| """Evaluate CartoLegendJSON predictions against JSONL gold rows.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import math |
| import re |
| from pathlib import Path |
| from statistics import mean, median |
| from typing import Any |
|
|
|
|
| FEATURE_TYPES = {"point", "line", "polygon", "area", "raster", "unknown"} |
| CENTER_THRESHOLDS = (16, 32, 64) |
|
|
|
|
| def bbox_iou(a: list[int] | None, b: list[int] | None) -> float: |
| if not a or not b: |
| return 0.0 |
| ax1, ay1, ax2, ay2 = a |
| bx1, by1, bx2, by2 = b |
| ix1, iy1 = max(ax1, bx1), max(ay1, by1) |
| ix2, iy2 = min(ax2, bx2), min(ay2, by2) |
| iw, ih = max(0, ix2 - ix1), max(0, iy2 - iy1) |
| inter = iw * ih |
| area_a = max(0, ax2 - ax1) * max(0, ay2 - ay1) |
| area_b = max(0, bx2 - bx1) * max(0, by2 - by1) |
| denom = area_a + area_b - inter |
| return inter / denom if denom else 0.0 |
|
|
|
|
| def bbox_center(bbox: list[int] | None) -> tuple[float, float] | None: |
| if not bbox or len(bbox) != 4: |
| return None |
| x1, y1, x2, y2 = bbox |
| return ((x1 + x2) / 2, (y1 + y2) / 2) |
|
|
|
|
| def bbox_center_distance(a: list[int] | None, b: list[int] | None) -> float | None: |
| ac = bbox_center(a) |
| bc = bbox_center(b) |
| if ac is None or bc is None: |
| return None |
| return math.hypot(ac[0] - bc[0], ac[1] - bc[1]) |
|
|
|
|
| def repeated_token_failure(text: str) -> bool: |
| return bool(re.search(r"([!?.#,\-_=])\1{12,}", text)) or bool(re.search(r"(\w{1,8})\1{20,}", text)) |
|
|
|
|
| def repair_json_text(text: str) -> str: |
| """Repair a few common near-JSON emissions without inventing content.""" |
| repaired = text.strip() |
| for key in ("symbol_bbox", "legend_bbox"): |
| repaired = re.sub(rf'"{key}=\[([0-9,\s]+)\]"', rf'"{key}":[\1]', repaired) |
| replacements = { |
| '"symbol_bbox=[': '"symbol_bbox":[', |
| '"legend_bbox=[': '"legend_bbox":[', |
| '"entries=[': '"entries":[', |
| } |
| for old, new in replacements.items(): |
| repaired = repaired.replace(old, new) |
| repaired = re.sub(r",\s*([}\]])", r"\1", repaired) |
| return repaired |
|
|
|
|
| def parse_json_text(text: str) -> dict[str, Any] | None: |
| try: |
| obj = json.loads(text) |
| except Exception: |
| return None |
| return obj if isinstance(obj, dict) else None |
|
|
|
|
| def parse_obj(text_or_obj: Any, *, repair: bool = False) -> tuple[dict[str, Any] | None, str, bool]: |
| if isinstance(text_or_obj, dict): |
| return text_or_obj, "", False |
| if not isinstance(text_or_obj, str): |
| return None, "prediction is not string/object", False |
| text = text_or_obj.strip() |
| obj = parse_json_text(text) |
| if obj is not None: |
| return obj, "", False |
|
|
| parse_error = "" |
| try: |
| json.loads(text) |
| except Exception as exc: |
| parse_error = str(exc) |
|
|
| |
| candidates = [] |
| start, end = text.find("{"), text.rfind("}") |
| if start >= 0 and end > start: |
| candidates.append(text[start : end + 1]) |
| if repair: |
| candidates.append(repair_json_text(text)) |
| if start >= 0 and end > start: |
| candidates.append(repair_json_text(text[start : end + 1])) |
| for candidate in candidates: |
| obj = parse_json_text(candidate) |
| if obj is not None: |
| return obj, "", candidate != text |
| return None, parse_error, False |
|
|
|
|
| def schema_errors(obj: Any) -> list[str]: |
| errors: list[str] = [] |
| if not isinstance(obj, dict): |
| return ["root not object"] |
| if not isinstance(obj.get("legend_detected"), bool): |
| errors.append("legend_detected not bool") |
| if not isinstance(obj.get("entries"), list): |
| errors.append("entries not list") |
| return errors |
| for i, entry in enumerate(obj["entries"]): |
| if not isinstance(entry, dict): |
| errors.append(f"entries[{i}] not object") |
| continue |
| if not isinstance(entry.get("label"), str) or not entry["label"]: |
| errors.append(f"entries[{i}].label invalid") |
| if entry.get("feature_type") not in FEATURE_TYPES: |
| errors.append(f"entries[{i}].feature_type invalid") |
| bbox = entry.get("symbol_bbox") |
| if bbox is not None: |
| if not isinstance(bbox, list) or len(bbox) != 4 or not all(isinstance(x, int) and x >= 0 for x in bbox): |
| errors.append(f"entries[{i}].symbol_bbox invalid") |
| return errors |
|
|
|
|
| def entry_key(entry: dict[str, Any]) -> str: |
| return str(entry.get("label") or entry.get("code") or "") |
|
|
|
|
| def normalize_label_key(label: str) -> str: |
| """Normalize labels for semantic-ish matching while preserving exact metrics.""" |
| text = str(label).lower().replace("&", " and ") |
| |
| |
| text = re.sub(r"\biron\s*formation(s?)\b", r"iron formation\1", text) |
| text = re.sub(r"\bsulphid(e|es|ic)\b", r"sulfid\1", text) |
| text = text.replace(".", "").replace("'", "") |
| tokens = re.findall(r"[a-z0-9]+", text) |
|
|
| normalized_tokens: list[str] = [] |
| initials: list[str] = [] |
| for token in tokens: |
| if len(token) == 1 and token.isalpha(): |
| initials.append(token) |
| continue |
| if initials: |
| normalized_tokens.append("".join(initials)) |
| initials = [] |
| normalized_tokens.append(token) |
| if initials: |
| normalized_tokens.append("".join(initials)) |
| return " ".join(normalized_tokens) |
|
|
|
|
| def normalized_entry_key(entry: dict[str, Any]) -> str: |
| return normalize_label_key(entry_key(entry)) |
|
|
|
|
| def gold_from_row(row: dict[str, Any]) -> dict[str, Any]: |
| return json.loads(row["messages"][-1]["content"]) |
|
|
|
|
| def pred_from_row(row: dict[str, Any]) -> Any: |
| if "prediction" in row: |
| return row["prediction"] |
| if "messages" in row: |
| return row["messages"][-1]["content"] |
| return row |
|
|
|
|
| def main() -> int: |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--gold", type=Path, required=True) |
| parser.add_argument("--pred", type=Path, required=True) |
| parser.add_argument("--out", type=Path, default=None) |
| parser.add_argument( |
| "--repair-json", |
| action="store_true", |
| help="Apply narrow syntax repairs before schema validation and report repair counts.", |
| ) |
| args = parser.parse_args() |
|
|
| gold_rows = [json.loads(line) for line in args.gold.read_text().splitlines() if line.strip()] |
| pred_rows = [json.loads(line) for line in args.pred.read_text().splitlines() if line.strip()] |
| n = min(len(gold_rows), len(pred_rows)) |
|
|
| valid = 0 |
| repeated = 0 |
| entry_count_close = 0 |
| type_correct = 0 |
| type_total = 0 |
| matched_total = 0 |
| gold_entry_total_all = 0 |
| empty_gold_rows_all = 0 |
| gold_entry_total = 0 |
| prediction_entry_total = 0 |
| empty_gold_rows = 0 |
| empty_gold_false_positive_rows = 0 |
| nonempty_gold_empty_prediction_rows = 0 |
| ious: list[float] = [] |
| center_distances: list[float] = [] |
| center_hits = {threshold: 0 for threshold in CENTER_THRESHOLDS} |
| normalized_type_correct = 0 |
| normalized_type_total = 0 |
| normalized_matched_total = 0 |
| normalized_recovered_exact_label_misses = 0 |
| normalized_ious: list[float] = [] |
| normalized_center_distances: list[float] = [] |
| normalized_center_hits = {threshold: 0 for threshold in CENTER_THRESHOLDS} |
| order_center_distances: list[float] = [] |
| order_center_hits = {threshold: 0 for threshold in CENTER_THRESHOLDS} |
| order_aligned_total = 0 |
| repaired_predictions = 0 |
| invalid_examples = [] |
|
|
| for i in range(n): |
| gold = gold_from_row(gold_rows[i]) |
| gold_entries_for_denominator = gold.get("entries", []) |
| gold_entry_total_all += len(gold_entries_for_denominator) |
| if not gold_entries_for_denominator: |
| empty_gold_rows_all += 1 |
| pred_raw = pred_from_row(pred_rows[i]) |
| pred_text = pred_raw if isinstance(pred_raw, str) else json.dumps(pred_raw) |
| if repeated_token_failure(pred_text): |
| repeated += 1 |
| pred, parse_error, repaired = parse_obj(pred_raw, repair=args.repair_json) |
| if pred is None: |
| invalid_examples.append({"index": i, "error": parse_error, "prediction": pred_text[:300]}) |
| continue |
| if repaired: |
| repaired_predictions += 1 |
| errors = schema_errors(pred) |
| if errors: |
| invalid_examples.append({"index": i, "error": "; ".join(errors), "prediction": pred_text[:300]}) |
| continue |
| valid += 1 |
|
|
| gold_entries = gold.get("entries", []) |
| pred_entries = pred.get("entries", []) |
| gold_entry_total += len(gold_entries) |
| prediction_entry_total += len(pred_entries) |
| if not gold_entries: |
| empty_gold_rows += 1 |
| if pred_entries: |
| empty_gold_false_positive_rows += 1 |
| elif not pred_entries: |
| nonempty_gold_empty_prediction_rows += 1 |
| if abs(len(gold_entries) - len(pred_entries)) <= 1: |
| entry_count_close += 1 |
|
|
| for g, p in zip(gold_entries, pred_entries): |
| dist = bbox_center_distance(g.get("symbol_bbox"), p.get("symbol_bbox")) |
| if dist is None: |
| continue |
| order_aligned_total += 1 |
| order_center_distances.append(dist) |
| for threshold in CENTER_THRESHOLDS: |
| if dist <= threshold: |
| order_center_hits[threshold] += 1 |
|
|
| pred_by_label = {entry_key(e): e for e in pred_entries if entry_key(e)} |
| pred_by_normalized_label: dict[str, list[dict[str, Any]]] = {} |
| for pred_entry in pred_entries: |
| normalized_key = normalized_entry_key(pred_entry) |
| if normalized_key: |
| pred_by_normalized_label.setdefault(normalized_key, []).append(pred_entry) |
|
|
| for g in gold_entries: |
| p = pred_by_label.get(entry_key(g)) |
| if not p: |
| continue |
| matched_total += 1 |
| ious.append(bbox_iou(g.get("symbol_bbox"), p.get("symbol_bbox"))) |
| dist = bbox_center_distance(g.get("symbol_bbox"), p.get("symbol_bbox")) |
| if dist is not None: |
| center_distances.append(dist) |
| for threshold in CENTER_THRESHOLDS: |
| if dist <= threshold: |
| center_hits[threshold] += 1 |
| type_total += 1 |
| if g.get("feature_type") == p.get("feature_type"): |
| type_correct += 1 |
|
|
| for g in gold_entries: |
| normalized_key = normalized_entry_key(g) |
| if not normalized_key: |
| continue |
| candidates = pred_by_normalized_label.get(normalized_key) |
| if not candidates: |
| continue |
| p = candidates.pop(0) |
| normalized_matched_total += 1 |
| if entry_key(g) != entry_key(p): |
| normalized_recovered_exact_label_misses += 1 |
| normalized_ious.append(bbox_iou(g.get("symbol_bbox"), p.get("symbol_bbox"))) |
| dist = bbox_center_distance(g.get("symbol_bbox"), p.get("symbol_bbox")) |
| if dist is not None: |
| normalized_center_distances.append(dist) |
| for threshold in CENTER_THRESHOLDS: |
| if dist <= threshold: |
| normalized_center_hits[threshold] += 1 |
| normalized_type_total += 1 |
| if g.get("feature_type") == p.get("feature_type"): |
| normalized_type_correct += 1 |
|
|
| result = { |
| "n": n, |
| "valid_json_rate": valid / n if n else 0.0, |
| "repeated_token_rate": repeated / n if n else 0.0, |
| "entry_count_close_rate": entry_count_close / n if n else 0.0, |
| "matched_entries": matched_total, |
| "gold_entries_total": gold_entry_total_all, |
| "label_recall": matched_total / gold_entry_total_all if gold_entry_total_all else 0.0, |
| "gold_entries_in_valid_predictions": gold_entry_total, |
| "prediction_entries_in_valid_predictions": prediction_entry_total, |
| "label_recall_in_valid_predictions": matched_total / gold_entry_total if gold_entry_total else 0.0, |
| "label_precision_in_valid_predictions": matched_total / prediction_entry_total if prediction_entry_total else 0.0, |
| "empty_gold_rows_total": empty_gold_rows_all, |
| "empty_gold_rows_in_valid_predictions": empty_gold_rows, |
| "empty_gold_false_positive_rows": empty_gold_false_positive_rows, |
| "empty_gold_false_positive_rate": empty_gold_false_positive_rows / empty_gold_rows if empty_gold_rows else 0.0, |
| "nonempty_gold_empty_prediction_rows": nonempty_gold_empty_prediction_rows, |
| "feature_type_accuracy": type_correct / type_total if type_total else 0.0, |
| "repaired_prediction_count": repaired_predictions, |
| "repaired_prediction_rate": repaired_predictions / n if n else 0.0, |
| "mean_symbol_bbox_iou": mean(ious) if ious else 0.0, |
| "median_symbol_bbox_iou": median(ious) if ious else 0.0, |
| "mean_symbol_center_distance_px": mean(center_distances) if center_distances else None, |
| "median_symbol_center_distance_px": median(center_distances) if center_distances else None, |
| "symbol_center_hit_rate_16px": center_hits[16] / matched_total if matched_total else 0.0, |
| "symbol_center_hit_rate_32px": center_hits[32] / matched_total if matched_total else 0.0, |
| "symbol_center_hit_rate_64px": center_hits[64] / matched_total if matched_total else 0.0, |
| "normalized_matched_entries": normalized_matched_total, |
| "normalized_recovered_exact_label_misses": normalized_recovered_exact_label_misses, |
| "normalized_label_recall": normalized_matched_total / gold_entry_total_all if gold_entry_total_all else 0.0, |
| "normalized_label_recall_in_valid_predictions": ( |
| normalized_matched_total / gold_entry_total if gold_entry_total else 0.0 |
| ), |
| "normalized_label_precision_in_valid_predictions": ( |
| normalized_matched_total / prediction_entry_total if prediction_entry_total else 0.0 |
| ), |
| "normalized_feature_type_accuracy": ( |
| normalized_type_correct / normalized_type_total if normalized_type_total else 0.0 |
| ), |
| "normalized_mean_symbol_bbox_iou": mean(normalized_ious) if normalized_ious else 0.0, |
| "normalized_median_symbol_bbox_iou": median(normalized_ious) if normalized_ious else 0.0, |
| "normalized_mean_symbol_center_distance_px": ( |
| mean(normalized_center_distances) if normalized_center_distances else None |
| ), |
| "normalized_median_symbol_center_distance_px": ( |
| median(normalized_center_distances) if normalized_center_distances else None |
| ), |
| "normalized_symbol_center_hit_rate_16px": ( |
| normalized_center_hits[16] / normalized_matched_total if normalized_matched_total else 0.0 |
| ), |
| "normalized_symbol_center_hit_rate_32px": ( |
| normalized_center_hits[32] / normalized_matched_total if normalized_matched_total else 0.0 |
| ), |
| "normalized_symbol_center_hit_rate_64px": ( |
| normalized_center_hits[64] / normalized_matched_total if normalized_matched_total else 0.0 |
| ), |
| "order_aligned_entries": order_aligned_total, |
| "order_aligned_mean_center_distance_px": mean(order_center_distances) if order_center_distances else None, |
| "order_aligned_median_center_distance_px": median(order_center_distances) if order_center_distances else None, |
| "order_aligned_center_hit_rate_16px": order_center_hits[16] / order_aligned_total if order_aligned_total else 0.0, |
| "order_aligned_center_hit_rate_32px": order_center_hits[32] / order_aligned_total if order_aligned_total else 0.0, |
| "order_aligned_center_hit_rate_64px": order_center_hits[64] / order_aligned_total if order_aligned_total else 0.0, |
| "invalid_examples": invalid_examples[:10], |
| } |
|
|
| text = json.dumps(result, indent=2) |
| print(text) |
| if args.out: |
| args.out.parent.mkdir(parents=True, exist_ok=True) |
| args.out.write_text(text + "\n") |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|