File size: 4,177 Bytes
76838d6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
# -*- coding: utf-8 -*-

"""
active_sampler.py
Selects uncertainty-based samples for (re-)labeling.
Input: val_predictions.json (or any predictions with fields path, boxes[{xyxy,cls,conf}])
Output: selection.json (list of samples with image path and rationale)

Heuristic:
    uncertainty = (1 - conf_calibrated) * size_factor
    size_factor = 1.0 for medium/large boxes, >1.0 for very small
Quota:
    max_per_class = K, global_max = N
"""

import argparse, json, math, sys
from pathlib import Path
from typing import Any, Dict, List, Tuple
import numpy as np

def read_json(p: Path):
    return json.loads(p.read_text(encoding="utf-8"))

def write_json(p: Path, obj: Any):
    p.parent.mkdir(parents=True, exist_ok=True)
    p.write_text(json.dumps(obj, ensure_ascii=False, indent=2), encoding="utf-8")

def temp_scale(p: float, T: float | None) -> float:
    if not T or T <= 0:
        return p
    p = min(max(p, 1e-8), 1 - 1e-8)
    z = math.log(p) - math.log(1 - p)
    return 1.0 / (1.0 + math.exp(-z / T))

def box_area_xyxy(box: List[float]) -> float:
    x1,y1,x2,y2 = box
    return max(0.0, x2-x1) * max(0.0, y2-y1)

def sample_candidates(preds: List[Dict[str,Any]],
                      names: List[str],
                      temperature: float | None,
                      small_area_px: float,
                      max_per_class: int,
                      global_max: int) -> List[Dict[str,Any]]:
    # Collect an uncertainty score per box
    candidates = []
    for r in preds:
        path = r.get("path")
        h, w = r.get("orig_shape", [None, None])
        boxes = r.get("boxes") or []
        for b in boxes:
            cls = int(b["cls"])
            conf = float(b["conf"])
            conf_cal = temp_scale(conf, temperature)
            unc = 1.0 - conf_cal

            # weight small boxes more
            area = box_area_xyxy(b["xyxy"])
            size_factor = 1.5 if area < small_area_px else 1.0
            score = unc * size_factor

            candidates.append({
                "path": path,
                "cls": cls,
                "label": names[cls] if 0 <= cls < len(names) else str(cls),
                "conf": conf,
                "conf_cal": conf_cal,
                "uncertainty": score,
                "area": area
            })

    # sort by uncertainty descending
    candidates.sort(key=lambda x: x["uncertainty"], reverse=True)

    # Quotas
    picked = []
    per_class = {i:0 for i in range(len(names))}
    for c in candidates:
        if len(picked) >= global_max:
            break
        ci = c["cls"]
        if per_class.get(ci, 0) >= max_per_class:
            continue
        picked.append(c)
        per_class[ci] = per_class.get(ci, 0) + 1
    return picked

def main():
    ap = argparse.ArgumentParser(description="Active Learning Sampler")
    ap.add_argument("--predictions", required=True, type=Path, help="Path to val_predictions.json or similar JSON")
    ap.add_argument("--names", required=True, type=Path, help="names.json (list of classes)")
    ap.add_argument("--calibration", type=Path, default=None, help="calibration_temp.json (optional)")
    ap.add_argument("--out", required=True, type=Path, help="Output selection.json")
    ap.add_argument("--small_area_px", type=float, default=32*32, help="Threshold for 'small' boxes (px^2)")
    ap.add_argument("--max_per_class", type=int, default=50)
    ap.add_argument("--global_max", type=int, default=300)
    args = ap.parse_args()

    preds = read_json(args.predictions)
    names = read_json(args.names)
    T = None
    if args.calibration and args.calibration.exists():
        try:
            T = float(read_json(args.calibration).get("temperature", None))
        except Exception:
            T = None

    picked = sample_candidates(preds, names, T, args.small_area_px, args.max_per_class, args.global_max)
    write_json(args.out, {
        "total": len(picked),
        "max_per_class": args.max_per_class,
        "global_max": args.global_max,
        "items": picked[:args.global_max]
    })
    print(f"selection -> {args.out} (n={len(picked)})")

if __name__ == "__main__":
    main()