lyte-codes commited on
Commit
a37b3ff
·
verified ·
1 Parent(s): 20a4178

Sync jointdecode.py for 263702

Browse files
Files changed (1) hide show
  1. jointdecode.py +201 -0
jointdecode.py ADDED
@@ -0,0 +1,201 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Decode a time that a real clock could actually show.
2
+
3
+ The old decoder read each head on its own -- the hour head named an angle, the
4
+ minute head named an angle, and a vernier step stitched them together. Nothing
5
+ in that path can notice that the two angles describe no time at all. On a real
6
+ clock the hands are geared: at time t the hour hand sits at t/720 of a turn and
7
+ the minute hand at (t mod 60)/60 of a turn, so ONE number determines BOTH. Any
8
+ pair of angles that does not satisfy that is a reading no clock has ever shown.
9
+
10
+ So instead of reading the hands and hoping they agree, score every time the
11
+ clock could be showing and keep the one that best explains both heads:
12
+
13
+ score(t) = log p_hour(t/720 turn) + log p_minute((t mod 60)/60 turn)
14
+
15
+ That is a search over 720 minutes rather than over a plane of angle pairs, and
16
+ it cannot return an impossible answer. It also does something the independent
17
+ decoder cannot: a confident minute head drags a confused hour head onto the
18
+ right hour, because the hour term only has to break the 12-way tie.
19
+
20
+ Hand swaps are the largest single class of error on real photographs -- bigger
21
+ than correct and near-correct readings combined -- so the obvious next step was
22
+ to also score the heads exchanged and keep whichever assignment fits better.
23
+ That does not work, and the measurement is worth keeping rather than quietly
24
+ deleting. On the 200 held-out photographs the exchanged hypothesis scored a
25
+ median 0.06 better on readings that really were swapped and 0.11 worse on the
26
+ rest: the two populations sit on top of each other. Flipping whenever the
27
+ exchange won touched 75 of 200 readings and made 42 of them worse. No margin
28
+ threshold recovered anything; the best available threshold flips nothing.
29
+
30
+ The reason is that a swap is not a clean exchange of two correct angles. When
31
+ the model mistakes which hand is which it is confidently wrong in both heads at
32
+ once, and both heads then agree on a wrong but perfectly geared time. The
33
+ evidence that distinguishes the hands lives in the image, not in the output
34
+ distributions, so no decoder can recover it. `resolve_swap` is kept, off by
35
+ default, so that claim stays cheap to re-test against a future model.
36
+ """
37
+
38
+ from __future__ import annotations
39
+
40
+ import math
41
+
42
+ import torch
43
+
44
+ TWO_PI = 2 * math.pi
45
+
46
+
47
+ def _log_p_at(logp: torch.Tensor, angles: torch.Tensor) -> torch.Tensor:
48
+ """Read a bin distribution at arbitrary angles, interpolating between bins.
49
+
50
+ logp is (B, bins) log-probabilities over angle; angles is (T,) radians.
51
+ Returns (B, T). Bin b is centred at (b + 0.5)/bins of a turn and the ring
52
+ wraps, so bin 0 and the last bin are neighbours.
53
+ """
54
+ bins = logp.shape[1]
55
+ pos = angles / TWO_PI * bins - 0.5 # fractional bin coordinate
56
+ lo = torch.floor(pos)
57
+ frac = (pos - lo).to(logp.dtype)
58
+ i0 = (lo.long()) % bins
59
+ i1 = (i0 + 1) % bins
60
+ a = logp.index_select(1, i0)
61
+ b = logp.index_select(1, i1)
62
+ return a * (1 - frac) + b * frac
63
+
64
+
65
+ def joint_decode(hour_logits: torch.Tensor, minute_logits: torch.Tensor,
66
+ time_logits: torch.Tensor = None, time_weight: float = 1.0,
67
+ grid: int = 2880, resolve_swap: bool = False):
68
+ """Best time on the geared manifold, plus what the search learned.
69
+
70
+ Returns (minutes, margin, swapped, consistency):
71
+ minutes (B,) best time in [0, 720)
72
+ margin (B,) how much better the winner scored than the best time at
73
+ least 30 minutes away -- a peakedness measure that reflects
74
+ BOTH hands, unlike either head's own sharpness
75
+ swapped (B,) bool, True where exchanging the heads explained the
76
+ image better and the reading was taken from the exchange.
77
+ All False unless resolve_swap is on, which is not advised --
78
+ see the module docstring for what it measured
79
+ consistency (B,) winning score minus the score of the swapped reading;
80
+ negative means the swap won by that margin
81
+ """
82
+ hp = torch.log_softmax(hour_logits.float(), dim=1)
83
+ mp = torch.log_softmax(minute_logits.float(), dim=1)
84
+
85
+ t = torch.arange(grid, dtype=torch.float32) / grid * 720.0 # candidates
86
+ th = t / 720.0 * TWO_PI
87
+ tm = (t % 60.0) / 60.0 * TWO_PI
88
+
89
+ direct = _log_p_at(hp, th) + _log_p_at(mp, tm)
90
+ if time_logits is not None:
91
+ # the whole-time head votes on t itself, so it is read at t's own
92
+ # position on the ring rather than at either hand's direction
93
+ tp = torch.log_softmax(time_logits.float(), dim=1)
94
+ direct = direct + time_weight * _log_p_at(tp, t / 720.0 * TWO_PI)
95
+ if not resolve_swap:
96
+ best = direct.argmax(dim=1)
97
+ return t[best], _margin(direct, best, t), \
98
+ torch.zeros(len(best), dtype=torch.bool), torch.zeros(len(best))
99
+
100
+ # the same clock read with the roles of the two hands exchanged
101
+ swap = _log_p_at(mp, th) + _log_p_at(hp, tm)
102
+
103
+ d_best, d_i = direct.max(dim=1)
104
+ s_best, s_i = swap.max(dim=1)
105
+ take_swap = s_best > d_best
106
+ idx = torch.where(take_swap, s_i, d_i)
107
+ scores = torch.where(take_swap.unsqueeze(1), swap, direct)
108
+ return t[idx], _margin(scores, idx, t), take_swap, d_best - s_best
109
+
110
+
111
+ def _margin(scores: torch.Tensor, best: torch.Tensor, t: torch.Tensor):
112
+ """Winner's score minus the best score at least 30 minutes away on the ring.
113
+
114
+ A single sharp peak scores high here. Two plausible readings -- the usual
115
+ shape when a hand is occluded or the dial is badly blurred -- score near
116
+ zero, which is the honest answer.
117
+ """
118
+ d = (t.unsqueeze(0) - t[best].unsqueeze(1)).abs()
119
+ far = torch.minimum(d, 720.0 - d) >= 30.0
120
+ top = scores.gather(1, best.unsqueeze(1)).squeeze(1)
121
+ runner = scores.masked_fill(~far, float("-inf")).max(dim=1).values
122
+ return top - runner
123
+
124
+
125
+ def _self_test():
126
+ """Check the decoder against times whose answer is known by construction."""
127
+ from model import decode_cls, ring_target, soft_targets
128
+
129
+ ok = fail = 0
130
+
131
+ def check(name, cond):
132
+ nonlocal ok, fail
133
+ if cond:
134
+ ok += 1
135
+ print(f" ok {name}")
136
+ else:
137
+ fail += 1
138
+ print(f" FAIL {name}")
139
+
140
+ L = lambda p: torch.log(p + 1e-9)
141
+ times = [0.0, 1.0, 90.0, 187.5, 359.0, 425.0, 719.5]
142
+
143
+ # heads that know the answer must decode to the answer
144
+ for true in times:
145
+ h, m = soft_targets(torch.tensor([true]), 180, 60.0)
146
+ t, margin, sw, _ = joint_decode(L(h), L(m))
147
+ e = min(abs(t.item() - true), 720 - abs(t.item() - true))
148
+ check(f"{true:6.1f} min decodes to itself within a quarter minute", e < 0.25)
149
+ check(f"{true:6.1f} min is not flagged as a swap", not bool(sw[0]))
150
+
151
+ # every decoded time must be one a clock could show: the two hand angles it
152
+ # implies have to be geared together, which is what the search guarantees
153
+ for true in times:
154
+ h, m = soft_targets(torch.tensor([true]), 180, 60.0)
155
+ t, _, _, _ = joint_decode(L(h), L(m))
156
+ implied_hour = t.item() / 720.0 * 360.0
157
+ implied_min = (t.item() % 60) / 60.0 * 360.0
158
+ check(f"{true:6.1f} min decodes onto the geared manifold",
159
+ abs((implied_hour * 12) % 360 - implied_min) < 1e-3)
160
+
161
+ # a confident minute hand should pull a hour hand that is merely vague onto
162
+ # the right hour -- the case the independent decoder cannot handle
163
+ true = 425.0
164
+ h, m = soft_targets(torch.tensor([true]), 180, 60.0)
165
+ vague = torch.full_like(h, 1.0 / h.shape[1])
166
+ vague = 0.75 * vague + 0.25 * h # a weak hint, not a peak
167
+ t, _, _, _ = joint_decode(L(vague), L(m))
168
+ e = min(abs(t.item() - true), 720 - abs(t.item() - true))
169
+ check("a vague hour hand plus a sharp minute hand still lands in the right hour", e < 1.0)
170
+
171
+ # the whole-time head should be able to break a twelve-way tie on its own
172
+ h_flat = torch.full((1, 180), 1.0 / 180)
173
+ _, m = soft_targets(torch.tensor([true]), 180, 60.0)
174
+ without = joint_decode(L(h_flat), L(m))[0].item()
175
+ ring = ring_target(torch.tensor([true]), 144, 20.0)
176
+ with_ = joint_decode(L(h_flat), L(m), L(ring))[0].item()
177
+ e_wo = min(abs(without - true), 720 - abs(without - true))
178
+ e_w = min(abs(with_ - true), 720 - abs(with_ - true))
179
+ check("with no hour hand at all the hour is a guess", e_wo > 30)
180
+ check("the whole-time head recovers the hour the missing hand cannot give", e_w < 1.0)
181
+
182
+ # the swap path stays off unless it is asked for, because it measured worse
183
+ h, m = soft_targets(torch.tensor([190.0]), 180, 60.0)
184
+ check("swap resolution is off by default", not bool(joint_decode(L(m), L(h))[2][0]))
185
+ check("swap resolution still works when asked for",
186
+ bool(joint_decode(L(m), L(h), resolve_swap=True)[2][0]))
187
+
188
+ # margin should be smaller when two readings are equally plausible
189
+ h, m = soft_targets(torch.tensor([180.0]), 180, 60.0)
190
+ sharp = joint_decode(L(h), L(m))[1].item()
191
+ h2, m2 = soft_targets(torch.tensor([540.0]), 180, 60.0)
192
+ two = joint_decode(L(0.5 * h + 0.5 * h2), L(0.5 * m + 0.5 * m2))[1].item()
193
+ check("one clear reading scores a wider margin than two rival readings", sharp > two)
194
+
195
+ print(f"\n{ok}/{ok + fail} checks passed")
196
+ return fail == 0
197
+
198
+
199
+ if __name__ == "__main__":
200
+ import sys
201
+ sys.exit(0 if _self_test() else 1)