import json import numpy as np import torch from torch import nn torch.manual_seed(7) np.random.seed(7) def fold_real(x, j, n=8, eps=1e-6): r, im = x[..., :n], x[..., n:] out = x.clone() a = torch.complex(r[..., (j-1) % n], im[..., (j-1) % n]) b = torch.complex(r[..., j], im[..., j]) c = torch.complex(r[..., (j+1) % n], im[..., (j+1) % n]) den = (b-c) - (a-b) den = den + eps * (den.abs() < eps).to(den.dtype) v = ((b-c)*a - (a-b)*c) / den out[..., j] = v.real out[..., n+j] = v.imag return out def fold_chain(x, schedule): for j in schedule: x = fold_real(x, j) return x def saved_bytes(fn, x): saved = [] def pack(t): saved.append(t.element_size() * t.numel()) return t with torch.autograd.graph.saved_tensors_hooks(pack, lambda t: t): y = fn(x) loss = (y*y).mean() loss.backward() return sum(saved), len(saved), float(loss.detach()) def one_run(device): n, batch, steps = 8, 32, 32 schedule = [0, 2, 4, 6, 1, 3, 5, 7] * (steps // 8) th = torch.arange(n, device=device) * (2*np.pi/n) base = torch.cat([torch.cos(th), torch.sin(th)]).repeat(batch, 1) x1 = (base + .01*torch.randn_like(base)).requires_grad_() x2 = x1.detach().clone().requires_grad_() gru = nn.GRU(2*n, 2*n, batch_first=True).to(device) inp = x2[:, None, :].expand(batch, steps, 2*n).contiguous() fold_mem, fold_tensors, _ = saved_bytes(lambda z: fold_chain(z, schedule), x1) gru_mem, gru_tensors, _ = saved_bytes(lambda z: gru(z)[0][:, -1, :], inp) with torch.no_grad(): y = fold_chain(x1.detach(), schedule) z = fold_chain(y, list(reversed(schedule))) recon = (torch.linalg.vector_norm(z-x1.detach()) / torch.linalg.vector_norm(x1.detach())).item() return { 'device': device, 'steps': steps, 'batch': batch, 'native_autograd_saved_bytes': {'fold': fold_mem, 'gru': gru_mem}, 'native_autograd_saved_tensor_count': {'fold': fold_tensors, 'gru': gru_tensors}, 'reverse_reconstruction_relative_error': recon, 'note': 'Fold measurement is native autograd; a custom reversible backward would be needed for constant-memory training.' } def run(): preferred = 'cuda' if torch.cuda.is_available() else 'cpu' try: result = one_run(preferred) except Exception as exc: if preferred != 'cuda': result = {'device': preferred, 'error': repr(exc)} else: result = {'cuda_error': repr(exc), 'fallback': one_run('cpu')} with open('benchmark_results.json', 'w') as f: json.dump(result, f, indent=2) print(json.dumps(result, indent=2)) if __name__ == '__main__': run()