painting-vision-robotics-kit / active_learning.py
constructelligence's picture
Upload active_learning.py with huggingface_hub
c2457da verified
Raw History Blame Contribute Delete
9.19 kB
#!/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()