File size: 9,186 Bytes
c2457da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Rank unlabeled frames by how much they would improve the dataset.

Label time is the scarce resource in the flywheel. Random sampling wastes it on
frames the model already handles. This scores each candidate on three axes and
returns a ranked labelling queue:

* **novelty** - 1 minus its perceptual-hash similarity to anything already
  labelled, so near-duplicates of existing walls sink to the bottom;
* **domain need** - rare metadata values (weather, lighting, material, paint
  stage, ...) score higher, filling the balance gaps the audit reports;
* **uncertainty** - optional, from a trained checkpoint (mean predictive
  entropy), so frames the model is unsure about are prioritised.

The output is a queue with per-frame reasons, so a human can label the top N and
get the most coverage per minute. numpy + Pillow for scoring; uncertainty is the
only part that needs torch.

Usage::

    python3 active_learning.py --data data --candidates unlabeled/ --top 50 \
        --candidates-metadata unlabeled/metadata.csv --out queue.csv
    python3 active_learning.py --data data --candidates unlabeled/ --checkpoint artifacts/best.pt
"""
from __future__ import annotations

import argparse
import csv
import json
from collections import Counter
from pathlib import Path

from flywheel import build_index, dhash, hamming

ROOT = Path(__file__).resolve().parent
IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png"}
DEFAULT_FIELDS = ("environment", "surface_material", "lighting", "weather", "paint_stage")
DEFAULT_WEIGHTS = {"novelty": 0.5, "domain": 0.3, "uncertainty": 0.2}
HASH_BITS = 64


def labelled_counts(data_dir, fields=DEFAULT_FIELDS):
    """Per-field value counts from the labelled dataset's metadata."""
    path = Path(data_dir) / "metadata.csv"
    counts = {field: Counter() for field in fields}
    if not path.is_file():
        return counts
    with path.open(newline="", encoding="utf-8-sig") as stream:
        for row in csv.DictReader(stream):
            for field in fields:
                value = (row.get(field) or "").strip().lower()
                if value:
                    counts[field][value] += 1
    return counts


def metadata_rows(path):
    path = Path(path)
    if not path.is_file():
        return {}
    with path.open(newline="", encoding="utf-8-sig") as stream:
        return {Path(row.get("image", "")).name: row for row in csv.DictReader(stream)}


def novelty_score(fingerprint, existing_hashes):
    if not existing_hashes:
        return 1.0, None
    best = min(hamming(fingerprint, other) for other in existing_hashes)
    return best / HASH_BITS, best


def domain_score(row, counts, fields=DEFAULT_FIELDS):
    scores, reasons = [], []
    for field in fields:
        value = (row.get(field) or "").strip().lower()
        if not value:
            continue
        column = counts.get(field, Counter())
        peak = max(column.values()) if column else 1
        presence = column.get(value, 0)
        scores.append(1.0 - presence / peak)
        if presence <= peak * 0.25:
            reasons.append(f"rare {field}={value}")
    return (sum(scores) / len(scores) if scores else 0.0), reasons


def score_candidates(candidates, index, counts, rows=None, fields=DEFAULT_FIELDS, weights=None,
                     uncertainties=None, top=None):
    rows = rows or {}
    weights = {**DEFAULT_WEIGHTS, **(weights or {})}
    existing_hashes = list(index["phash"].keys())
    ranked = []
    for path in candidates:
        row = rows.get(Path(path).name, {})
        novelty, distance = novelty_score(dhash(path), existing_hashes)
        domain, reasons = domain_score(row, counts, fields)
        uncertainty = float(uncertainties.get(str(path), 0.0)) if uncertainties else None
        components = {"novelty": novelty, "domain": domain}
        active_weights = {k: weights[k] for k in ("novelty", "domain")}
        if uncertainty is not None:
            components["uncertainty"] = uncertainty
            active_weights["uncertainty"] = weights["uncertainty"]
        total = sum(active_weights.values()) or 1.0
        score = sum(components[k] * w for k, w in active_weights.items()) / total
        if novelty > 0.6:
            reasons = ["novel scene"] + reasons
        if uncertainty is not None and uncertainty > 0.3:
            reasons = ["high model uncertainty"] + reasons
        ranked.append({"image": str(path), "score": round(score, 4), "novelty": round(novelty, 4),
                       "nearest_hash_distance": distance, "domain": round(domain, 4),
                       "uncertainty": (round(uncertainty, 4) if uncertainty is not None else None),
                       "reasons": reasons})
    ranked.sort(key=lambda item: item["score"], reverse=True)
    return ranked[:top] if top else ranked


def uncertainty_for_images(images, checkpoint, size, device):
    """Mean predictive entropy per image from a Painting Vision checkpoint."""
    import numpy as np
    import torch

    from models import DEFAULT_ARCH, build_model

    state = torch.load(checkpoint, map_location=device, weights_only=True)
    model = build_model(str(state.get("arch", DEFAULT_ARCH)), len(state["classes"]), pretrained_backbone=False)
    model.load_state_dict(state["model"])
    model.to(device).eval()
    mean = np.asarray(state["mean"], np.float32)[None, None, :]
    std = np.asarray(state["std"], np.float32)[None, None, :]
    scores = {}
    from PIL import Image
    import cv2
    for path in images:
        bgr = cv2.imread(str(path))
        if bgr is None:
            continue
        rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)
        height, width = rgb.shape[:2]
        scale = size / max(height, width)
        resized = cv2.resize(rgb, (max(1, round(width * scale)), max(1, round(height * scale))))
        canvas = np.full((size, size, 3), 114, np.uint8)
        top, left = (size - resized.shape[0]) // 2, (size - resized.shape[1]) // 2
        canvas[top:top + resized.shape[0], left:left + resized.shape[1]] = resized
        tensor = torch.from_numpy(((canvas.astype(np.float32) / 255.0 - mean) / std)
                                  .transpose(2, 0, 1).copy()).unsqueeze(0)
        with torch.inference_mode():
            probabilities = torch.softmax(model(tensor.to(device))["semantic"], dim=1)[0].cpu().numpy()
        entropy = -np.sum(probabilities * np.log(probabilities + 1e-9), axis=0)
        scores[str(path)] = float(entropy.mean() / np.log(probabilities.shape[0]))
    return scores


def main():
    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("--data", type=Path, required=True, help="labelled dataset root")
    parser.add_argument("--candidates", type=Path, required=True, help="directory of unlabeled frames")
    parser.add_argument("--candidates-metadata", type=Path, help="optional CSV with domain columns per candidate")
    parser.add_argument("--fields", default=",".join(DEFAULT_FIELDS))
    parser.add_argument("--top", type=int, default=50)
    parser.add_argument("--checkpoint", type=Path, help="optional checkpoint for uncertainty scoring")
    parser.add_argument("--size", type=int, default=640)
    parser.add_argument("--device", default="cuda")
    parser.add_argument("--out", type=Path, help="write the ranked queue CSV here")
    args = parser.parse_args()

    fields = tuple(f.strip() for f in args.fields.split(",") if f.strip())
    rows = metadata_rows(args.candidates_metadata)

    candidates = sorted(p for p in args.candidates.iterdir() if p.suffix.lower() in IMAGE_EXTENSIONS)
    if not candidates:
        raise SystemExit(f"no images found in {args.candidates}")
    index = build_index(args.data)
    counts = labelled_counts(args.data, fields)
    uncertainties = None
    if args.checkpoint:
        import torch
        device = torch.device(args.device if torch.cuda.is_available() or args.device == "cpu" else "cpu")
        uncertainties = uncertainty_for_images(candidates, args.checkpoint, args.size, device)

    ranked = score_candidates(candidates, index, counts, rows=rows, fields=fields,
                              uncertainties=uncertainties, top=args.top)
    summary = {"labelled_images": len(index["sha256"]), "candidates": len(candidates),
               "queue": len(ranked), "domain_gaps": {field: dict(counts[field]) for field in fields}, "ranked": ranked}
    print(json.dumps({k: v for k, v in summary.items() if k != "ranked"}, indent=2))
    for item in ranked:
        print(f"{item['score']:.3f}  {item['image']}  {', '.join(item['reasons']) or '-'}")
    if args.out:
        args.out.parent.mkdir(parents=True, exist_ok=True)
        with args.out.open("w", newline="", encoding="utf-8") as stream:
            writer = csv.DictWriter(stream, fieldnames=["image", "score", "novelty", "nearest_hash_distance",
                                                        "domain", "uncertainty", "reasons"])
            writer.writeheader()
            for item in ranked:
                writer.writerow({**item, "reasons": "; ".join(item["reasons"])})


if __name__ == "__main__":
    main()