lyte-codes commited on
Commit
20a4178
·
verified ·
1 Parent(s): 24e3898

Sync predict.py for 263702

Browse files
Files changed (1) hide show
  1. 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("mps" if torch.backends.mps.is_available() else "cpu")
 
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 args.use_gt_dial and r.get("render", {}).get("dial"):
 
 
 
 
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
- hl, ml = reader(x)
101
- t, hour_only, dis, conf = decode_cls(hl.float().cpu(), ml.float().cpu())
 
 
 
 
 
 
 
 
 
 
 
 
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