lyte-codes commited on
Commit
4837dcb
·
verified ·
1 Parent(s): 30e5fb4

Add model.py

Browse files
Files changed (1) hide show
  1. model.py +195 -0
model.py ADDED
@@ -0,0 +1,195 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The clockface model: a small CNN that reads both hands.
2
+
3
+ Two angles come out, not one time. The hour hand alone determines the time, and
4
+ the minute hand alone determines it modulo an hour; predicting both lets them
5
+ be checked against each other, which is where the confidence signal comes from.
6
+
7
+ Angles are represented as (sin, cos), never as raw degrees. Degrees wrap at 360,
8
+ so 359 and 1 are neighbours on a dial but maximally distant in the loss, and a
9
+ model trained on raw degrees blows up at the seam.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import math
15
+
16
+ import torch
17
+ import torch.nn as nn
18
+
19
+
20
+ class ConvBlock(nn.Module):
21
+ def __init__(self, cin, cout, stride=1):
22
+ super().__init__()
23
+ self.conv = nn.Conv2d(cin, cout, 3, stride=stride, padding=1, bias=False)
24
+ # BatchNorm rather than GroupNorm: measured on this M2, GroupNorm runs a
25
+ # forward+backward of this net at 39 img/s against BatchNorm's 65, and
26
+ # the batch is large enough (64) for batch statistics to be stable.
27
+ self.norm = nn.BatchNorm2d(cout)
28
+ self.act = nn.SiLU(inplace=True)
29
+
30
+ def forward(self, x):
31
+ return self.act(self.norm(self.conv(x)))
32
+
33
+
34
+ class ClockNet(nn.Module):
35
+ """Small CNN -> 4 numbers: (sin, cos) for the hour hand and the minute hand."""
36
+
37
+ def __init__(self, width=32, in_res=256):
38
+ super().__init__()
39
+ w = width
40
+ self.stem = ConvBlock(3, w, stride=2) # 112
41
+ self.stage1 = nn.Sequential(ConvBlock(w, w), ConvBlock(w, w * 2, stride=2)) # 56
42
+ self.stage2 = nn.Sequential(ConvBlock(w * 2, w * 2), ConvBlock(w * 2, w * 4, stride=2)) # 28
43
+ self.stage3 = nn.Sequential(ConvBlock(w * 4, w * 4), ConvBlock(w * 4, w * 8, stride=2)) # 14
44
+ self.stage4 = nn.Sequential(ConvBlock(w * 8, w * 8), ConvBlock(w * 8, w * 8, stride=2)) # 7
45
+ # Keep a 4x4 spatial grid rather than collapsing to 1x1. Reading a hand
46
+ # angle is a question about WHERE something points, and global average
47
+ # pooling discards exactly that: it is the right head for "is there a
48
+ # clock" and the wrong one for "which way does the hand point".
49
+ # 256px in -> 8x8 final map -> pool to 4x4. MPS cannot adaptive-pool
50
+ # when the input size is not divisible by the output size, which 7->4
51
+ # is not; 8->4 is.
52
+ self.pool = nn.AdaptiveAvgPool2d(4)
53
+ self.head = nn.Sequential(
54
+ nn.Flatten(),
55
+ nn.Linear(w * 8 * 16, 256), nn.SiLU(inplace=True),
56
+ nn.Dropout(0.1),
57
+ nn.Linear(256, 4),
58
+ )
59
+
60
+ def forward(self, x):
61
+ x = self.stem(x)
62
+ x = self.stage1(x); x = self.stage2(x); x = self.stage3(x); x = self.stage4(x)
63
+ out = self.head(self.pool(x))
64
+ # Return the RAW (sin, cos) pairs. Normalising here divides by a
65
+ # magnitude that is near zero at initialisation, so the gradient through
66
+ # the division scales as 1/||x|| and the early steps thrash. The loss
67
+ # against unit-length targets pulls the magnitude to 1 on its own, and
68
+ # decode() normalises when it needs a direction. The magnitude is still
69
+ # a usable confidence signal: a short vector means the model is unsure.
70
+ h_mag = out[:, 0:2].norm(dim=1, keepdim=True)
71
+ m_mag = out[:, 2:4].norm(dim=1, keepdim=True)
72
+ return out, torch.cat([h_mag, m_mag], dim=1)
73
+
74
+
75
+ def angles_to_targets(minutes: torch.Tensor) -> torch.Tensor:
76
+ """minutes on the 720 ring -> (sin,cos) of each hand's angle."""
77
+ hour_ang = minutes / 720.0 * 2 * math.pi
78
+ min_ang = (minutes % 60.0) / 60.0 * 2 * math.pi
79
+ return torch.stack([torch.sin(hour_ang), torch.cos(hour_ang),
80
+ torch.sin(min_ang), torch.cos(min_ang)], dim=1)
81
+
82
+
83
+ def decode(pred: torch.Tensor):
84
+ """(sin,cos) pairs -> a time in minutes, plus the two hands' disagreement.
85
+
86
+ The minute hand is precise but says nothing about which hour it is; the hour
87
+ hand says which hour but reads the minutes coarsely. Combine them the way a
88
+ vernier scale does: take the minute-of-hour from the minute hand, and take
89
+ only the hour count from the hour hand.
90
+ """
91
+ # atan2 is scale-invariant, so raw (unnormalised) outputs decode correctly.
92
+ h_ang = torch.atan2(pred[:, 0], pred[:, 1]) % (2 * math.pi)
93
+ m_ang = torch.atan2(pred[:, 2], pred[:, 3]) % (2 * math.pi)
94
+
95
+ hour_minutes = h_ang / (2 * math.pi) * 720.0 # what the hour hand alone says
96
+ minute_of_hour = m_ang / (2 * math.pi) * 60.0 # what the minute hand alone says
97
+
98
+ # which hour does the minute hand belong to, given the hour hand
99
+ k = torch.round((hour_minutes - minute_of_hour) / 60.0)
100
+ combined = (k * 60.0 + minute_of_hour) % 720.0
101
+
102
+ # disagreement in minutes between the two readings, on the 720 ring
103
+ d = (hour_minutes - combined).abs() % 720.0
104
+ disagreement = torch.minimum(d, 720.0 - d)
105
+ return combined, hour_minutes, disagreement
106
+
107
+
108
+ def count_params(model):
109
+ n = sum(p.numel() for p in model.parameters())
110
+ return n, n * 4 / 1e6 # float32 megabytes
111
+
112
+
113
+ if __name__ == "__main__":
114
+ m = ClockNet()
115
+ n, mb = count_params(m)
116
+ x = torch.randn(2, 3, 256, 256)
117
+ pred, mag = m(x)
118
+ t, hm, dis = decode(pred)
119
+ print(f"ClockNet: {n:,} params, {mb:.2f} MB float32")
120
+ print(f" forward {tuple(x.shape)} -> pred {tuple(pred.shape)}, magnitudes {tuple(mag.shape)}")
121
+ print(f" decoded times {t.tolist()}")
122
+ print(f" disagreement {dis.tolist()}")
123
+
124
+
125
+ # ---------------------------------------------------------------------------
126
+ # Classification head.
127
+ #
128
+ # Regressing (sin, cos) under MSE collapses to the mean: predicting the zero
129
+ # vector scores 0.5 against unit-length targets, and on this data the optimiser
130
+ # settles there rather than finding the hands (measured: loss 0.4986, val MAE
131
+ # 180 min = chance). Yang/Xie/Zisserman report the same thing, 5.4% against
132
+ # 59.6% for classification, and give the reason: too little penalty for being
133
+ # slightly wrong.
134
+ #
135
+ # So predict a DISTRIBUTION over angles for each hand instead. Bins are
136
+ # circular, the target is a von Mises bump rather than a one-hot (so being one
137
+ # bin out is genuinely cheaper than being ten out), and decoding takes a
138
+ # circular soft-argmax, which recovers sub-bin precision.
139
+ # ---------------------------------------------------------------------------
140
+
141
+ class ClockNetCls(nn.Module):
142
+ def __init__(self, width=32, bins=180):
143
+ super().__init__()
144
+ w = width
145
+ self.bins = bins
146
+ self.stem = ConvBlock(3, w, stride=2)
147
+ self.stage1 = nn.Sequential(ConvBlock(w, w), ConvBlock(w, w * 2, stride=2))
148
+ self.stage2 = nn.Sequential(ConvBlock(w * 2, w * 2), ConvBlock(w * 2, w * 4, stride=2))
149
+ self.stage3 = nn.Sequential(ConvBlock(w * 4, w * 4), ConvBlock(w * 4, w * 8, stride=2))
150
+ self.stage4 = nn.Sequential(ConvBlock(w * 8, w * 8), ConvBlock(w * 8, w * 8, stride=2))
151
+ self.pool = nn.AdaptiveAvgPool2d(4)
152
+ self.trunk = nn.Sequential(nn.Flatten(), nn.Linear(w * 8 * 16, 512), nn.SiLU(inplace=True),
153
+ nn.Dropout(0.1))
154
+ self.hour_head = nn.Linear(512, bins)
155
+ self.minute_head = nn.Linear(512, bins)
156
+
157
+ def forward(self, x):
158
+ x = self.stem(x)
159
+ x = self.stage1(x); x = self.stage2(x); x = self.stage3(x); x = self.stage4(x)
160
+ f = self.trunk(self.pool(x))
161
+ return self.hour_head(f), self.minute_head(f)
162
+
163
+
164
+ def soft_targets(minutes: torch.Tensor, bins: int, kappa: float = 40.0):
165
+ """von Mises bumps over circular bins, for the hour and minute hands."""
166
+ dev = minutes.device
167
+ centres = (torch.arange(bins, device=dev, dtype=torch.float32) + 0.5) / bins * 2 * math.pi
168
+ out = []
169
+ for ang in (minutes / 720.0 * 2 * math.pi, (minutes % 60.0) / 60.0 * 2 * math.pi):
170
+ d = centres.unsqueeze(0) - ang.unsqueeze(1)
171
+ t = torch.exp(kappa * (torch.cos(d) - 1.0))
172
+ out.append(t / t.sum(dim=1, keepdim=True))
173
+ return out
174
+
175
+
176
+ def soft_argmax_angle(logits: torch.Tensor):
177
+ """Circular expectation of a distribution over angle bins -> radians."""
178
+ bins = logits.shape[1]
179
+ p = torch.softmax(logits, dim=1)
180
+ centres = (torch.arange(bins, device=logits.device, dtype=torch.float32) + 0.5) / bins * 2 * math.pi
181
+ s = (p * torch.sin(centres)).sum(dim=1)
182
+ c = (p * torch.cos(centres)).sum(dim=1)
183
+ return torch.atan2(s, c) % (2 * math.pi), torch.sqrt(s ** 2 + c ** 2)
184
+
185
+
186
+ def decode_cls(hour_logits, minute_logits):
187
+ """Same vernier decode as the regression head, from two distributions."""
188
+ h_ang, h_conf = soft_argmax_angle(hour_logits)
189
+ m_ang, m_conf = soft_argmax_angle(minute_logits)
190
+ hour_minutes = h_ang / (2 * math.pi) * 720.0
191
+ minute_of_hour = m_ang / (2 * math.pi) * 60.0
192
+ k = torch.round((hour_minutes - minute_of_hour) / 60.0)
193
+ combined = (k * 60.0 + minute_of_hour) % 720.0
194
+ d = (hour_minutes - combined).abs() % 720.0
195
+ return combined, hour_minutes, torch.minimum(d, 720.0 - d), torch.stack([h_conf, m_conf], 1)