| """Decode a time that a real clock could actually show. |
| |
| The old decoder read each head on its own -- the hour head named an angle, the |
| minute head named an angle, and a vernier step stitched them together. Nothing |
| in that path can notice that the two angles describe no time at all. On a real |
| clock the hands are geared: at time t the hour hand sits at t/720 of a turn and |
| the minute hand at (t mod 60)/60 of a turn, so ONE number determines BOTH. Any |
| pair of angles that does not satisfy that is a reading no clock has ever shown. |
| |
| So instead of reading the hands and hoping they agree, score every time the |
| clock could be showing and keep the one that best explains both heads: |
| |
| score(t) = log p_hour(t/720 turn) + log p_minute((t mod 60)/60 turn) |
| |
| That is a search over 720 minutes rather than over a plane of angle pairs, and |
| it cannot return an impossible answer. It also does something the independent |
| decoder cannot: a confident minute head drags a confused hour head onto the |
| right hour, because the hour term only has to break the 12-way tie. |
| |
| Hand swaps are the largest single class of error on real photographs -- bigger |
| than correct and near-correct readings combined -- so the obvious next step was |
| to also score the heads exchanged and keep whichever assignment fits better. |
| That does not work, and the measurement is worth keeping rather than quietly |
| deleting. On the 200 held-out photographs the exchanged hypothesis scored a |
| median 0.06 better on readings that really were swapped and 0.11 worse on the |
| rest: the two populations sit on top of each other. Flipping whenever the |
| exchange won touched 75 of 200 readings and made 42 of them worse. No margin |
| threshold recovered anything; the best available threshold flips nothing. |
| |
| The reason is that a swap is not a clean exchange of two correct angles. When |
| the model mistakes which hand is which it is confidently wrong in both heads at |
| once, and both heads then agree on a wrong but perfectly geared time. The |
| evidence that distinguishes the hands lives in the image, not in the output |
| distributions, so no decoder can recover it. `resolve_swap` is kept, off by |
| default, so that claim stays cheap to re-test against a future model. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import math |
|
|
| import torch |
|
|
| TWO_PI = 2 * math.pi |
|
|
|
|
| def _log_p_at(logp: torch.Tensor, angles: torch.Tensor) -> torch.Tensor: |
| """Read a bin distribution at arbitrary angles, interpolating between bins. |
| |
| logp is (B, bins) log-probabilities over angle; angles is (T,) radians. |
| Returns (B, T). Bin b is centred at (b + 0.5)/bins of a turn and the ring |
| wraps, so bin 0 and the last bin are neighbours. |
| """ |
| bins = logp.shape[1] |
| pos = angles / TWO_PI * bins - 0.5 |
| lo = torch.floor(pos) |
| frac = (pos - lo).to(logp.dtype) |
| i0 = (lo.long()) % bins |
| i1 = (i0 + 1) % bins |
| a = logp.index_select(1, i0) |
| b = logp.index_select(1, i1) |
| return a * (1 - frac) + b * frac |
|
|
|
|
| def joint_decode(hour_logits: torch.Tensor, minute_logits: torch.Tensor, |
| time_logits: torch.Tensor = None, time_weight: float = 1.0, |
| grid: int = 2880, resolve_swap: bool = False): |
| """Best time on the geared manifold, plus what the search learned. |
| |
| Returns (minutes, margin, swapped, consistency): |
| minutes (B,) best time in [0, 720) |
| margin (B,) how much better the winner scored than the best time at |
| least 30 minutes away -- a peakedness measure that reflects |
| BOTH hands, unlike either head's own sharpness |
| swapped (B,) bool, True where exchanging the heads explained the |
| image better and the reading was taken from the exchange. |
| All False unless resolve_swap is on, which is not advised -- |
| see the module docstring for what it measured |
| consistency (B,) winning score minus the score of the swapped reading; |
| negative means the swap won by that margin |
| """ |
| hp = torch.log_softmax(hour_logits.float(), dim=1) |
| mp = torch.log_softmax(minute_logits.float(), dim=1) |
|
|
| t = torch.arange(grid, dtype=torch.float32) / grid * 720.0 |
| th = t / 720.0 * TWO_PI |
| tm = (t % 60.0) / 60.0 * TWO_PI |
|
|
| direct = _log_p_at(hp, th) + _log_p_at(mp, tm) |
| if time_logits is not None: |
| |
| |
| tp = torch.log_softmax(time_logits.float(), dim=1) |
| direct = direct + time_weight * _log_p_at(tp, t / 720.0 * TWO_PI) |
| if not resolve_swap: |
| best = direct.argmax(dim=1) |
| return t[best], _margin(direct, best, t), \ |
| torch.zeros(len(best), dtype=torch.bool), torch.zeros(len(best)) |
|
|
| |
| swap = _log_p_at(mp, th) + _log_p_at(hp, tm) |
|
|
| d_best, d_i = direct.max(dim=1) |
| s_best, s_i = swap.max(dim=1) |
| take_swap = s_best > d_best |
| idx = torch.where(take_swap, s_i, d_i) |
| scores = torch.where(take_swap.unsqueeze(1), swap, direct) |
| return t[idx], _margin(scores, idx, t), take_swap, d_best - s_best |
|
|
|
|
| def _margin(scores: torch.Tensor, best: torch.Tensor, t: torch.Tensor): |
| """Winner's score minus the best score at least 30 minutes away on the ring. |
| |
| A single sharp peak scores high here. Two plausible readings -- the usual |
| shape when a hand is occluded or the dial is badly blurred -- score near |
| zero, which is the honest answer. |
| """ |
| d = (t.unsqueeze(0) - t[best].unsqueeze(1)).abs() |
| far = torch.minimum(d, 720.0 - d) >= 30.0 |
| top = scores.gather(1, best.unsqueeze(1)).squeeze(1) |
| runner = scores.masked_fill(~far, float("-inf")).max(dim=1).values |
| return top - runner |
|
|
|
|
| def _self_test(): |
| """Check the decoder against times whose answer is known by construction.""" |
| from model import decode_cls, ring_target, soft_targets |
|
|
| ok = fail = 0 |
|
|
| def check(name, cond): |
| nonlocal ok, fail |
| if cond: |
| ok += 1 |
| print(f" ok {name}") |
| else: |
| fail += 1 |
| print(f" FAIL {name}") |
|
|
| L = lambda p: torch.log(p + 1e-9) |
| times = [0.0, 1.0, 90.0, 187.5, 359.0, 425.0, 719.5] |
|
|
| |
| for true in times: |
| h, m = soft_targets(torch.tensor([true]), 180, 60.0) |
| t, margin, sw, _ = joint_decode(L(h), L(m)) |
| e = min(abs(t.item() - true), 720 - abs(t.item() - true)) |
| check(f"{true:6.1f} min decodes to itself within a quarter minute", e < 0.25) |
| check(f"{true:6.1f} min is not flagged as a swap", not bool(sw[0])) |
|
|
| |
| |
| for true in times: |
| h, m = soft_targets(torch.tensor([true]), 180, 60.0) |
| t, _, _, _ = joint_decode(L(h), L(m)) |
| implied_hour = t.item() / 720.0 * 360.0 |
| implied_min = (t.item() % 60) / 60.0 * 360.0 |
| check(f"{true:6.1f} min decodes onto the geared manifold", |
| abs((implied_hour * 12) % 360 - implied_min) < 1e-3) |
|
|
| |
| |
| true = 425.0 |
| h, m = soft_targets(torch.tensor([true]), 180, 60.0) |
| vague = torch.full_like(h, 1.0 / h.shape[1]) |
| vague = 0.75 * vague + 0.25 * h |
| t, _, _, _ = joint_decode(L(vague), L(m)) |
| e = min(abs(t.item() - true), 720 - abs(t.item() - true)) |
| check("a vague hour hand plus a sharp minute hand still lands in the right hour", e < 1.0) |
|
|
| |
| h_flat = torch.full((1, 180), 1.0 / 180) |
| _, m = soft_targets(torch.tensor([true]), 180, 60.0) |
| without = joint_decode(L(h_flat), L(m))[0].item() |
| ring = ring_target(torch.tensor([true]), 144, 20.0) |
| with_ = joint_decode(L(h_flat), L(m), L(ring))[0].item() |
| e_wo = min(abs(without - true), 720 - abs(without - true)) |
| e_w = min(abs(with_ - true), 720 - abs(with_ - true)) |
| check("with no hour hand at all the hour is a guess", e_wo > 30) |
| check("the whole-time head recovers the hour the missing hand cannot give", e_w < 1.0) |
|
|
| |
| h, m = soft_targets(torch.tensor([190.0]), 180, 60.0) |
| check("swap resolution is off by default", not bool(joint_decode(L(m), L(h))[2][0])) |
| check("swap resolution still works when asked for", |
| bool(joint_decode(L(m), L(h), resolve_swap=True)[2][0])) |
|
|
| |
| h, m = soft_targets(torch.tensor([180.0]), 180, 60.0) |
| sharp = joint_decode(L(h), L(m))[1].item() |
| h2, m2 = soft_targets(torch.tensor([540.0]), 180, 60.0) |
| two = joint_decode(L(0.5 * h + 0.5 * h2), L(0.5 * m + 0.5 * m2))[1].item() |
| check("one clear reading scores a wider margin than two rival readings", sharp > two) |
|
|
| print(f"\n{ok}/{ok + fail} checks passed") |
| return fail == 0 |
|
|
|
|
| if __name__ == "__main__": |
| import sys |
| sys.exit(0 if _self_test() else 1) |
|
|