Download predict.py from lyte-codes/clockface: direct link, hf CLI and curl.
- Browser
- Download file 4.74 kB
-
https://huggingface.co/lyte-codes/clockface/resolve/98f945c8702fd79792ed3cbc1718052d3cb12fd2/predict.py
- Command line
-
hf download hf://lyte-codes/clockface@98f945c8702fd79792ed3cbc1718052d3cb12fd2/predict.py
-
curl -L -o predict.py https://huggingface.co/lyte-codes/clockface/resolve/98f945c8702fd79792ed3cbc1718052d3cb12fd2/predict.py
4.74 kB
| #!/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 model import ClockNetCls, decode_cls | |
| from twostage import 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") | |
| model = (ClockNetCls(width=a.get("width", 32), bins=a.get("bins", 180)) | |
| if backbone == "scratch" else | |
| PretrainedReader(bins=a.get("bins", 180), arch=backbone)) | |
| 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") | |
| 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) | |
| args = ap.parse_args() | |
| device = torch.device("mps" if torch.backends.mps.is_available() else "cpu") | |
| reader, res = build_reader(args.stage2, device) | |
| 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 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) | |
| hl, ml = reader(x) | |
| t, hour_only, dis, conf = decode_cls(hl.float().cpu(), ml.float().cpu()) | |
| 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), | |
| }) + "\n") | |
| n += 1 | |
| out.close() | |
| print(f"wrote {n} predictions -> {args.out}") | |
| if __name__ == "__main__": | |
| main() | |