File size: 7,684 Bytes
1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 20a4178 1f7b895 | 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 | #!/usr/bin/env python3
"""Run the two-stage model over a dataset and write predictions for eval.py.
python3 predict.py --labels clockface-external/labels.jsonl \
--root clockface-external --stage2 checkpoints/v2mnv3/stage2_best.pt \
--out preds.jsonl
Each prediction carries `agreement_minutes`: how far apart the hour hand's
reading and the minute hand's reading are. On a real clock they agree, so
disagreement is a confidence signal that costs nothing to produce.
With --stage1 the dial is located first and the crop comes from that. Without
it, the whole image is used, which is what you want only when the clock already
fills the frame.
"""
from __future__ import annotations
import argparse
import json
import os
import numpy as np
import torch
from PIL import Image
from jointdecode import joint_decode
from model import ClockNetCls, decode_cls
from twostage import CocoDialDetector, DialLocator, PretrainedReader, crop_dial
MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)
STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)
def to_tensor(img):
x = torch.from_numpy(np.asarray(img, dtype=np.float32).copy() / 255.0).permute(2, 0, 1)
return (x - MEAN) / STD
def build_reader(ckpt, device):
blob = torch.load(ckpt, map_location="cpu", weights_only=False)
a = blob["args"]
backbone = a.get("backbone", "scratch")
# checkpoints written before the whole-time head have no weights for it,
# so take the head's size from the file rather than from today's default
tb = blob["model"].get("time_head.weight")
model = (ClockNetCls(width=a.get("width", 32), bins=a.get("bins", 180))
if backbone == "scratch" else
PretrainedReader(bins=a.get("bins", 180), arch=backbone,
time_bins=0 if tb is None else tb.shape[0]))
model.load_state_dict(blob["model"])
return model.to(device).eval(), a.get("res", 256)
def main():
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--labels", required=True)
ap.add_argument("--root", required=True, help="directory the 'file' fields are relative to")
ap.add_argument("--stage2", required=True)
ap.add_argument("--stage1", help="a trained DialLocator checkpoint (deprecated)")
ap.add_argument("--coco", action="store_true",
help="locate the dial with a COCO-pretrained detector. Measured at "
"MAE 73.6 min against 120.7 for the trained locator on the same "
"200 real photographs.")
ap.add_argument("--coco-arch", default="mobilenet", choices=["mobilenet", "resnet50"])
ap.add_argument("--use-gt-dial", action="store_true",
help="crop with the label's own dial geometry (synthetic only)")
ap.add_argument("--margin", type=float, default=1.25)
ap.add_argument("--batch", type=int, default=32)
ap.add_argument("--out", required=True)
ap.add_argument("--decode", choices=["joint", "independent"], default="joint",
help="joint scores every time the clock could show and keeps the "
"best; independent reads each head alone (the old behaviour)")
ap.add_argument("--dump-logits", help="write raw head distributions here, so "
"decoders can be compared without re-running the network")
ap.add_argument("--time-weight", type=float, default=1.0,
help="how loudly the whole-time head votes in the joint decode")
ap.add_argument("--cpu", action="store_true",
help="stay off the GPU, so a prediction run can share the "
"machine with a training run")
args = ap.parse_args()
logit_dump = open(args.dump_logits, "w") if args.dump_logits else None
device = torch.device("cpu" if args.cpu else
"mps" if torch.backends.mps.is_available() else "cpu")
reader, res = build_reader(args.stage2, device)
coco = CocoDialDetector(args.coco_arch, device=device) if args.coco else None
locator = None
if args.stage1:
blob = torch.load(args.stage1, map_location="cpu", weights_only=False)
locator = DialLocator().to(device).eval()
locator.load_state_dict(blob["model"])
loc_res = blob["args"].get("res", 256)
rows = [json.loads(l) for l in open(args.labels) if l.strip()]
out = open(args.out, "w")
n = 0
with torch.no_grad():
for i in range(0, len(rows), args.batch):
chunk = rows[i:i + args.batch]
crops, ids = [], []
for r in chunk:
path = os.path.join(args.root, r["file"])
if not os.path.exists(path):
continue
img = Image.open(path).convert("RGB")
if coco is not None:
got = coco.locate(img)
crop = (crop_dial(img, got[0], got[1], got[2], res, 1.15)
if got else img.resize((res, res)))
elif args.use_gt_dial and r.get("render", {}).get("dial"):
d = r["render"]["dial"]
crop = crop_dial(img, d["cx"], d["cy"], d["r_max"], res, args.margin)
elif locator is not None:
small = to_tensor(img.resize((loc_res, loc_res))).unsqueeze(0).to(device)
p = locator(small)[0].float().cpu()
cx, cy, r_ = p[0].item(), p[1].item(), float(np.exp(p[2].item()))
crop = crop_dial(img, cx, cy, r_, res, args.margin)
else:
crop = img.resize((res, res))
crops.append(to_tensor(crop))
ids.append(r["id"])
if not crops:
continue
x = torch.stack(crops).to(device)
out_heads = reader(x)
hl, ml, tl = (out_heads if len(out_heads) == 3
else (out_heads[0], out_heads[1], None))
hl, ml = hl.float().cpu(), ml.float().cpu()
tl = tl.float().cpu() if tl is not None else None
t, hour_only, dis, conf = decode_cls(hl, ml)
if args.decode == "joint":
t, margin, swapped, consistency = joint_decode(hl, ml, tl,
args.time_weight)
if logit_dump is not None:
for j, rid in enumerate(ids):
logit_dump.write(json.dumps({"id": rid,
"hour": [round(v, 4) for v in hl[j].tolist()],
"minute": [round(v, 4) for v in ml[j].tolist()]}) + "\n")
for j, rid in enumerate(ids):
hh = int(t[j].item() // 60) or 12
mm = t[j].item() - (t[j].item() // 60) * 60
out.write(json.dumps({
"id": rid,
"time": f"{hh}:{int(round(mm)) % 60:02d}",
"minutes": round(t[j].item(), 3),
"agreement_minutes": round(dis[j].item(), 3),
"sharpness": round(float(conf[j].mean()), 4),
**({"margin": round(margin[j].item(), 4),
"hands_swapped": bool(swapped[j]),
"consistency": round(consistency[j].item(), 4)}
if args.decode == "joint" else {}),
}) + "\n")
n += 1
out.close()
if logit_dump:
logit_dump.close()
print(f"wrote {n} predictions -> {args.out}")
if __name__ == "__main__":
main()
|