Download active_learning.py from constructelligence/painting-vision-robotics-kit: direct link, hf CLI and curl.
- Browser
- Download file 9.19 kB
-
https://huggingface.co/constructelligence/painting-vision-robotics-kit/resolve/main/active_learning.py
- Command line
-
hf download hf://constructelligence/painting-vision-robotics-kit/active_learning.py
-
curl -L -o active_learning.py https://huggingface.co/constructelligence/painting-vision-robotics-kit/resolve/main/active_learning.py
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() | |