File size: 8,655 Bytes
087643a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
# Finds the test pages the model got most wrong and renders them side by
# side with ground truth. The 5 failure cases in the memo come from this,
# not from guessing - the brief flags "no acknowledged failure cases" as a
# red flag, so these need to be real, pointable-at examples.
#
# Scoring is deliberately crude (greedy IoU 0.5 matching) - just needs to
# rank pages well enough to surface the interesting ones, not be a precise
# metric (that's evaluate.py's job).
#
# Usage:
#   python scripts/mine_failures.py --weights runs/detect/.../best.pt \
#       --data data/doclaynet --out reports/failures --top 25
from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from app.constants import ID_TO_CLASS  # noqa: E402
from scripts._render_utils import load_label_font as _label_font  # noqa: E402

PALETTE = [
    "#e6194b", "#3cb44b", "#ffe119", "#4363d8", "#f58231", "#911eb4",
    "#46f0f0", "#f032e6", "#bcf60c", "#fabebe", "#008080",
]


def iou(a: tuple[float, ...], b: tuple[float, ...]) -> float:
    """IoU for two xyxy boxes."""
    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)
    if ix2 <= ix1 or iy2 <= iy1:
        return 0.0

    intersection = (ix2 - ix1) * (iy2 - iy1)
    union = (ax2 - ax1) * (ay2 - ay1) + (bx2 - bx1) * (by2 - by1) - intersection
    return intersection / union if union > 0 else 0.0


def _read_ground_truth(label_path: Path, width: int, height: int) -> list[tuple[int, tuple[float, ...]]]:
    """Reads a YOLO label file back into absolute xyxy boxes."""
    if not label_path.exists():
        return []

    boxes = []
    for line in label_path.read_text(encoding="utf-8").splitlines():
        if not line.strip():
            continue
        class_id, xc, yc, w, h = line.split()
        xc, yc, w, h = float(xc) * width, float(yc) * height, float(w) * width, float(h) * height
        boxes.append((int(class_id), (xc - w / 2, yc - h / 2, xc + w / 2, yc + h / 2)))
    return boxes


def score_page(
    ground_truth: list[tuple[int, tuple[float, ...]]],
    predictions: list[tuple[int, tuple[float, ...], float]],
    iou_threshold: float = 0.5,
) -> dict:
    """Greedy match + counts. Splits misses from misclassifications since
    they have different causes - a miss is usually small/thin/query-
    starved, a misclassification is usually genuine label ambiguity
    (Text vs List-item, Title vs Section-header). That split is what lets
    the memo talk about annotation ambiguity instead of just "missed stuff"."""
    unmatched_gt = list(range(len(ground_truth)))
    matched_predictions: set[int] = set()

    correct = 0
    misclassified = 0

    for gt_index in list(unmatched_gt):
        gt_class, gt_box = ground_truth[gt_index]

        best_index, best_iou = None, iou_threshold
        for pred_index, (_, pred_box, _) in enumerate(predictions):
            if pred_index in matched_predictions:
                continue
            overlap = iou(gt_box, pred_box)
            if overlap >= best_iou:
                best_index, best_iou = pred_index, overlap

        if best_index is not None:
            matched_predictions.add(best_index)
            unmatched_gt.remove(gt_index)
            if predictions[best_index][0] == gt_class:
                correct += 1
            else:
                misclassified += 1

    false_positives = len(predictions) - len(matched_predictions)

    return {
        "ground_truth_regions": len(ground_truth),
        "predicted_regions": len(predictions),
        "correct": correct,
        "misclassified": misclassified,
        "missed": len(unmatched_gt),
        "false_positives": false_positives,
        # misclassifications weighted higher - they're the interesting ones
        "error_score": len(unmatched_gt) + false_positives + 1.5 * misclassified,
    }


def _draw_labelled_box(draw, colour: str, x1, y1, x2, y2, label: str, font) -> None:
    """Box outline plus a solid background behind the label, legible over
    any page content."""
    draw.rectangle([x1, y1, x2, y2], outline=colour, width=4)
    text_box = draw.textbbox((x1, y1), label, font=font)
    draw.rectangle(
        [text_box[0] - 2, text_box[1] - 2, text_box[2] + 2, text_box[3] + 2],
        fill=colour,
    )
    draw.text((x1, y1), label, font=font, fill="white")


def _render_comparison(image, ground_truth, predictions, out_path: Path) -> None:
    """Ground truth on the left, prediction on the right."""
    from PIL import Image, ImageDraw

    width, height = image.size
    canvas = Image.new("RGB", (width * 2 + 20, height + 40), "white")
    canvas.paste(image, (0, 40))
    canvas.paste(image, (width + 20, 40))

    draw = ImageDraw.Draw(canvas)
    header_font = _label_font(24)
    label_font = _label_font(22)
    draw.text((10, 8), "GROUND TRUTH", font=header_font, fill="black")
    draw.text((width + 30, 8), "PREDICTION", font=header_font, fill="black")

    for class_id, (x1, y1, x2, y2) in ground_truth:
        colour = PALETTE[class_id % len(PALETTE)]
        _draw_labelled_box(draw, colour, x1, y1 + 40, x2, y2 + 40,
                            ID_TO_CLASS.get(class_id, "?"), label_font)

    offset = width + 20
    for class_id, (x1, y1, x2, y2), confidence in predictions:
        colour = PALETTE[class_id % len(PALETTE)]
        label = f"{ID_TO_CLASS.get(class_id, '?')} {confidence:.2f}"
        _draw_labelled_box(draw, colour, x1 + offset, y1 + 40, x2 + offset, y2 + 40,
                            label, label_font)

    canvas.save(out_path)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--weights", required=True)
    parser.add_argument("--data", required=True, help="dataset root containing images/ and labels/")
    parser.add_argument("--out", default="reports/failures")
    parser.add_argument("--top", type=int, default=25)
    parser.add_argument("--conf", type=float, default=0.25)
    args = parser.parse_args()

    from PIL import Image
    from ultralytics import RTDETR

    data_root = Path(args.data)
    image_dir = data_root / "images" / "test"
    label_dir = data_root / "labels" / "test"
    out_dir = Path(args.out)
    out_dir.mkdir(parents=True, exist_ok=True)

    model = RTDETR(args.weights)
    images = sorted(image_dir.glob("*.png"))
    print(f"Scoring {len(images)} test pages...")

    manifest_path = data_root / "manifest_test.json"
    categories = {}
    if manifest_path.exists():
        categories = {
            row["stem"]: row["doc_category"]
            for row in json.loads(manifest_path.read_text(encoding="utf-8"))
        }

    scored = []
    for index, image_path in enumerate(images):
        image = Image.open(image_path).convert("RGB")
        width, height = image.size

        result = model.predict(image, conf=args.conf, verbose=False)[0]
        predictions = [
            (int(cls), tuple(float(v) for v in box), float(conf))
            for cls, box, conf in zip(
                result.boxes.cls.tolist(),
                result.boxes.xyxy.tolist(),
                result.boxes.conf.tolist(),
            )
        ]

        ground_truth = _read_ground_truth(label_dir / f"{image_path.stem}.txt", width, height)
        score = score_page(ground_truth, predictions)
        score["stem"] = image_path.stem
        score["doc_category"] = categories.get(image_path.stem, "unknown")
        scored.append((score, image, ground_truth, predictions))

        if (index + 1) % 100 == 0:
            print(f"  {index + 1}/{len(images)}", flush=True)

    scored.sort(key=lambda item: item[0]["error_score"], reverse=True)

    summary = []
    for rank, (score, image, ground_truth, predictions) in enumerate(scored[: args.top], start=1):
        out_path = out_dir / f"{rank:02d}_{score['stem']}_err{score['error_score']:.0f}.png"
        _render_comparison(image, ground_truth, predictions, out_path)
        score["render"] = out_path.name
        summary.append(score)

    (out_dir / "failure_summary.json").write_text(
        json.dumps(summary, indent=2), encoding="utf-8"
    )

    print(f"\nWrote {len(summary)} comparison renders to {out_dir}")
    print("\nWorst pages:")
    for score in summary[:8]:
        print(
            f"  {score['stem']:20s} {score['doc_category']:20s} "
            f"missed={score['missed']:3d} misclassified={score['misclassified']:3d} "
            f"fp={score['false_positives']:3d}"
        )


if __name__ == "__main__":
    main()