"""Reference run for the guide "A shock wave: Burgers' equation".

u_t + u u_x = nu u_xx on x in [-1,1], t in [0,1], nu = 0.01/pi
u(x,0) = -sin(pi x),  u(-1,t) = u(1,t) = 0

No closed form is needed here: the Cole-Hopf transform turns Burgers into the
heat equation, and the resulting integral is evaluated by numerical quadrature
to give a reference solution (reference_burgers() below).

Hand-written for checking the guide; not the Studio's generated output.
Usage: python burgers_1d_reference.py [adam_epochs] [lbfgs_iters] [hidden_layers] [width] [n_interior]
"""
import sys, math
import numpy as np, torch, torch.nn as nn

NU = 0.01 / math.pi

def reference_burgers(x, t, n_eta=20001, half_width=4.0):
    """Cole-Hopf reference, x and t arrays of equal shape, t > 0 (t = 0 gives the initial data)."""
    x = np.asarray(x, float); t = np.asarray(t, float)
    out = np.empty_like(x)
    eta = np.linspace(-half_width, half_width, n_eta)
    for k, (xi, ti) in enumerate(zip(x.ravel(), t.ravel())):
        if ti <= 0:
            out.ravel()[k] = -math.sin(math.pi * xi); continue
        y = xi - eta
        expo = -np.cos(math.pi * y) / (2 * math.pi * NU) - eta**2 / (4 * NU * ti)
        w = np.exp(expo - expo.max())
        out.ravel()[k] = -np.trapezoid(np.sin(math.pi * y) * w, eta) / np.trapezoid(w, eta)
    return out

if __name__ == "__main__":
    adam_epochs = int(sys.argv[1]) if len(sys.argv) > 1 else 6000
    lbfgs_iters = int(sys.argv[2]) if len(sys.argv) > 2 else 500
    n_hidden = int(sys.argv[3]) if len(sys.argv) > 3 else 4
    width = int(sys.argv[4]) if len(sys.argv) > 4 else 32
    Nr = int(sys.argv[5]) if len(sys.argv) > 5 else 6000
    torch.manual_seed(0); np.random.seed(0)

    layers = [nn.Linear(2, width), nn.Tanh()]
    for _ in range(n_hidden - 1): layers += [nn.Linear(width, width), nn.Tanh()]
    net = nn.Sequential(*layers, nn.Linear(width, 1))
    for m in net:
        if isinstance(m, nn.Linear):
            nn.init.xavier_normal_(m.weight); nn.init.zeros_(m.bias)

    Nb, Ni = 200, 400
    xt = torch.rand(Nr, 2); xt[:, 0] = xt[:, 0] * 2 - 1
    xi = torch.rand(Ni, 1) * 2 - 1
    tb = torch.rand(Nb, 1)
    u_of = lambda x, t: 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); one = torch.ones_like(u)
        u_t = torch.autograd.grad(u, t, one, create_graph=True)[0]
        u_x = torch.autograd.grad(u, x, one, create_graph=True)[0]
        u_xx = torch.autograd.grad(u_x, x, one, create_graph=True)[0]
        l_pde = ((u_t + u * u_x - NU * u_xx) ** 2).mean()
        l_ic = ((u_of(xi, torch.zeros_like(xi)) + torch.sin(math.pi * xi)) ** 2).mean()
        l_bc = (u_of(-torch.ones_like(tb), tb) ** 2).mean() + (u_of(torch.ones_like(tb), tb) ** 2).mean()
        return l_pde, l_ic, l_bc

    ng_x, ng_t = 201, 101
    xg, tg = np.meshgrid(np.linspace(-1, 1, ng_x), np.linspace(0, 1, ng_t))
    XT = torch.tensor(np.stack([xg.ravel(), tg.ravel()], 1), dtype=torch.float32)
    UE = reference_burgers(xg, tg).ravel()
    def rel_err():
        with torch.no_grad(): up = net(XT).numpy().ravel()
        return float(np.linalg.norm(up - UE) / np.linalg.norm(UE)), up

    W_PDE, W_IC, W_BC = 1.0, 10.0, 10.0
    hist = []
    def total():
        lp, li, lb = losses()
        return W_PDE * lp + W_IC * li + W_BC * lb, lp, li, lb

    opt = torch.optim.Adam(net.parameters(), lr=1e-3)
    for ep in range(adam_epochs + 1):
        opt.zero_grad(); loss, lp, li, lb = total()
        if ep % 100 == 0 or ep == adam_epochs:
            e, _ = rel_err(); hist.append((ep, loss.item(), lp.item(), li.item(), lb.item(), e))
            if ep % 500 == 0:
                print(f"adam {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 < adam_epochs: loss.backward(); opt.step()
    e_adam, up_adam = rel_err()

    lb_opt = torch.optim.LBFGS(net.parameters(), lr=1.0, max_iter=lbfgs_iters, max_eval=int(lbfgs_iters * 1.25),
                               history_size=50, tolerance_grad=1e-12, tolerance_change=1e-14,
                               line_search_fn="strong_wolfe")
    def closure():
        lb_opt.zero_grad(); loss, *_ = total(); loss.backward(); return loss
    lb_opt.step(closure)
    e, up = rel_err()
    loss, lp, li, lb = total()
    hist.append((adam_epochs + lbfgs_iters, loss.item(), lp.item(), li.item(), lb.item(), e))
    print(f"after adam relL2 {e_adam:.3e} | after L-BFGS relL2 {e:.3e}  max abs err {np.abs(up-UE).max():.3e}")
    np.savez(sys.argv[6] if len(sys.argv) > 6 else "burgers.npz", hist=np.array(hist), up=up, up_adam=up_adam, ue=UE, xg=xg, tg=tg, e_adam=e_adam)
