Coxeter Folding Reversible Recurrence / benchmark.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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()