Download model.py from lyte-codes/clockface: direct link, hf CLI and curl.
- Browser
- Download file 9.02 kB
-
https://huggingface.co/lyte-codes/clockface/resolve/98f945c8702fd79792ed3cbc1718052d3cb12fd2/model.py
- Command line
-
hf download hf://lyte-codes/clockface@98f945c8702fd79792ed3cbc1718052d3cb12fd2/model.py
-
curl -L -o model.py https://huggingface.co/lyte-codes/clockface/resolve/98f945c8702fd79792ed3cbc1718052d3cb12fd2/model.py
9.02 kB
| """The clockface model: a small CNN that reads both hands. | |
| Two angles come out, not one time. The hour hand alone determines the time, and | |
| the minute hand alone determines it modulo an hour; predicting both lets them | |
| be checked against each other, which is where the confidence signal comes from. | |
| Angles are represented as (sin, cos), never as raw degrees. Degrees wrap at 360, | |
| so 359 and 1 are neighbours on a dial but maximally distant in the loss, and a | |
| model trained on raw degrees blows up at the seam. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| import torch | |
| import torch.nn as nn | |
| class ConvBlock(nn.Module): | |
| def __init__(self, cin, cout, stride=1): | |
| super().__init__() | |
| self.conv = nn.Conv2d(cin, cout, 3, stride=stride, padding=1, bias=False) | |
| # BatchNorm rather than GroupNorm: measured on this M2, GroupNorm runs a | |
| # forward+backward of this net at 39 img/s against BatchNorm's 65, and | |
| # the batch is large enough (64) for batch statistics to be stable. | |
| self.norm = nn.BatchNorm2d(cout) | |
| self.act = nn.SiLU(inplace=True) | |
| def forward(self, x): | |
| return self.act(self.norm(self.conv(x))) | |
| class ClockNet(nn.Module): | |
| """Small CNN -> 4 numbers: (sin, cos) for the hour hand and the minute hand.""" | |
| def __init__(self, width=32, in_res=256): | |
| super().__init__() | |
| w = width | |
| self.stem = ConvBlock(3, w, stride=2) # 112 | |
| self.stage1 = nn.Sequential(ConvBlock(w, w), ConvBlock(w, w * 2, stride=2)) # 56 | |
| self.stage2 = nn.Sequential(ConvBlock(w * 2, w * 2), ConvBlock(w * 2, w * 4, stride=2)) # 28 | |
| self.stage3 = nn.Sequential(ConvBlock(w * 4, w * 4), ConvBlock(w * 4, w * 8, stride=2)) # 14 | |
| self.stage4 = nn.Sequential(ConvBlock(w * 8, w * 8), ConvBlock(w * 8, w * 8, stride=2)) # 7 | |
| # Keep a 4x4 spatial grid rather than collapsing to 1x1. Reading a hand | |
| # angle is a question about WHERE something points, and global average | |
| # pooling discards exactly that: it is the right head for "is there a | |
| # clock" and the wrong one for "which way does the hand point". | |
| # 256px in -> 8x8 final map -> pool to 4x4. MPS cannot adaptive-pool | |
| # when the input size is not divisible by the output size, which 7->4 | |
| # is not; 8->4 is. | |
| self.pool = nn.AdaptiveAvgPool2d(4) | |
| self.head = nn.Sequential( | |
| nn.Flatten(), | |
| nn.Linear(w * 8 * 16, 256), nn.SiLU(inplace=True), | |
| nn.Dropout(0.1), | |
| nn.Linear(256, 4), | |
| ) | |
| def forward(self, x): | |
| x = self.stem(x) | |
| x = self.stage1(x); x = self.stage2(x); x = self.stage3(x); x = self.stage4(x) | |
| out = self.head(self.pool(x)) | |
| # Return the RAW (sin, cos) pairs. Normalising here divides by a | |
| # magnitude that is near zero at initialisation, so the gradient through | |
| # the division scales as 1/||x|| and the early steps thrash. The loss | |
| # against unit-length targets pulls the magnitude to 1 on its own, and | |
| # decode() normalises when it needs a direction. The magnitude is still | |
| # a usable confidence signal: a short vector means the model is unsure. | |
| h_mag = out[:, 0:2].norm(dim=1, keepdim=True) | |
| m_mag = out[:, 2:4].norm(dim=1, keepdim=True) | |
| return out, torch.cat([h_mag, m_mag], dim=1) | |
| def angles_to_targets(minutes: torch.Tensor) -> torch.Tensor: | |
| """minutes on the 720 ring -> (sin,cos) of each hand's angle.""" | |
| hour_ang = minutes / 720.0 * 2 * math.pi | |
| min_ang = (minutes % 60.0) / 60.0 * 2 * math.pi | |
| return torch.stack([torch.sin(hour_ang), torch.cos(hour_ang), | |
| torch.sin(min_ang), torch.cos(min_ang)], dim=1) | |
| def decode(pred: torch.Tensor): | |
| """(sin,cos) pairs -> a time in minutes, plus the two hands' disagreement. | |
| The minute hand is precise but says nothing about which hour it is; the hour | |
| hand says which hour but reads the minutes coarsely. Combine them the way a | |
| vernier scale does: take the minute-of-hour from the minute hand, and take | |
| only the hour count from the hour hand. | |
| """ | |
| # atan2 is scale-invariant, so raw (unnormalised) outputs decode correctly. | |
| h_ang = torch.atan2(pred[:, 0], pred[:, 1]) % (2 * math.pi) | |
| m_ang = torch.atan2(pred[:, 2], pred[:, 3]) % (2 * math.pi) | |
| hour_minutes = h_ang / (2 * math.pi) * 720.0 # what the hour hand alone says | |
| minute_of_hour = m_ang / (2 * math.pi) * 60.0 # what the minute hand alone says | |
| # which hour does the minute hand belong to, given the hour hand | |
| k = torch.round((hour_minutes - minute_of_hour) / 60.0) | |
| combined = (k * 60.0 + minute_of_hour) % 720.0 | |
| # disagreement in minutes between the two readings, on the 720 ring | |
| d = (hour_minutes - combined).abs() % 720.0 | |
| disagreement = torch.minimum(d, 720.0 - d) | |
| return combined, hour_minutes, disagreement | |
| def count_params(model): | |
| n = sum(p.numel() for p in model.parameters()) | |
| return n, n * 4 / 1e6 # float32 megabytes | |
| if __name__ == "__main__": | |
| m = ClockNet() | |
| n, mb = count_params(m) | |
| x = torch.randn(2, 3, 256, 256) | |
| pred, mag = m(x) | |
| t, hm, dis = decode(pred) | |
| print(f"ClockNet: {n:,} params, {mb:.2f} MB float32") | |
| print(f" forward {tuple(x.shape)} -> pred {tuple(pred.shape)}, magnitudes {tuple(mag.shape)}") | |
| print(f" decoded times {t.tolist()}") | |
| print(f" disagreement {dis.tolist()}") | |
| # --------------------------------------------------------------------------- | |
| # Classification head. | |
| # | |
| # Regressing (sin, cos) under MSE collapses to the mean: predicting the zero | |
| # vector scores 0.5 against unit-length targets, and on this data the optimiser | |
| # settles there rather than finding the hands (measured: loss 0.4986, val MAE | |
| # 180 min = chance). Yang/Xie/Zisserman report the same thing, 5.4% against | |
| # 59.6% for classification, and give the reason: too little penalty for being | |
| # slightly wrong. | |
| # | |
| # So predict a DISTRIBUTION over angles for each hand instead. Bins are | |
| # circular, the target is a von Mises bump rather than a one-hot (so being one | |
| # bin out is genuinely cheaper than being ten out), and decoding takes a | |
| # circular soft-argmax, which recovers sub-bin precision. | |
| # --------------------------------------------------------------------------- | |
| class ClockNetCls(nn.Module): | |
| def __init__(self, width=32, bins=180): | |
| super().__init__() | |
| w = width | |
| self.bins = bins | |
| self.stem = ConvBlock(3, w, stride=2) | |
| self.stage1 = nn.Sequential(ConvBlock(w, w), ConvBlock(w, w * 2, stride=2)) | |
| self.stage2 = nn.Sequential(ConvBlock(w * 2, w * 2), ConvBlock(w * 2, w * 4, stride=2)) | |
| self.stage3 = nn.Sequential(ConvBlock(w * 4, w * 4), ConvBlock(w * 4, w * 8, stride=2)) | |
| self.stage4 = nn.Sequential(ConvBlock(w * 8, w * 8), ConvBlock(w * 8, w * 8, stride=2)) | |
| self.pool = nn.AdaptiveAvgPool2d(4) | |
| self.trunk = nn.Sequential(nn.Flatten(), nn.Linear(w * 8 * 16, 512), nn.SiLU(inplace=True), | |
| nn.Dropout(0.1)) | |
| self.hour_head = nn.Linear(512, bins) | |
| self.minute_head = nn.Linear(512, bins) | |
| def forward(self, x): | |
| x = self.stem(x) | |
| x = self.stage1(x); x = self.stage2(x); x = self.stage3(x); x = self.stage4(x) | |
| f = self.trunk(self.pool(x)) | |
| return self.hour_head(f), self.minute_head(f) | |
| def soft_targets(minutes: torch.Tensor, bins: int, kappa: float = 40.0): | |
| """von Mises bumps over circular bins, for the hour and minute hands.""" | |
| dev = minutes.device | |
| centres = (torch.arange(bins, device=dev, dtype=torch.float32) + 0.5) / bins * 2 * math.pi | |
| out = [] | |
| for ang in (minutes / 720.0 * 2 * math.pi, (minutes % 60.0) / 60.0 * 2 * math.pi): | |
| d = centres.unsqueeze(0) - ang.unsqueeze(1) | |
| t = torch.exp(kappa * (torch.cos(d) - 1.0)) | |
| out.append(t / t.sum(dim=1, keepdim=True)) | |
| return out | |
| def soft_argmax_angle(logits: torch.Tensor): | |
| """Circular expectation of a distribution over angle bins -> radians.""" | |
| bins = logits.shape[1] | |
| p = torch.softmax(logits, dim=1) | |
| centres = (torch.arange(bins, device=logits.device, dtype=torch.float32) + 0.5) / bins * 2 * math.pi | |
| s = (p * torch.sin(centres)).sum(dim=1) | |
| c = (p * torch.cos(centres)).sum(dim=1) | |
| return torch.atan2(s, c) % (2 * math.pi), torch.sqrt(s ** 2 + c ** 2) | |
| def decode_cls(hour_logits, minute_logits): | |
| """Same vernier decode as the regression head, from two distributions.""" | |
| h_ang, h_conf = soft_argmax_angle(hour_logits) | |
| m_ang, m_conf = soft_argmax_angle(minute_logits) | |
| hour_minutes = h_ang / (2 * math.pi) * 720.0 | |
| minute_of_hour = m_ang / (2 * math.pi) * 60.0 | |
| k = torch.round((hour_minutes - minute_of_hour) / 60.0) | |
| combined = (k * 60.0 + minute_of_hour) % 720.0 | |
| d = (hour_minutes - combined).abs() % 720.0 | |
| return combined, hour_minutes, torch.minimum(d, 720.0 - d), torch.stack([h_conf, m_conf], 1) | |