Certified Coarse-to-Fine Coordinate Refinement / registered_bench.py
Failed on benchmark
1import json
2import sys
3import numpy as np
4import torch
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
8
9TRACK = 'two_view_gauge_localization_v2'
10MODEL = 'mlp_tiny'
11EPOCHS = 12
12NTRAIN, NTEST = 400, 200
13
14class CoarseToFine(torch.nn.Module):
15 """Differentiable coarse proposal followed by one fixed refinement step.
16
17 The scalar target is the localized x-coordinate in the registered localization
18 track. The coarse grid is fixed in the known bounded target domain [-2, 2].
19 """
20 def __init__(self, base, temperature=18.0, refinement=0.35):
21 super().__init__()
22 self.base = base
23 self.temperature = temperature
24 self.refinement = refinement
25 self.register_buffer('grid', torch.linspace(-2.0, 2.0, nine if False else 9))
26
27 def forward(self, x):
28 raw = self.base(x).reshape(-1, 1)
29 dist = (raw - self.grid.reshape(1, -1)) ** 2
30 weights = torch.softmax(-self.temperature * dist, dim=1)
31 proposal = (weights * self.grid.reshape(1, -1)).sum(dim=1, keepdim=True)
32 return raw + self.refinement * (proposal - raw)
33
34
35def run_one(seed, lr, idea):
36 d = get_dataset(TRACK, int(seed), n_train=NTRAIN, n_test=NTEST)
37 torch.manual_seed(10000 + int(seed))
38 base = make_model(MODEL, d['input_shape'], d['out_dim'])
39 net = CoarseToFine(base) if idea else base
40 net, _, _ = train_model(net, d, epochs=EPOCHS, lr=float(lr), log=lambda *_: None)
41 if net is None:
42 return float('nan'), {'raw_mae': float('nan'), 'refined_mae': float('nan'), 'proposal_shift': float('nan')}
43 device = next(net.parameters()).device
44 xte, yte = d['xte'].to(device), d['yte'].to(device)
45 with torch.no_grad():
46 pred = net(xte)
47 if idea:
48 raw = net.base(xte)
49 refined_mae = float(torch.mean(torch.abs(pred - yte)).cpu())
50 raw_mae = float(torch.mean(torch.abs(raw - yte)).cpu())
51 shift = float(torch.mean(torch.abs(pred - raw)).cpu())
52 else:
53 refined_mae = float(torch.mean(torch.abs(pred - yte)).cpu())
54 raw_mae, shift = refined_mae, 0.0
55 mse = float(torch.mean((pred - yte) ** 2).cpu())
56 # The benchmark metric remains the task MSE; signature numbers are measured
57 # from predictions of the trained systems, not from an analytic toy identity.
58 return mse, {'raw_mae': raw_mae, 'refined_mae': refined_mae, 'proposal_shift': shift}
59
60
61def main():
62 grid = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}]
63 cache = {}
64 def baseline_factory(cfg):
65 return lambda seed: run_one(seed, cfg['lr'], False)[0]
66 base_block = sweep_baseline(baseline_factory, grid)
67 best_lr = float(base_block['best_cfg']['lr'])
68 idea_res = evaluate(lambda seed: run_one(seed, best_lr, True)[0])
69 nearby = {}
70 for lr in (1e-3, 1e-2):
71 r = evaluate(lambda seed, lr=lr: run_one(seed, lr, True)[0], seeds=(0,1,2,3))
72 nearby[str(lr)] = r
73 sig = evaluate(lambda seed: run_one(seed, best_lr, True)[1]['proposal_shift'])
74 raw = evaluate(lambda seed: run_one(seed, best_lr, True)[1]['raw_mae'])
75 refined = evaluate(lambda seed: run_one(seed, best_lr, True)[1]['refined_mae'])
76 extra = {
77 'mechanism_signature': {
78 'prediction': 'coarse proposal plus fixed refinement should reduce raw localization error',
79 'raw_mae_mean': raw['mean'],
80 'refined_mae_mean': refined['mean'],
81 'mean_proposal_shift': sig['mean'],
82 'confirmed': bool(np.isfinite(raw['mean']) and refined['mean'] < raw['mean'])
83 },
84 'idea_nearby_sweep': nearby,
85 'custom_track': {'name': TRACK, 'registered': True, 'domain': 'geometry/localization'}
86 }
87 report = make_report(TRACK, MODEL, base_block, idea_res, extra)
88 with open('registered_bench_report.json', 'w') as f:
89 json.dump(report, f, indent=2)
90 print(json.dumps(report, indent=2))
91
92if __name__ == '__main__':
93 main()