#!/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()