"""Reference PINN for the 1D Poisson guide:   -u''(x) = pi^2 sin(pi x),  u(0) = u(1) = 0.
Exact solution: u(x) = sin(pi x).  Run:  python poisson_1d_reference.py
(Written by hand to check the numbers quoted in the guide; it is not the Lab's generated script.)"""
import json, math, torch
torch.manual_seed(0)

# --- Domain: 100 interior collocation points, 2 boundary points -------------
N_PDE = 100
x_pde = torch.linspace(0, 1, N_PDE + 2)[1:-1].reshape(-1, 1).requires_grad_(True)
x_bc  = torch.tensor([[0.0], [1.0]])

# --- Network: 1 -> 32 -> 32 -> 32 -> 1, tanh --------------------------------
net = torch.nn.Sequential(
    torch.nn.Linear(1, 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))

def f(x):  return math.pi ** 2 * torch.sin(math.pi * x)
def exact(x): return torch.sin(math.pi * x)

def residual(x):
    u = net(x)
    du  = torch.autograd.grad(u,  x, torch.ones_like(u),  create_graph=True)[0]
    d2u = torch.autograd.grad(du, x, torch.ones_like(du), create_graph=True)[0]
    return -d2u - f(x)                       # = 0 where the equation holds

opt = torch.optim.Adam(net.parameters(), lr=1e-3)
x_test = torch.linspace(0, 1, 401).reshape(-1, 1)
hist = []
EPOCHS = 5000
for ep in range(EPOCHS + 1):
    opt.zero_grad()
    l_pde = residual(x_pde).pow(2).mean()
    l_bc  = net(x_bc).pow(2).mean()
    loss  = l_pde + l_bc                     # both weights = 1
    loss.backward(); opt.step()
    if ep % 50 == 0:
        with torch.no_grad():
            err = ((net(x_test) - exact(x_test)).norm() / exact(x_test).norm()).item()
        hist.append(dict(epoch=ep, loss=loss.item(), pde=l_pde.item(), bc=l_bc.item(), l2=err))
        if ep % 500 == 0:
            print(f"epoch {ep:5d}  loss {loss.item():.3e}  pde {l_pde.item():.3e}  bc {l_bc.item():.3e}  rel-L2 {err:.3e}")

with torch.no_grad():
    u = net(x_test); 
print("max |u - exact| =", (u - exact(x_test)).abs().max().item())
json.dump(dict(hist=hist, x=x_test.flatten().tolist(), u=u.flatten().tolist(),
               ex=exact(x_test).flatten().tolist(), pts=x_pde.detach().flatten().tolist()),
          open("poisson_1d_run.json", "w"))
