r"""
Scientific ML Studio -- generated code

Licence : Creative Commons Attribution-NonCommercial 4.0 International (CC BY-NC 4.0)
          https://creativecommons.org/licenses/by-nc/4.0/
(c) 2026 Scientific ML Studio. All rights are reserved to the Scientific ML Studio platform
except the permissions granted by the licence above.

You may share and adapt this code for NON-COMMERCIAL purposes, provided you credit
"Scientific ML Studio" with a link to https://scimlstudio.com, and indicate any changes you make.
Commercial use, paid client work and production deployment need prior written permission,
which you can ask for through the site's Contact page.
-------------------------------------------------------------------------------

Poisson 2D

Generated by the Scientific ML Studio Lab (engine 2.0) from a block diagram.

Geometry    : Rectangle (A 2-D box [x_min, x_max] x [y_min, y_max].)
Coordinates : (x, y)
Equation    : \nabla^2 u = f(\mathbf{x})
Conditions  : 4 boundary, 0 initial

Run it directly (python this_file.py), or import it and call train().
Command line: --epochs N  --device cpu|cuda  --checkpoint PATH  --no-plots
Run target  : auto (paste this file into a Colab cell or run it as a local script -- both work unmodified).
"""

from __future__ import annotations

import argparse
import json
import math
import time
from pathlib import Path

import torch
import torch.nn as nn

# TARGET_ENV is fixed at generation time from the Train block; IN_COLAB is
# decided at import time by whether `google.colab` is importable, so this
# same file runs unmodified as `python this_file.py` on a local machine and
# as a pasted-in cell on Colab, without anything to edit by hand.
TARGET_ENV = 'auto'  # "auto" | "local" | "colab"
try:
    import google.colab  # noqa: F401
    _COLAB_MODULE_PRESENT = True
except ImportError:
    _COLAB_MODULE_PRESENT = False
IN_COLAB = _COLAB_MODULE_PRESENT if TARGET_ENV == "auto" else TARGET_ENV == "colab"

try:
    import matplotlib
    if not IN_COLAB:
        matplotlib.use("Agg")  # headless-safe backend; Colab keeps its own inline one
    import matplotlib.pyplot as plt
except ImportError:  # plotting is optional; training is not
    plt = None


# --------------------------------------------------------------------------
# Configuration -- every value here came from a block on the canvas.
# --------------------------------------------------------------------------
TITLE = 'Poisson 2D'
SEED = 0
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CHECKPOINT = 'model.pt'

COORDS = ['x', 'y']
DIM = 2
SPATIAL_DIM = 2
X_MIN = 0.0
X_MAX = 1.0
Y_MIN = 0.0
Y_MAX = 1.0

N_INTERIOR = 2000
INTERIOR_STRATEGY = 'uniform'
N_INITIAL = 200

EPOCHS = 5000
LEARNING_RATE = 0.001
LOG_EVERY = 200
RESAMPLE_EVERY = 0


# --------------------------------------------------------------------------
# Output location
# --------------------------------------------------------------------------
OUTPUT_DIR = Path(".")
CHECKPOINT = str(OUTPUT_DIR / CHECKPOINT)


# --------------------------------------------------------------------------
# Sampling
#
# Every strategy produces points in the unit cube; the geometry maps them.
# That is what lets a Sobol sequence on a disc be the same code as a Sobol
# sequence on an interval.
# --------------------------------------------------------------------------
def unit_samples(n: int, dim: int, strategy: str = "uniform") -> torch.Tensor:
    """n x dim points in [0, 1]^dim."""
    n = max(int(n), 1)
    if strategy == "grid":
        per = max(int(round(n ** (1.0 / dim))), 2)
        axes = [torch.linspace(0.0, 1.0, per) for _ in range(dim)]
        mesh = torch.meshgrid(*axes, indexing="ij")
        return torch.stack([m.reshape(-1) for m in mesh], dim=1)
    if strategy == "sobol":
        engine = torch.quasirandom.SobolEngine(dimension=dim, scramble=True, seed=SEED)
        return engine.draw(n).to(torch.get_default_dtype())
    if strategy == "latin_hypercube":
        cols = []
        for _ in range(dim):
            perm = torch.randperm(n).to(torch.get_default_dtype())
            cols.append((perm + torch.rand(n)) / n)
        return torch.stack(cols, dim=1)
    return torch.rand(n, dim)


def as_points(columns: list[torch.Tensor], grad: bool = True) -> torch.Tensor:
    """Assemble coordinate columns into one leaf tensor on the run device."""
    pts = torch.cat([c.reshape(-1, 1) for c in columns], dim=1).to(DEVICE)
    return pts.requires_grad_(grad)


def sample_interior(n: int = N_INTERIOR, strategy: str = INTERIOR_STRATEGY) -> torch.Tensor:
    """Collocation points strictly inside the domain."""
    u = unit_samples(n, 2, strategy)
    x = X_MIN + (X_MAX - X_MIN) * u[:, 0]
    y = Y_MIN + (Y_MAX - Y_MIN) * u[:, 1]
    return as_points([x, y])


def sample_bc1(n: int = 80, strategy: str = 'uniform'):
    """BC1: x = x_min (left). Returns (points, outward normals)."""
    u = unit_samples(n, 1, strategy)
    x = torch.full((u.shape[0],), X_MIN)
    y = Y_MIN + (Y_MAX - Y_MIN) * u[:, 0]
    normals = torch.tensor([[-1.0, 0.0]], device=DEVICE).expand(u.shape[0], SPATIAL_DIM)
    return as_points([x, y]), normals


def sample_bc2(n: int = 80, strategy: str = 'uniform'):
    """BC2: x = x_max (right). Returns (points, outward normals)."""
    u = unit_samples(n, 1, strategy)
    x = torch.full((u.shape[0],), X_MAX)
    y = Y_MIN + (Y_MAX - Y_MIN) * u[:, 0]
    normals = torch.tensor([[1.0, 0.0]], device=DEVICE).expand(u.shape[0], SPATIAL_DIM)
    return as_points([x, y]), normals


def sample_bc3(n: int = 80, strategy: str = 'uniform'):
    """BC3: y = y_min (bottom). Returns (points, outward normals)."""
    u = unit_samples(n, 1, strategy)
    x = X_MIN + (X_MAX - X_MIN) * u[:, 0]
    y = torch.full((u.shape[0],), Y_MIN)
    normals = torch.tensor([[0.0, -1.0]], device=DEVICE).expand(u.shape[0], SPATIAL_DIM)
    return as_points([x, y]), normals


def sample_bc4(n: int = 80, strategy: str = 'uniform'):
    """BC4: y = y_max (top). Returns (points, outward normals)."""
    u = unit_samples(n, 1, strategy)
    x = X_MIN + (X_MAX - X_MIN) * u[:, 0]
    y = torch.full((u.shape[0],), Y_MAX)
    normals = torch.tensor([[0.0, 1.0]], device=DEVICE).expand(u.shape[0], SPATIAL_DIM)
    return as_points([x, y]), normals


# --------------------------------------------------------------------------
# Network
# --------------------------------------------------------------------------
LOWER = torch.tensor([0.0, 0.0])
UPPER = torch.tensor([1.0, 1.0])

class PINN(nn.Module):
    """Trial function u((x, y)) -> R^1."""

    def __init__(self):
        super().__init__()
        self.encoding = None
        width_in = DIM
        widths = [width_in] + [40, 40, 40, 40] + [1]
        layers = []
        for i in range(len(widths) - 1):
            layers.append(nn.Linear(widths[i], widths[i + 1]))
            if i < len(widths) - 2:
                layers.append(nn.Tanh())
        self.net = nn.Sequential(*layers)
        self.apply(self._init_weights)

    @staticmethod
    def _init_weights(module):
        if not isinstance(module, nn.Linear):
            return
        nn.init.xavier_normal_(module.weight)
        if module.bias is not None:
            nn.init.zeros_(module.bias)

    def forward(self, pts):
        lower = LOWER.to(pts.device)
        upper = UPPER.to(pts.device)
        z = 2.0 * (pts - lower) / torch.clamp(upper - lower, min=1e-12) - 1.0
        return self.net(z)


net = PINN().to(DEVICE)


def grad_of(y: torch.Tensor, pts: torch.Tensor) -> torch.Tensor:
    """d y / d pts, as one tensor with a column per coordinate."""
    if not y.requires_grad:
        return pts * 0.0
    gradient = torch.autograd.grad(
        y, pts, grad_outputs=torch.ones_like(y), create_graph=True, retain_graph=True, allow_unused=True
    )[0]
    return pts * 0.0 if gradient is None else gradient + pts * 0.0


def field(pts: torch.Tensor, component: int = 0) -> torch.Tensor:
    """One scalar output of the network, shaped (n, 1)."""
    return net(pts)[:, component:component + 1]


# --------------------------------------------------------------------------
# Loss terms
# --------------------------------------------------------------------------
def reduce_term(r: torch.Tensor) -> torch.Tensor:
    """Mean square residual."""
    return torch.mean(r ** 2)


def residual_pde1(pts: torch.Tensor) -> torch.Tensor:
    """Interior residual of the poisson equation; zero when it is satisfied."""
    u = field(pts, 0)
    x = pts[:, 0:1]
    y = pts[:, 1:2]
    g1 = grad_of(u, pts)
    u_x = g1[:, 0:1]
    u_y = g1[:, 1:2]
    g2_x = grad_of(u_x, pts)
    g2_y = grad_of(u_y, pts)
    u_xx = g2_x[:, 0:1]
    u_yy = g2_y[:, 1:2]
    f = ((((-2.0) * (math.pi ** 2.0)) * torch.sin((math.pi * x))) * torch.sin((math.pi * y)))
    return (u_xx + u_yy) - f


def loss_pde1(pts: torch.Tensor) -> torch.Tensor:
    return reduce_term(residual_pde1(pts))


def loss_bc1() -> torch.Tensor:
    """Dirichlet condition on BC1."""
    pts, normals = sample_bc1()
    u = field(pts, 0)
    target = 0.0
    return reduce_term(u - target)


def loss_bc2() -> torch.Tensor:
    """Dirichlet condition on BC2."""
    pts, normals = sample_bc2()
    u = field(pts, 0)
    target = 0.0
    return reduce_term(u - target)


def loss_bc3() -> torch.Tensor:
    """Dirichlet condition on BC3."""
    pts, normals = sample_bc3()
    u = field(pts, 0)
    target = 0.0
    return reduce_term(u - target)


def loss_bc4() -> torch.Tensor:
    """Dirichlet condition on BC4."""
    pts, normals = sample_bc4()
    u = field(pts, 0)
    target = 0.0
    return reduce_term(u - target)


# --------------------------------------------------------------------------
# Objective
# --------------------------------------------------------------------------
TERM_GROUPS = {'pde1': 'pde', 'bc1': 'bc', 'bc2': 'bc', 'bc3': 'bc', 'bc4': 'bc'}
TERM_LABELS = {'pde1': 'poisson residual', 'bc1': 'BC1 dirichlet', 'bc2': 'BC2 dirichlet', 'bc3': 'BC3 dirichlet', 'bc4': 'BC4 dirichlet'}
TERM_SCALES = {'pde1': 1.0, 'bc1': 100.0, 'bc2': 100.0, 'bc3': 100.0, 'bc4': 100.0}
GROUP_WEIGHTS = {
    "pde": 1.0,
    "bc": 1.0,
    "ic": 1.0,
}


def loss_terms(pts: torch.Tensor) -> dict[str, torch.Tensor]:
    """Every term of the objective, unweighted, keyed by name."""
    return {
        'pde1': loss_pde1(pts),
        'bc1': loss_bc1(),
        'bc2': loss_bc2(),
        'bc3': loss_bc3(),
        'bc4': loss_bc4(),
    }


class Balancer:
    """Fixed weights: the numbers from the Loss block, unchanged."""

    def __init__(self):
        self.weights = dict(GROUP_WEIGHTS)

    def update(self, epoch: int, terms: dict[str, torch.Tensor]) -> None:
        return


balancer = Balancer()


def total_loss(pts: torch.Tensor):
    """The weighted objective, plus the individual terms for logging."""
    terms = loss_terms(pts)
    total = torch.zeros((), device=DEVICE)
    for name, value in terms.items():
        group = TERM_GROUPS.get(name, "pde")
        total = total + balancer.weights.get(group, 1.0) * TERM_SCALES.get(name, 1.0) * value
    return total, terms


# --------------------------------------------------------------------------
# Training
# --------------------------------------------------------------------------
def build_optimiser(lr: float = LEARNING_RATE):
    return torch.optim.Adam(net.parameters(), lr=lr, weight_decay=0.0)


def build_scheduler(optimiser, epochs: int):
    return None


def train(epochs: int = EPOCHS, log_every: int = LOG_EVERY, verbose: bool = True) -> dict:
    """Run the optimisation and return the loss history."""
    torch.manual_seed(SEED)
    optimiser = build_optimiser()
    pts = sample_interior()
    history = {"epoch": [], "total": []}
    history['pde1'] = []
    history['bc1'] = []
    history['bc2'] = []
    history['bc3'] = []
    history['bc4'] = []
    started = time.time()

    for epoch in range(1, epochs + 1):
        if RESAMPLE_EVERY and epoch % RESAMPLE_EVERY == 0:
            pts = sample_interior()
        optimiser.zero_grad(set_to_none=True)
        total, terms_now = total_loss(pts)
        if total.requires_grad:
            total.backward()
        optimiser.step()
        balancer.update(epoch, terms_now)

        value = float(total.detach())
        history["epoch"].append(epoch)
        history["total"].append(value)
        for name, term in terms_now.items():
            history.setdefault(name, []).append(float(term.detach()))

        if verbose and (epoch % log_every == 0 or epoch == 1):
            parts = "  ".join(f"{TERM_LABELS.get(k, k)}={float(v.detach()):.3e}"
                for k, v in terms_now.items())
            print(f"epoch {epoch:>7d}  total={value:.6e}  {parts}")

    if verbose:
        print(f"finished in {time.time() - started:.1f}s")
    return history


def save_checkpoint(history: dict, path: str = CHECKPOINT) -> str:
    """Weights plus enough metadata to reload without the diagram."""
    payload = {
        "model": net.state_dict(),
        "coords": COORDS,
        "title": TITLE,
        "generator": "Omega Lab 2.0",
        "history": history,
    }
    torch.save(payload, path)
    Path(path).with_suffix(".history.json").write_text(json.dumps(history))
    return path


# --------------------------------------------------------------------------
# Visualisation
# --------------------------------------------------------------------------
FIGURE_DIR = OUTPUT_DIR / 'figures'
FIGURE_DPI = 150


def _finish(fig, name: str) -> None:
    """Save and/or show one figure, then release it."""
    if plt is None:
        return
    FIGURE_DIR.mkdir(parents=True, exist_ok=True)
    fig.savefig(FIGURE_DIR / f"{name}.png", dpi=FIGURE_DPI, bbox_inches="tight")
    plt.close(fig)


def grid_points(resolution: int = 200, t_value: float | None = None) -> tuple:
    """A regular grid over the domain, for plotting and tabulating."""
    x = torch.linspace(X_MIN, X_MAX, resolution)
    y = torch.linspace(Y_MIN, Y_MAX, resolution)
    gx, gy = torch.meshgrid(x, y, indexing="ij")
    cols = [gx.reshape(-1), gy.reshape(-1)]
    shape = (resolution, resolution)
    return as_points(cols, grad=False), shape


def plot_field() -> None:
    """The trained field over the domain."""
    if plt is None:
        return
    slices = [None]
    for t_value in slices:
        pts, shape = grid_points(t_value=t_value)
        with torch.no_grad():
            values = field(pts).reshape(shape).cpu()
        fig = plt.figure(figsize=(6.4, 4.2))
        ax = fig.add_subplot(111)
        gx = pts[:, 0].detach().cpu().reshape(shape)
        gy = pts[:, 1].detach().cpu().reshape(shape)
        mesh = ax.contourf(gx, gy, values, levels=50, cmap="viridis")
        fig.colorbar(mesh, ax=ax)
        ax.set_xlabel("x")
        ax.set_ylabel("y")
        label = TITLE if t_value is None else f"{TITLE}  (t = {t_value:g})"
        ax.set_title(label)
        name = "field" if t_value is None else f"field_t{t_value:g}"
        _finish(fig, name)


def plot_losses(history: dict) -> None:
    """Total objective and each term against epoch."""
    if plt is None:
        return
    fig, ax = plt.subplots(figsize=(6.4, 4.2))
    epochs = history.get("epoch", [])
    ax.plot(epochs, history.get("total", []), linewidth=2, label="total")
    for name, label in TERM_LABELS.items():
        series = history.get(name)
        if series:
            ax.plot(epochs[:len(series)], series, linewidth=1, alpha=0.75, label=label)
    ax.set_yscale("log")
    ax.set_xlabel("epoch")
    ax.set_ylabel("loss")
    ax.set_title("Convergence")
    ax.legend(fontsize=8)
    ax.grid(alpha=0.25)
    _finish(fig, "loss")


# --------------------------------------------------------------------------
# Report
# --------------------------------------------------------------------------
REPORT_PATH = str(OUTPUT_DIR / 'results.csv')
REPORT_DECIMALS = 6


def report_points() -> torch.Tensor:
    """The rows of the table."""
    return sample_interior()


def write_report(history: dict | None = None, path: str = REPORT_PATH) -> str:
    """Tabulate the field over the collocation points."""
    pts = report_points()
    values = net(pts)
    columns = ['x', 'y', 'u1']
    data = [pts[:, i].detach().cpu() for i in range(DIM)]
    data += [values[:, i].detach().cpu() for i in range(values.shape[1])]
    columns.append("residual")
    data.append(residual_pde1(pts).detach().cpu().reshape(-1))

    fmt = "{:.6g}"
    rows = []
    for i in range(pts.shape[0]):
        rows.append([fmt.format(float(col[i])) for col in data])

    lines = [",".join(columns)]
    lines += [",".join(row) for row in rows]
    lines.append("")
    lines.append("# summary")
    lines.append(f"# title,{TITLE}")
    lines.append(f"# rows,{len(rows)}")
    if history:
        for key, series in history.items():
            if key == "epoch" or not series:
                continue
            lines.append(f"# final_{key},{series[-1]:.6e}")
    Path(path).write_text("\n".join(lines) + "\n")
    return path


# --------------------------------------------------------------------------
# Entry point
# --------------------------------------------------------------------------
def main(argv: list[str] | None = None) -> None:
    parser = argparse.ArgumentParser(description=TITLE)
    parser.add_argument("--epochs", type=int, default=EPOCHS)
    parser.add_argument("--device", default=None, choices=["cpu", "cuda"])
    parser.add_argument("--checkpoint", default=CHECKPOINT)
    parser.add_argument("--log-every", type=int, default=LOG_EVERY)
    parser.add_argument("--quiet", action="store_true")
    parser.add_argument("--no-plots", action="store_true")
    args, _unknown = parser.parse_known_args(argv)

    global DEVICE
    if args.device:
        DEVICE = torch.device(args.device)
        net.to(DEVICE)

    history = train(epochs=args.epochs, log_every=args.log_every, verbose=not args.quiet)
    save_checkpoint(history, args.checkpoint)
    if not args.no_plots:
        plot_field()
        plot_losses(history)
    write_report(history)
    if not args.quiet:
        print(f"saved {args.checkpoint}")


if __name__ == "__main__":
    main()
