"""Reference run for the guide "Heat in a bar" (1D heat equation).

u_t = a u_xx on x in [0,1], t in [0,1], a = 0.1
u(x,0) = sin(pi x),  u(0,t) = u(1,t) = 0
Exact: u = exp(-a pi^2 t) sin(pi x)

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

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

def exact(x, t):
    return np.exp(-A * math.pi**2 * t) * np.sin(math.pi * x)

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)

# training points (fixed for the whole run)
Nr, Nb, Ni = 2000, 100, 100
xt = torch.rand(Nr, 2)                               # (x, t) interior, incl. t in (0,1)
xi = torch.rand(Ni, 1); ti = torch.zeros(Ni, 1)      # initial line t = 0
tb = torch.rand(Nb, 1)
xl = torch.zeros(Nb, 1); xr = torch.ones(Nb, 1)

def u_of(x, t):
    return net(torch.cat([x, t], 1))

def losses():
    x = xt[:, :1].clone().requires_grad_(True); t = xt[:, 1:].clone().requires_grad_(True)
    u = u_of(x, t)
    ones = torch.ones_like(u)
    u_t = torch.autograd.grad(u, t, ones, create_graph=True)[0]
    u_x = torch.autograd.grad(u, x, ones, create_graph=True)[0]
    u_xx = torch.autograd.grad(u_x, x, torch.ones_like(u_x), create_graph=True)[0]
    l_pde = ((u_t - A * u_xx) ** 2).mean()
    l_ic = ((u_of(xi, ti) - torch.sin(math.pi * xi)) ** 2).mean()
    l_bc = (u_of(xl, tb) ** 2).mean() + (u_of(xr, tb) ** 2).mean()
    return l_pde, l_ic, l_bc

# evaluation grid, different from training points
ng = 101
xg, tg = np.meshgrid(np.linspace(0, 1, ng), np.linspace(0, 1, ng))
XT = torch.tensor(np.stack([xg.ravel(), tg.ravel()], 1), dtype=torch.float32)
UE = exact(xg.ravel(), tg.ravel())

def rel_err():
    with torch.no_grad():
        up = net(XT).numpy().ravel()
    return np.linalg.norm(up - UE) / np.linalg.norm(UE), up

opt = torch.optim.Adam(net.parameters(), lr=1e-3)
hist = []; best = (1e9, None)
for ep in range(epochs + 1):
    opt.zero_grad()
    lp, li, lb = losses()
    loss = lp + w_bc * (li + lb)
    if ep % 100 == 0 or ep == epochs:
        e, up = rel_err()
        hist.append((ep, loss.item(), lp.item(), li.item(), lb.item(), float(e)))
        if e < best[0]: best = (float(e), up.copy())
        if ep % 500 == 0: print(f"ep {ep:5d} loss {loss.item():.3e} pde {lp.item():.2e} ic {li.item():.2e} bc {lb.item():.2e} relL2 {e:.3e}", flush=True)
    if ep < epochs:
        loss.backward(); opt.step()

e, up = rel_err()
print("final relL2", e, "best", best[0], "max abs err", np.abs(up - UE).max())
np.savez(f"heat_w{int(w_bc)}.npz", hist=np.array(hist), up=up, ue=UE, best=best[1], xg=xg, tg=tg)
