clockface / twostage.py
lyte-codes's picture
Sync twostage.py for 263702
24e3898 verified
Raw History Blame Contribute Delete
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
@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)