Coxeter Folding Reversible Recurrence / benchmark.py
Mechanism confirmed, baseline not beaten
1import json
2import numpy as np
3import torch
4from torch import nn
5
6torch.manual_seed(7)
7np.random.seed(7)
8
9
10def fold_real(x, j, n=8, eps=1e-6):
11 r, im = x[..., :n], x[..., n:]
12 out = x.clone()
13 a = torch.complex(r[..., (j-1) % n], im[..., (j-1) % n])
14 b = torch.complex(r[..., j], im[..., j])
15 c = torch.complex(r[..., (j+1) % n], im[..., (j+1) % n])
16 den = (b-c) - (a-b)
17 den = den + eps * (den.abs() < eps).to(den.dtype)
18 v = ((b-c)*a - (a-b)*c) / den
19 out[..., j] = v.real
20 out[..., n+j] = v.imag
21 return out
22
23
24def fold_chain(x, schedule):
25 for j in schedule:
26 x = fold_real(x, j)
27 return x
28
29
30def saved_bytes(fn, x):
31 saved = []
32 def pack(t):
33 saved.append(t.element_size() * t.numel())
34 return t
35 with torch.autograd.graph.saved_tensors_hooks(pack, lambda t: t):
36 y = fn(x)
37 loss = (y*y).mean()
38 loss.backward()
39 return sum(saved), len(saved), float(loss.detach())
40
41
42def one_run(device):
43 n, batch, steps = 8, 32, 32
44 schedule = [0, 2, 4, 6, 1, 3, 5, 7] * (steps // 8)
45 th = torch.arange(n, device=device) * (2*np.pi/n)
46 base = torch.cat([torch.cos(th), torch.sin(th)]).repeat(batch, 1)
47 x1 = (base + .01*torch.randn_like(base)).requires_grad_()
48 x2 = x1.detach().clone().requires_grad_()
49 gru = nn.GRU(2*n, 2*n, batch_first=True).to(device)
50 inp = x2[:, None, :].expand(batch, steps, 2*n).contiguous()
51 fold_mem, fold_tensors, _ = saved_bytes(lambda z: fold_chain(z, schedule), x1)
52 gru_mem, gru_tensors, _ = saved_bytes(lambda z: gru(z)[0][:, -1, :], inp)
53 with torch.no_grad():
54 y = fold_chain(x1.detach(), schedule)
55 z = fold_chain(y, list(reversed(schedule)))
56 recon = (torch.linalg.vector_norm(z-x1.detach()) /
57 torch.linalg.vector_norm(x1.detach())).item()
58 return {
59 'device': device, 'steps': steps, 'batch': batch,
60 'native_autograd_saved_bytes': {'fold': fold_mem, 'gru': gru_mem},
61 'native_autograd_saved_tensor_count': {'fold': fold_tensors, 'gru': gru_tensors},
62 'reverse_reconstruction_relative_error': recon,
63 'note': 'Fold measurement is native autograd; a custom reversible backward would be needed for constant-memory training.'
64 }
65
66
67def run():
68 preferred = 'cuda' if torch.cuda.is_available() else 'cpu'
69 try:
70 result = one_run(preferred)
71 except Exception as exc:
72 if preferred != 'cuda':
73 result = {'device': preferred, 'error': repr(exc)}
74 else:
75 result = {'cuda_error': repr(exc), 'fallback': one_run('cpu')}
76 with open('benchmark_results.json', 'w') as f:
77 json.dump(result, f, indent=2)
78 print(json.dumps(result, indent=2))
79
80if __name__ == '__main__':
81 run()