"""Reference run for the guide "Laplace's equation on a disc".

laplacian(u) = 0 on the unit disc x^2 + y^2 < 1,  u = exp(x) cos(y) on the circle.
Exact: u = exp(x) cos(y)  (it is harmonic: u_xx = u, u_yy = -u).

Hand-written for checking the guide; not the Studio's generated output.
Usage: python laplace_disc_reference.py [bc_weight] [epochs]
"""
import sys, math
import numpy as np, torch, torch.nn as nn

w_bc = float(sys.argv[1]) if len(sys.argv) > 1 else 100.0
epochs = int(sys.argv[2]) if len(sys.argv) > 2 else 5000
torch.manual_seed(0); np.random.seed(0)

net = nn.Sequential(nn.Linear(2, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh(),
                    nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 1))
for m in net:
    if isinstance(m, nn.Linear):
        nn.init.xavier_normal_(m.weight); nn.init.zeros_(m.bias)

exact = lambda x, y: np.exp(x) * np.cos(y)
Nr, Nb = 1000, 200
r = torch.sqrt(torch.rand(Nr, 1)); th = 2 * math.pi * torch.rand(Nr, 1)    # uniform in area
Pr = torch.cat([r * torch.cos(th), r * torch.sin(th)], 1)
thb = torch.linspace(0, 2 * math.pi, Nb + 1)[:-1].reshape(-1, 1)
Pb = torch.cat([torch.cos(thb), torch.sin(thb)], 1)
Tb = torch.exp(Pb[:, :1]) * torch.cos(Pb[:, 1:])

def losses():
    p = Pr.clone().requires_grad_(True)
    u = net(p)
    g = torch.autograd.grad(u, p, torch.ones_like(u), create_graph=True)[0]
    uxx = torch.autograd.grad(g[:, :1], p, torch.ones_like(g[:, :1]), create_graph=True)[0][:, :1]
    uyy = torch.autograd.grad(g[:, 1:], p, torch.ones_like(g[:, 1:]), create_graph=True)[0][:, 1:]
    return ((uxx + uyy) ** 2).mean(), ((net(Pb) - Tb) ** 2).mean()

ng = 201
xg, yg = np.meshgrid(np.linspace(-1, 1, ng), np.linspace(-1, 1, ng))
mask = (xg**2 + yg**2) <= 1.0
XY = torch.tensor(np.stack([xg[mask], yg[mask]], 1), dtype=torch.float32)
UE = exact(xg[mask], yg[mask])
def rel_err():
    with torch.no_grad(): up = net(XY).numpy().ravel()
    return float(np.linalg.norm(up - UE) / np.linalg.norm(UE)), up

opt = torch.optim.Adam(net.parameters(), lr=1e-3)
hist = []
for ep in range(epochs + 1):
    opt.zero_grad(); lp, lb = losses(); loss = lp + w_bc * lb
    if ep % 100 == 0 or ep == epochs:
        e, up = rel_err(); hist.append((ep, loss.item(), lp.item(), lb.item(), e))
        if ep % 500 == 0:
            print(f"ep {ep:5d} loss {loss.item():.3e} pde {lp.item():.2e} bc {lb.item():.2e} relL2 {e:.3e}", flush=True)
    if ep < epochs: loss.backward(); opt.step()
e, up = rel_err()
err = np.abs(up - UE)
rr = np.sqrt(xg[mask]**2 + yg[mask]**2)
print("final relL2", e, "max abs err", err.max(), "at r =", rr[err.argmax()])
print("max err r<0.5:", err[rr < 0.5].max(), "| max err r>0.9:", err[rr > 0.9].max())
np.savez(f"disc_w{int(w_bc)}.npz", hist=np.array(hist), up=up, ue=UE, xg=xg, yg=yg, mask=mask)
