import json, time import torch import torch.nn as nn SEED = 2370 torch.manual_seed(SEED) torch.set_num_threads(4) DEVICE = "cuda" if torch.cuda.is_available() else "cpu" try: if DEVICE == "cuda": torch.cuda.set_device(0) torch.empty(1, device="cuda") except Exception: DEVICE = "cpu" def dx(a): return (torch.roll(a, -1, dims=-1) - torch.roll(a, 1, dims=-1)) * 0.5 def dy(a): return (torch.roll(a, -1, dims=-2) - torch.roll(a, 1, dims=-2)) * 0.5 def curl2(a): return torch.stack((dy(a), -dx(a)), dim=1) def divergence(b): return dx(b[:, 0]) + dy(b[:, 1]) class TinyNet(nn.Module): def __init__(self, out_channels): super().__init__() self.net = nn.Sequential( nn.Conv2d(1, 16, 3, padding=1, padding_mode="circular"), nn.Tanh(), nn.Conv2d(16, 16, 3, padding=1, padding_mode="circular"), nn.Tanh(), nn.Conv2d(16, out_channels, 1)) def forward(self, x): return self.net(x) def make_data(n, h=16, w=16): kx = torch.fft.fftfreq(w).reshape(1, 1, 1, w) ky = torch.fft.fftfreq(h).reshape(1, 1, h, 1) x = torch.randn(n, 1, h, w) f = torch.fft.fft2(x) filt = torch.exp(-((kx / .22) ** 2 + (ky / .22) ** 2)) a = torch.fft.ifft2(f * filt).real a = a / (a.std(dim=(-2, -1), keepdim=True) + 1e-7) return a, curl2(a[:, 0]) def train(kind, xtr, ytr, xte, yte, epochs=100): torch.manual_seed(SEED + (1 if kind == "curl" else 2)) net = TinyNet(1 if kind == "curl" else 2).to(DEVICE) opt = torch.optim.Adam(net.parameters(), lr=3e-3) xtr, ytr, xte, yte = [z.to(DEVICE) for z in (xtr, ytr, xte, yte)] for _ in range(epochs): pred = net(xtr) if kind == "curl": field = curl2(pred[:, 0]) loss = ((field - ytr) ** 2).mean() else: field = pred loss = ((field - ytr) ** 2).mean() + (divergence(field) ** 2).mean() opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): pred = net(xte) field = curl2(pred[:, 0]) if kind == "curl" else pred div = divergence(field) return {"mse": ((field - yte) ** 2).mean().item(), "max_abs_divergence": div.abs().max().item(), "rms_divergence": torch.sqrt((div ** 2).mean()).item(), "parameters": sum(p.numel() for p in net.parameters())} def run_models(xtr, ytr, xte, yte): return train("direct", xtr, ytr, xte, yte), train("curl", xtr, ytr, xte, yte) def main(): global DEVICE torch.manual_seed(SEED) a = torch.randn(3, 11, 13, dtype=torch.float64) exact_div = dx(dy(a)) - dy(dx(a)) generic = torch.randn(3, 2, 11, 13, dtype=torch.float64) gd = divergence(generic) identity = {"max_abs_DCA": exact_div.abs().max().item(), "rms_DCA": torch.sqrt((exact_div ** 2).mean()).item(), "generic_random_rms_divergence": torch.sqrt((gd ** 2).mean()).item()} xtr, ytr = make_data(192); xte, yte = make_data(64) t0 = time.time() try: baseline, idea = run_models(xtr, ytr, xte, yte) except Exception as exc: if DEVICE != "cuda": raise DEVICE = "cpu" torch.manual_seed(SEED) baseline, idea = run_models(xtr, ytr, xte, yte) identity["cuda_fallback"] = type(exc).__name__ + ": " + str(exc).split("\n")[0] result = {"device": DEVICE, "seed": SEED, "identity_check": identity, "baseline_direct_plus_penalty": baseline, "idea_exact_curl": idea, "elapsed_sec": time.time() - t0} print(json.dumps(result, indent=2)) with open("results.json", "w") as f: json.dump(result, f, indent=2) if __name__ == "__main__": main()