"""Reference PINN for the 2D Poisson guide:  -(u_xx + u_yy) = 2 pi^2 sin(pi x) sin(pi y) on (0,1)^2,
u = 0 on the four edges.  Exact solution: u = sin(pi x) sin(pi y).
Run: python poisson_2d_reference.py     (hand-written to check the guide's numbers; not the Lab's generated script)"""
import json, math, sys, torch
torch.manual_seed(0)
EPOCHS = int(sys.argv[1]) if len(sys.argv) > 1 else 10000
BC_W   = float(sys.argv[2]) if len(sys.argv) > 2 else 1.0
NI, NB = 30, 50                                   # 30x30 interior grid, 50 points per edge
g = torch.linspace(0, 1, NI + 2)[1:-1]
X, Y = torch.meshgrid(g, g, indexing="ij")
xy = torch.stack([X.flatten(), Y.flatten()], 1).requires_grad_(True)       # 900 interior points
t = torch.linspace(0, 1, NB).reshape(-1, 1); z, o = torch.zeros_like(t), torch.ones_like(t)
xy_bc = torch.cat([torch.cat([t, z], 1), torch.cat([t, o], 1), torch.cat([z, t], 1), torch.cat([o, t], 1)])  # 200 points

net = torch.nn.Sequential(torch.nn.Linear(2, 32), torch.nn.Tanh(), torch.nn.Linear(32, 32), torch.nn.Tanh(),
                          torch.nn.Linear(32, 32), torch.nn.Tanh(), torch.nn.Linear(32, 1))
f     = lambda p: 2 * math.pi ** 2 * torch.sin(math.pi * p[:, :1]) * torch.sin(math.pi * p[:, 1:])
exact = lambda p: torch.sin(math.pi * p[:, :1]) * torch.sin(math.pi * p[:, 1:])

def residual(p):
    u = net(p)
    du = torch.autograd.grad(u, p, torch.ones_like(u), create_graph=True)[0]
    uxx = torch.autograd.grad(du[:, :1], p, torch.ones_like(du[:, :1]), create_graph=True)[0][:, :1]
    uyy = torch.autograd.grad(du[:, 1:], p, torch.ones_like(du[:, 1:]), create_graph=True)[0][:, 1:]
    return -(uxx + uyy) - f(p)

n = 101; gt = torch.linspace(0, 1, n); XT, YT = torch.meshgrid(gt, gt, indexing="ij")
test = torch.stack([XT.flatten(), YT.flatten()], 1)
opt = torch.optim.Adam(net.parameters(), lr=1e-3)
hist = []
for ep in range(EPOCHS + 1):
    opt.zero_grad()
    lp = residual(xy).pow(2).mean(); lb = net(xy_bc).pow(2).mean()
    loss = lp + BC_W * lb; loss.backward(); opt.step()
    if ep % 100 == 0:
        with torch.no_grad():
            e = ((net(test) - exact(test)).norm() / exact(test).norm()).item()
        hist.append(dict(epoch=ep, loss=loss.item(), pde=lp.item(), bc=lb.item(), l2=e))
        if ep % 1000 == 0: print(f"epoch {ep:6d} loss {loss.item():.3e} pde {lp.item():.3e} bc {lb.item():.3e} relL2 {e:.3e}")
with torch.no_grad():
    U = net(test).reshape(n, n); UE = exact(test).reshape(n, n)
print("max abs err", (U - UE).abs().max().item())
json.dump(dict(hist=hist, n=n, u=U.tolist(), ue=UE.tolist(), epochs=EPOCHS, bc_w=BC_W), open("poisson_2d_run.json", "w"))
