Download twostage.py from lyte-codes/clockface: direct link, hf CLI and curl.
- Browser
- Download file 10.5 kB
-
https://huggingface.co/lyte-codes/clockface/resolve/main/twostage.py
- Command line
-
hf download hf://lyte-codes/clockface/twostage.py
-
curl -L -o twostage.py https://huggingface.co/lyte-codes/clockface/resolve/main/twostage.py
10.5 kB
| """Two-stage clock reading: find the dial, then read it. | |
| A single-stage CNN has to locate the face and measure two angles at once, from | |
| an image where the clock may fill 30% of the frame. Splitting those apart gives | |
| each stage an easier job, and the supervision for stage 1 is free because the | |
| renderer knows exactly where it put the dial. | |
| stage 1 DialLocator full image -> (cx, cy, r) of the dial | |
| stage 2 ClockReader a crop around the dial -> two angular distributions | |
| Stage 2 is trained on ground-truth crops with jitter, so it tolerates the error | |
| stage 1 will actually make at inference. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| import torch | |
| import torch.nn as nn | |
| from PIL import Image | |
| from model import ConvBlock, ClockNetCls, decode_cls, soft_targets, soft_argmax_angle # noqa: F401 | |
| class DialLocator(nn.Module): | |
| """Predicts the dial's centre and its 2x2 shape matrix. | |
| A centre and a radius describe a circle, and a dial viewed off-axis is an | |
| ellipse. The 2x2 matrix carries foreshortening, rotation and scale together, | |
| which is what makes rectification possible: invert it and the dial comes | |
| back frontal and upright. | |
| Output is (cx, cy, a, b, c, d) where the matrix maps dial-local coordinates | |
| to image offsets. | |
| """ | |
| def __init__(self, width=16, outputs=6): | |
| super().__init__() | |
| # outputs=6 is (cx, cy) + the 2x2 matrix. outputs=3 is the older | |
| # (cx, cy, log r) head, kept so existing checkpoints still load. | |
| self.outputs = outputs | |
| w = width | |
| self.net = nn.Sequential( | |
| ConvBlock(3, w, stride=2), | |
| ConvBlock(w, w * 2, stride=2), ConvBlock(w * 2, w * 2), | |
| ConvBlock(w * 2, w * 4, stride=2), ConvBlock(w * 4, w * 4), | |
| ConvBlock(w * 4, w * 8, stride=2), ConvBlock(w * 8, w * 8, stride=2), | |
| nn.AdaptiveAvgPool2d(4), nn.Flatten(), | |
| nn.Linear(w * 8 * 16, 128), nn.SiLU(inplace=True), nn.Linear(128, outputs), | |
| ) | |
| def forward(self, x): | |
| out = self.net(x) | |
| cx = torch.sigmoid(out[:, 0]) # centre stays inside the frame | |
| cy = torch.sigmoid(out[:, 1]) | |
| return torch.cat([cx.unsqueeze(1), cy.unsqueeze(1), out[:, 2:]], dim=1) | |
| def rectify_dial(img, dial, out_res, margin=1.15): | |
| """Warp a tilted dial back to a frontal, upright circle. | |
| The dial is a circle in the world and an ellipse in the image. Its 2x2 | |
| matrix M maps dial-local coordinates to image offsets, so M inverse undoes | |
| the foreshortening AND the in-image rotation in one step, handing stage 2 a | |
| canonical face instead of a squashed one at an arbitrary angle. | |
| PIL's AFFINE transform maps OUTPUT coordinates back into the INPUT, which is | |
| exactly the direction we already have, so no inversion is needed: output | |
| pixel -> dial-local -> input pixel. | |
| """ | |
| W, H = img.size | |
| M = dial["M"] | |
| cx, cy = dial["cx"] * W, dial["cy"] * H | |
| # pixels per unit of dial-local coordinate in the output | |
| k = out_res / (2.0 * margin) | |
| a, b = M[0][0] * W / k, M[0][1] * W / k | |
| d, e = M[1][0] * H / k, M[1][1] * H / k | |
| half = out_res / 2.0 | |
| c = cx - (a + b) * half | |
| f = cy - (d + e) * half | |
| return img.transform((out_res, out_res), Image.AFFINE, (a, b, c, d, e, f), | |
| resample=Image.BILINEAR) | |
| def crop_box(cx, cy, r, margin=1.25): | |
| """Square box around a dial, in normalised coordinates.""" | |
| half = r * margin | |
| return cx - half, cy - half, cx + half, cy + half | |
| def crop_dial(img, cx, cy, r, out_res, margin=1.25, rng=None, jitter=0.0): | |
| """Crop a PIL image around the dial and resize. | |
| With jitter > 0 the box is perturbed, so stage 2 sees the kind of imperfect | |
| box stage 1 will hand it rather than only perfect ones. | |
| """ | |
| W, H = img.size | |
| if rng is not None and jitter > 0: | |
| cx += rng.uniform(-jitter, jitter) * r | |
| cy += rng.uniform(-jitter, jitter) * r | |
| r *= 1.0 + rng.uniform(-jitter, jitter) | |
| x0, y0, x1, y1 = crop_box(cx * W, cy * H, r * W, margin) | |
| return img.crop((int(x0), int(y0), int(x1), int(y1))).resize((out_res, out_res)) | |
| def locator_targets(dials, device): | |
| rows = [] | |
| for d in dials: | |
| M = d["M"] | |
| rows.append([d["cx"], d["cy"], M[0][0], M[0][1], M[1][0], M[1][1]]) | |
| return torch.tensor(rows, dtype=torch.float32, device=device) | |
| def locator_loss(pred, tgt): | |
| """Centre and matrix, weighted so neither dominates. | |
| Matrix entries run about 0.1 to 0.8 and the centre runs 0 to 1, so the two | |
| are already on comparable scales; the centre is weighted up because getting | |
| it wrong moves the whole crop while a slightly wrong matrix only skews it. | |
| """ | |
| centre = ((pred[:, :2] - tgt[:, :2]) ** 2).mean() | |
| shape = ((pred[:, 2:6] - tgt[:, 2:6]) ** 2).mean() | |
| return centre * 4.0 + shape, centre.detach(), shape.detach() | |
| class PretrainedReader(nn.Module): | |
| """ImageNet-pretrained backbone + two angular heads and a whole-time head. | |
| A from-scratch CNN memorises 400 crops but sits on a plateau at 10.39 | |
| (= 2*ln(180), a uniform distribution) on 10,800. Yang/Xie/Zisserman used | |
| ImageNet-pretrained ResNet-50 for the same task; the features that make a | |
| dial's hands salient are exactly the low-level edge and orientation filters | |
| ImageNet training already provides. | |
| The third head exists because of what the two angular heads get wrong. On | |
| real photographs 27.5% of readings put the minute right and the hour wrong, | |
| split between hands read in the wrong roles and hours simply missed. Both | |
| angular heads describe a hand, so a face whose hands are hard to tell apart | |
| corrupts both of them at once and no amount of decoding recovers the hour -- | |
| measured, in jointdecode.py. `time_head` is asked instead for the time | |
| itself over the whole 12-hour ring, coarsely, with no notion of a hand. It | |
| cannot express a role confusion because it never assigns roles, and it only | |
| has to be good enough to break the twelve-way tie: the minute hand, which | |
| the model already reads well, supplies the precision. | |
| """ | |
| def __init__(self, bins=180, arch="mobilenet_v3_small", pretrained=True, time_bins=144): | |
| super().__init__() | |
| import torchvision.models as tvm | |
| self.bins = bins | |
| if arch == "mobilenet_v3_small": | |
| net = tvm.mobilenet_v3_small(weights=tvm.MobileNet_V3_Small_Weights.DEFAULT | |
| if pretrained else None) | |
| self.features = net.features | |
| feat_dim = 576 | |
| elif arch == "resnet18": | |
| net = tvm.resnet18(weights=tvm.ResNet18_Weights.DEFAULT if pretrained else None) | |
| self.features = nn.Sequential(*list(net.children())[:-2]) | |
| feat_dim = 512 | |
| elif arch == "resnet50": | |
| net = tvm.resnet50(weights=tvm.ResNet50_Weights.DEFAULT if pretrained else None) | |
| self.features = nn.Sequential(*list(net.children())[:-2]) | |
| feat_dim = 2048 | |
| else: | |
| raise ValueError(arch) | |
| self.pool = nn.AdaptiveAvgPool2d(2) # keep some spatial layout | |
| self.trunk = nn.Sequential(nn.Flatten(), nn.Linear(feat_dim * 4, 512), | |
| nn.SiLU(inplace=True), nn.Dropout(0.1)) | |
| self.hour_head = nn.Linear(512, bins) | |
| self.minute_head = nn.Linear(512, bins) | |
| self.time_bins = time_bins | |
| self.time_head = nn.Linear(512, time_bins) if time_bins else None | |
| def forward(self, x): | |
| f = self.trunk(self.pool(self.features(x))) | |
| t = self.time_head(f) if self.time_head is not None else None | |
| return self.hour_head(f), self.minute_head(f), t | |
| class CocoDialDetector: | |
| """Stage 1: find the dial with a COCO-pretrained detector. | |
| This replaces DialLocator. Measured on 200 held-out real photographs, same | |
| reader, three ways of producing a crop: | |
| whole image, no locator MAE 120.6 min median 88.3 | |
| DialLocator (2.2 MB) MAE 120.7 min median 67.9 | |
| this (78 MB) MAE 73.6 min median 18.4 | |
| DialLocator was trained on renders and contributes nothing over having no | |
| locator at all -- 120.7 against 120.6. Distilling this detector into it | |
| could not have helped either: a student cannot beat its teacher, and the | |
| training would have run through the renderer whose domain gap is the thing | |
| being escaped. | |
| It costs 75 MB. That buys halving the error, so the single-digit-MB target | |
| loses. | |
| """ | |
| COCO_CLOCK = 85 | |
| def __init__(self, arch="mobilenet", score_thresh=0.3, device=None, det_res=640): | |
| import torchvision.models.detection as tvd | |
| if arch == "mobilenet": | |
| # ResNet50-FPN trips the M2 GPU watchdog under sustained load | |
| # (kIOGPUCommandBufferCallbackErrorImpactingInteractivity). | |
| w = tvd.FasterRCNN_MobileNet_V3_Large_FPN_Weights.DEFAULT | |
| self.net = tvd.fasterrcnn_mobilenet_v3_large_fpn(weights=w) | |
| else: | |
| w = tvd.FasterRCNN_ResNet50_FPN_V2_Weights.DEFAULT | |
| self.net = tvd.fasterrcnn_resnet50_fpn_v2(weights=w) | |
| self.device = device or torch.device("mps" if torch.backends.mps.is_available() else "cpu") | |
| self.net = self.net.eval().to(self.device) | |
| self.score_thresh = score_thresh | |
| self.det_res = det_res | |
| def locate(self, img): | |
| """(cx, cy, r) normalised to image width, or None when nothing is found.""" | |
| import numpy as np | |
| scale = min(1.0, self.det_res / max(img.size)) | |
| small = (img.resize((max(32, round(img.width * scale)), | |
| max(32, round(img.height * scale)))) if scale < 1.0 else img) | |
| x = torch.from_numpy(np.asarray(small.convert("RGB"), dtype=np.float32) / 255.0) | |
| x = x.permute(2, 0, 1).to(self.device) | |
| out = self.net([x])[0] | |
| best, best_score = None, 0.0 | |
| for box, label, score in zip(out["boxes"], out["labels"], out["scores"]): | |
| s = float(score) | |
| if int(label) == self.COCO_CLOCK and s >= self.score_thresh and s > best_score: | |
| best, best_score = [float(v) / scale for v in box], s | |
| if best is None: | |
| return None | |
| x0, y0, x1, y1 = best | |
| return ((x0 + x1) / 2 / img.width, (y0 + y1) / 2 / img.height, | |
| max(x1 - x0, y1 - y0) / 2 / img.width, best_score) | |