Sync predict.py for 263702
Browse files- predict.py +52 -7
predict.py
CHANGED
|
@@ -24,8 +24,9 @@ import numpy as np
|
|
| 24 |
import torch
|
| 25 |
from PIL import Image
|
| 26 |
|
|
|
|
| 27 |
from model import ClockNetCls, decode_cls
|
| 28 |
-
from twostage import DialLocator, PretrainedReader, crop_dial
|
| 29 |
|
| 30 |
MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)
|
| 31 |
STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)
|
|
@@ -40,9 +41,13 @@ def build_reader(ckpt, device):
|
|
| 40 |
blob = torch.load(ckpt, map_location="cpu", weights_only=False)
|
| 41 |
a = blob["args"]
|
| 42 |
backbone = a.get("backbone", "scratch")
|
|
|
|
|
|
|
|
|
|
| 43 |
model = (ClockNetCls(width=a.get("width", 32), bins=a.get("bins", 180))
|
| 44 |
if backbone == "scratch" else
|
| 45 |
-
PretrainedReader(bins=a.get("bins", 180), arch=backbone
|
|
|
|
| 46 |
model.load_state_dict(blob["model"])
|
| 47 |
return model.to(device).eval(), a.get("res", 256)
|
| 48 |
|
|
@@ -53,16 +58,34 @@ def main():
|
|
| 53 |
ap.add_argument("--labels", required=True)
|
| 54 |
ap.add_argument("--root", required=True, help="directory the 'file' fields are relative to")
|
| 55 |
ap.add_argument("--stage2", required=True)
|
| 56 |
-
ap.add_argument("--stage1")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
ap.add_argument("--use-gt-dial", action="store_true",
|
| 58 |
help="crop with the label's own dial geometry (synthetic only)")
|
| 59 |
ap.add_argument("--margin", type=float, default=1.25)
|
| 60 |
ap.add_argument("--batch", type=int, default=32)
|
| 61 |
ap.add_argument("--out", required=True)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 62 |
args = ap.parse_args()
|
|
|
|
| 63 |
|
| 64 |
-
device = torch.device("
|
|
|
|
| 65 |
reader, res = build_reader(args.stage2, device)
|
|
|
|
| 66 |
locator = None
|
| 67 |
if args.stage1:
|
| 68 |
blob = torch.load(args.stage1, map_location="cpu", weights_only=False)
|
|
@@ -82,7 +105,11 @@ def main():
|
|
| 82 |
if not os.path.exists(path):
|
| 83 |
continue
|
| 84 |
img = Image.open(path).convert("RGB")
|
| 85 |
-
if
|
|
|
|
|
|
|
|
|
|
|
|
|
| 86 |
d = r["render"]["dial"]
|
| 87 |
crop = crop_dial(img, d["cx"], d["cy"], d["r_max"], res, args.margin)
|
| 88 |
elif locator is not None:
|
|
@@ -97,8 +124,20 @@ def main():
|
|
| 97 |
if not crops:
|
| 98 |
continue
|
| 99 |
x = torch.stack(crops).to(device)
|
| 100 |
-
|
| 101 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 102 |
for j, rid in enumerate(ids):
|
| 103 |
hh = int(t[j].item() // 60) or 12
|
| 104 |
mm = t[j].item() - (t[j].item() // 60) * 60
|
|
@@ -108,9 +147,15 @@ def main():
|
|
| 108 |
"minutes": round(t[j].item(), 3),
|
| 109 |
"agreement_minutes": round(dis[j].item(), 3),
|
| 110 |
"sharpness": round(float(conf[j].mean()), 4),
|
|
|
|
|
|
|
|
|
|
|
|
|
| 111 |
}) + "\n")
|
| 112 |
n += 1
|
| 113 |
out.close()
|
|
|
|
|
|
|
| 114 |
print(f"wrote {n} predictions -> {args.out}")
|
| 115 |
|
| 116 |
|
|
|
|
| 24 |
import torch
|
| 25 |
from PIL import Image
|
| 26 |
|
| 27 |
+
from jointdecode import joint_decode
|
| 28 |
from model import ClockNetCls, decode_cls
|
| 29 |
+
from twostage import CocoDialDetector, DialLocator, PretrainedReader, crop_dial
|
| 30 |
|
| 31 |
MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)
|
| 32 |
STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)
|
|
|
|
| 41 |
blob = torch.load(ckpt, map_location="cpu", weights_only=False)
|
| 42 |
a = blob["args"]
|
| 43 |
backbone = a.get("backbone", "scratch")
|
| 44 |
+
# checkpoints written before the whole-time head have no weights for it,
|
| 45 |
+
# so take the head's size from the file rather than from today's default
|
| 46 |
+
tb = blob["model"].get("time_head.weight")
|
| 47 |
model = (ClockNetCls(width=a.get("width", 32), bins=a.get("bins", 180))
|
| 48 |
if backbone == "scratch" else
|
| 49 |
+
PretrainedReader(bins=a.get("bins", 180), arch=backbone,
|
| 50 |
+
time_bins=0 if tb is None else tb.shape[0]))
|
| 51 |
model.load_state_dict(blob["model"])
|
| 52 |
return model.to(device).eval(), a.get("res", 256)
|
| 53 |
|
|
|
|
| 58 |
ap.add_argument("--labels", required=True)
|
| 59 |
ap.add_argument("--root", required=True, help="directory the 'file' fields are relative to")
|
| 60 |
ap.add_argument("--stage2", required=True)
|
| 61 |
+
ap.add_argument("--stage1", help="a trained DialLocator checkpoint (deprecated)")
|
| 62 |
+
ap.add_argument("--coco", action="store_true",
|
| 63 |
+
help="locate the dial with a COCO-pretrained detector. Measured at "
|
| 64 |
+
"MAE 73.6 min against 120.7 for the trained locator on the same "
|
| 65 |
+
"200 real photographs.")
|
| 66 |
+
ap.add_argument("--coco-arch", default="mobilenet", choices=["mobilenet", "resnet50"])
|
| 67 |
ap.add_argument("--use-gt-dial", action="store_true",
|
| 68 |
help="crop with the label's own dial geometry (synthetic only)")
|
| 69 |
ap.add_argument("--margin", type=float, default=1.25)
|
| 70 |
ap.add_argument("--batch", type=int, default=32)
|
| 71 |
ap.add_argument("--out", required=True)
|
| 72 |
+
ap.add_argument("--decode", choices=["joint", "independent"], default="joint",
|
| 73 |
+
help="joint scores every time the clock could show and keeps the "
|
| 74 |
+
"best; independent reads each head alone (the old behaviour)")
|
| 75 |
+
ap.add_argument("--dump-logits", help="write raw head distributions here, so "
|
| 76 |
+
"decoders can be compared without re-running the network")
|
| 77 |
+
ap.add_argument("--time-weight", type=float, default=1.0,
|
| 78 |
+
help="how loudly the whole-time head votes in the joint decode")
|
| 79 |
+
ap.add_argument("--cpu", action="store_true",
|
| 80 |
+
help="stay off the GPU, so a prediction run can share the "
|
| 81 |
+
"machine with a training run")
|
| 82 |
args = ap.parse_args()
|
| 83 |
+
logit_dump = open(args.dump_logits, "w") if args.dump_logits else None
|
| 84 |
|
| 85 |
+
device = torch.device("cpu" if args.cpu else
|
| 86 |
+
"mps" if torch.backends.mps.is_available() else "cpu")
|
| 87 |
reader, res = build_reader(args.stage2, device)
|
| 88 |
+
coco = CocoDialDetector(args.coco_arch, device=device) if args.coco else None
|
| 89 |
locator = None
|
| 90 |
if args.stage1:
|
| 91 |
blob = torch.load(args.stage1, map_location="cpu", weights_only=False)
|
|
|
|
| 105 |
if not os.path.exists(path):
|
| 106 |
continue
|
| 107 |
img = Image.open(path).convert("RGB")
|
| 108 |
+
if coco is not None:
|
| 109 |
+
got = coco.locate(img)
|
| 110 |
+
crop = (crop_dial(img, got[0], got[1], got[2], res, 1.15)
|
| 111 |
+
if got else img.resize((res, res)))
|
| 112 |
+
elif args.use_gt_dial and r.get("render", {}).get("dial"):
|
| 113 |
d = r["render"]["dial"]
|
| 114 |
crop = crop_dial(img, d["cx"], d["cy"], d["r_max"], res, args.margin)
|
| 115 |
elif locator is not None:
|
|
|
|
| 124 |
if not crops:
|
| 125 |
continue
|
| 126 |
x = torch.stack(crops).to(device)
|
| 127 |
+
out_heads = reader(x)
|
| 128 |
+
hl, ml, tl = (out_heads if len(out_heads) == 3
|
| 129 |
+
else (out_heads[0], out_heads[1], None))
|
| 130 |
+
hl, ml = hl.float().cpu(), ml.float().cpu()
|
| 131 |
+
tl = tl.float().cpu() if tl is not None else None
|
| 132 |
+
t, hour_only, dis, conf = decode_cls(hl, ml)
|
| 133 |
+
if args.decode == "joint":
|
| 134 |
+
t, margin, swapped, consistency = joint_decode(hl, ml, tl,
|
| 135 |
+
args.time_weight)
|
| 136 |
+
if logit_dump is not None:
|
| 137 |
+
for j, rid in enumerate(ids):
|
| 138 |
+
logit_dump.write(json.dumps({"id": rid,
|
| 139 |
+
"hour": [round(v, 4) for v in hl[j].tolist()],
|
| 140 |
+
"minute": [round(v, 4) for v in ml[j].tolist()]}) + "\n")
|
| 141 |
for j, rid in enumerate(ids):
|
| 142 |
hh = int(t[j].item() // 60) or 12
|
| 143 |
mm = t[j].item() - (t[j].item() // 60) * 60
|
|
|
|
| 147 |
"minutes": round(t[j].item(), 3),
|
| 148 |
"agreement_minutes": round(dis[j].item(), 3),
|
| 149 |
"sharpness": round(float(conf[j].mean()), 4),
|
| 150 |
+
**({"margin": round(margin[j].item(), 4),
|
| 151 |
+
"hands_swapped": bool(swapped[j]),
|
| 152 |
+
"consistency": round(consistency[j].item(), 4)}
|
| 153 |
+
if args.decode == "joint" else {}),
|
| 154 |
}) + "\n")
|
| 155 |
n += 1
|
| 156 |
out.close()
|
| 157 |
+
if logit_dump:
|
| 158 |
+
logit_dump.close()
|
| 159 |
print(f"wrote {n} predictions -> {args.out}")
|
| 160 |
|
| 161 |
|