"""Reference run for the guide "Manufactured solutions" (2D Poisson, mixed boundaries).

Pick u*(x,y) = exp(x) sin(pi y) + x y, then let sympy work out everything else:
  f = laplacian(u*)            (Studio convention:  laplacian(u) = f)
  boundary data on each edge, including the normal derivative on x = 1.

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

# ---- 1. manufacture the problem with sympy ---------------------------------
x, y = sp.symbols("x y")
u_star = sp.exp(x) * sp.sin(sp.pi * y) + x * y
f_sym = sp.simplify(sp.diff(u_star, x, 2) + sp.diff(u_star, y, 2))
g_left = sp.simplify(u_star.subs(x, 0))
g_bottom = sp.simplify(u_star.subs(y, 0))
g_top = sp.simplify(u_star.subs(y, 1))
h_right = sp.simplify(sp.diff(u_star, x).subs(x, 1))       # outward normal on x = 1 is +x
print("f        =", f_sym)
print("left     u =", g_left, "| bottom u =", g_bottom, "| top u =", g_top)
print("right  du/dn =", h_right)

f_fn = sp.lambdify((x, y), f_sym, "numpy")
u_fn = sp.lambdify((x, y), u_star, "numpy")

# ---- 2. the PINN ----------------------------------------------------------
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 6000
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)

g = np.linspace(0.0, 1.0, 32)[1:-1]
XI, YI = np.meshgrid(g, g)                                  # 30 x 30 interior grid
Xr = torch.tensor(np.stack([XI.ravel(), YI.ravel()], 1), dtype=torch.float32)
Fr = torch.tensor(f_fn(Xr[:, 0].numpy(), Xr[:, 1].numpy()), dtype=torch.float32).reshape(-1, 1)
s = torch.linspace(0, 1, 50).reshape(-1, 1)
z, o = torch.zeros_like(s), torch.ones_like(s)
P_left = torch.cat([z, s], 1); P_bottom = torch.cat([s, z], 1)
P_top = torch.cat([s, o], 1); P_right = torch.cat([o, s], 1)
T_left = torch.sin(math.pi * s)
T_bottom = torch.zeros_like(s)
T_top = s.clone()                                            # u = x on y = 1
T_right = math.e * torch.sin(math.pi * s) + s                # du/dx on x = 1

def pde_loss():
    p = Xr.clone().requires_grad_(True)
    u = net(p)
    gr = torch.autograd.grad(u, p, torch.ones_like(u), create_graph=True)[0]
    uxx = torch.autograd.grad(gr[:, :1], p, torch.ones_like(gr[:, :1]), create_graph=True)[0][:, :1]
    uyy = torch.autograd.grad(gr[:, 1:], p, torch.ones_like(gr[:, 1:]), create_graph=True)[0][:, 1:]
    return ((uxx + uyy - Fr) ** 2).mean()

def bc_losses():
    d = ((net(P_left) - T_left) ** 2).mean() + ((net(P_bottom) - T_bottom) ** 2).mean() \
        + ((net(P_top) - T_top) ** 2).mean()
    p = P_right.clone().requires_grad_(True)
    ux = torch.autograd.grad(net(p), p, torch.ones(p.shape[0], 1), create_graph=True)[0][:, :1]
    n = ((ux - T_right) ** 2).mean()
    return d, n

ng = 101
xg, yg = np.meshgrid(np.linspace(0, 1, ng), np.linspace(0, 1, ng))
XY = torch.tensor(np.stack([xg.ravel(), yg.ravel()], 1), dtype=torch.float32)
UE = u_fn(xg.ravel(), yg.ravel())
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 = pde_loss(); ld, ln = bc_losses()
    loss = lp + w_bc * (ld + ln)
    if ep % 100 == 0 or ep == epochs:
        e, up = rel_err()
        hist.append((ep, loss.item(), lp.item(), ld.item(), ln.item(), e))
        if ep % 500 == 0:
            print(f"ep {ep:5d} loss {loss.item():.3e} pde {lp.item():.2e} dir {ld.item():.2e} neu {ln.item():.2e} relL2 {e:.3e}", flush=True)
    if ep < epochs:
        loss.backward(); opt.step()
e, up = rel_err()
print("final relL2", e, "max abs err", np.abs(up - UE).max())
# error on the Neumann edge only vs the rest
err = np.abs(up - UE).reshape(ng, ng)
print("max err on right edge", err[:, -1].max(), "| max err on left edge", err[:, 0].max())
np.savez(f"manuf_w{int(w_bc)}.npz", hist=np.array(hist), up=up, ue=UE, xg=xg, yg=yg)
