"""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 @torch.no_grad() 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)