r"""U-Net as a learned time-stepper: predict :math:`u^{n+1}` from :math:`u^n`. The training data is not manufactured here, it is SOLVED: a Q1 CG-FEM discretization of .. math:: \partial_t u - \nabla\cdot(D\,\nabla u) = f, \qquad u = 0 \text{ on } \partial(0,1)^2, \qquad u(\cdot, 0) = 0, stepped by an L-stable SDIRK2 (Pareschi-Russo). ``f`` is what varies across the dataset -- a mixture of Gaussians with random centres, amplitudes and widths, fixed in time within one trajectory -- so the map to learn is .. math:: (u^n, f) \longmapsto u^{n+1}, one step of the *implicit* solver. Both inputs are needed: ``u^{n+1}`` depends on ``f``, so a network shown only ``u^n`` would be fitting an ill-posed map. **The whole batch is solved by one ``jax.vmap``.** ``TimeDiscreteFEscheme`` is a ``ScimbaPytree``, and once the source is one too (``GaussianMixtureSource`` below, with ``jnp`` leaves rather than a closure), the scheme carries the sources as genuine pytree CHILDREN -- five leaves in total, since mesh, basis and tableau all live in ``aux_data``. Three practical points, each of which cost a debugging round: * Every callable handed to the weak form, the basis or the Dirichlet datum lands in ``aux_data``, compared by IDENTITY. A ``lambda`` written inside a factory function is a fresh object per call, so two otherwise identical schemes get different treedefs, share no executable, and cannot be stacked at all. They are hoisted to module level below, once, for that reason. * Map over the SOURCE, not over the scheme. Handing ``vmap`` a whole scheme as ``in_axes`` does not work (it wants a tree prefix and compares custom-node metadata), and flattening the scheme to dodge that is unnecessary anyway: ``uq_batched_transport_1d.py`` already has the idiom, which is to map a function that drops a different model into a ``copy.copy`` of one reference scheme. One pytree argument, ``in_axes=0``, nothing else to declare. * ``_before_solve`` runs once, on that reference. The shallow copies inherit its warmed boundary-geometry cache, so no stage recomputes ``boundary_dof_positions`` inside the trace -- which is a ``TracerArrayConversionError``. Measured on this problem (8 trajectories, 16 steps): the vmapped solve agrees with a Python loop over the same jitted function to ``1.1e-16``, and runs in 0.09 s against 0.15 s. Q1 nodal DOFs on a Cartesian mesh are numbered by ``np.unravel_index`` over the node grid (see ``VariablesFE.boundary_dof_positions``), so ``dofsl`` reshapes straight to the U-Net's grid -- no interpolation anywhere between the solver and the network. The script asserts this rather than trusting it. **The network learns the INCREMENT, and that is the whole difference between this working and not working.** One implicit step barely moves the solution -- ``u^{n+1} - u^n`` is some forty times smaller than ``u^n`` here -- so a network asked for ``u^{n+1}`` directly spends its capacity copying its own input. Measured, first version of this script: one-step relative error 1.85e-01, against 9.94e-02 for the do-nothing predictor ``u^{n+1} = u^n``. It was WORSE than doing nothing. Writing ``u^{n+1} = u^n + s * net(u^n, f)`` puts the copy in the architecture, where it is free and exact, and leaves the network only the part that is actually dynamics. **What is reported.** A one-step error is easy to make look good for the reason above, so the identity map is printed next to it as the control it deserves to be, and both are measured on ``u^{n+1}`` itself rather than on the scaled increment. Every one of them is given on TRAIN as well as TEST: their gap is the only thing that separates a network short on capacity from one that has memorised its training set, and the two look alike if only the test error is read (the convention of ``unet_source_to_solution.py``). It is worth having here, since 487k parameters against 768 training pairs is a heavily over-parameterised regime. The real question for a learned stepper, though, is whether it can be ITERATED, so the network is also rolled out for the full horizon from ``u^0 = 0`` against the FEM trajectory. Measured with the settings below: train test one-step error on u^{n+1}, U-Net 2.48e-03 2.79e-03 one-step error on u^{n+1}, identity 9.96e-02 9.94e-02 (36x worse) one-step error on the increment 2.49e-02 2.81e-02 rollout, 1 / 4 / 8 / 16 steps 4.09e-02, 2.28e-02, 2.02e-02, 1.98e-02 A test/train ratio of 1.13 on the increment says the over-parameterisation is not costing anything here: the network is not memorising its 768 pairs. What that rules IN is as useful as what it rules out -- the residual 2.8e-02 is a capacity or optimisation limit, not a generalisation one, so more epochs or a wider network would move it and more data would not. The rollout error falling and then flattening, rather than compounding, is the part that is not automatic: a stepper accurate for one step can still diverge when fed its own output. """ import copy import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_nd import HypercubeND from scimba_jax.linear_approximation.basis.analytic_bases import ( local_lagrange_basis, ) from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.fem.time_discrete_fe_scheme import ( TimeDiscreteFEscheme, ) from scimba_jax.linear_approximation.galerkin.time_discrete_galerkin_scheme import ( _ConstantFactory, ) from scimba_jax.linear_approximation.meshes.cartesian_mesh import cartesian_mesh from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.neural_operator.data_for_no.grid_data import GridData from scimba_jax.neural_operator.discrete_no.grid_based.unet import UNet from scimba_jax.nonlinear_approximation.numerical_solvers.no_projectors import ( NOProjector, ) from scimba_jax.physical_models.classical_weakform.diffusion_advection_reaction_weak_form import ( # noqa: E501 EllipticWeakForm, ) from scimba_jax.time_discrete.butcher_tableau import build_pareschi_russo_tableau from scimba_jax.utils.scimba_pytree import ScimbaPytree jax.config.update("jax_enable_x64", True) DIM = 2 N_NODES = 32 # grid side; N_NODES - 1 Q1 cells, so the nodes ARE the grid N_CELLS = N_NODES - 1 DIFFUSION = 0.05 DT, N_STEPS = 0.02, 16 # D*dt/h^2 = 1.0: one step diffuses about one cell N_GAUSSIANS = 3 N_TRAIN_TRAJ, N_TEST_TRAJ = 48, 12 LEVELS, BASE_CHANNELS = 3, 16 N_EPOCHS, BATCH_SIZE, LEARNING_RATE = 400, 32, 2.0e-3 # ── The parametric source, as a pytree ─────────────────────────────────────── class GaussianMixtureSource(ScimbaPytree): """``f(x) = sum_i A_i exp(-|x - c_i|^2 / 2 sigma_i^2)``, batchable. A ``ScimbaPytree`` rather than a closure: auto-classification looks for ``jnp.ndarray`` LEAVES, and a bare Python callable has none, so a closure would be filed as opaque ``aux_data`` -- static, un-batchable, and fatal the moment ``jax.vmap`` parks a tracer there. Holding the parameters as fields gives ``tree_flatten`` something real to find, which is exactly what makes one draw per batch element possible. Args: centers: Gaussian centres, shape ``(n_gaussians, dim)``. amplitudes: Per-Gaussian amplitude, shape ``(n_gaussians,)``. sigmas: Per-Gaussian width, shape ``(n_gaussians,)``. """ def __init__(self, centers, amplitudes, sigmas): self.centers = jnp.asarray(centers) self.amplitudes = jnp.asarray(amplitudes) self.sigmas = jnp.asarray(sigmas) def __call__(self, x: jnp.ndarray) -> jnp.ndarray: squared = jnp.sum((x - self.centers) ** 2, axis=1) return jnp.sum(self.amplitudes * jnp.exp(-squared / (2.0 * self.sigmas**2))) # ── Everything that lands in aux_data, built ONCE ──────────────────────────── # # See the module docstring: a fresh lambda per scheme is a fresh identity, so # the schemes would not share a treedef and could not be batched. def _q1_basis(y, i, m): """The Q1 local basis, as a named module-level function.""" return local_lagrange_basis(y, i, m, order=1, out_dim=1) def _diffusion_tensor(_x): """Constant isotropic diffusion ``D I``.""" return DIFFUSION * jnp.eye(DIM) def _no_advection(_x): return jnp.zeros(DIM) def _no_reaction(_x): return jnp.zeros(()) def _zero_dirichlet(_x): return jnp.zeros(1) _TABLEAU = build_pareschi_russo_tableau() def make_variables() -> VariablesFE: """The Q1 continuous FE space whose nodes are the U-Net's grid.""" mesh = cartesian_mesh( n_cells=(N_CELLS,) * DIM, quad_order=3, bounds=[(0.0, 1.0)] * DIM ) basis = AnalyticBasis( nb_basis=2**DIM, out_dim=1, mesh=mesh, local_basis=_q1_basis, basis_type="scalar", ) return VariablesFE(basis=basis, nb_variables=1) def weak_form_for(source: GaussianMixtureSource) -> EllipticWeakForm: """Pure diffusion driven by this source: ``a(u, v) = int D grad u . grad v``.""" return EllipticWeakForm( dim=DIM, A=_diffusion_tensor, b=_no_advection, c=_no_reaction, f=source, ) def make_scheme(source: GaussianMixtureSource, variables: VariablesFE): """The implicit FEM time-stepper for one source.""" return TimeDiscreteFEscheme( spatial_weak_form_factory=weak_form_for(source), variables=variables, butcher_tableau=_TABLEAU, dt=DT, dirichlet=_zero_dirichlet, tol=1e-10, max_iter=400, ) # ── Solving the whole dataset in one vmapped program ───────────────────────── def draw_sources(key, n_traj: int): """Random Gaussian-mixture parameters, one mixture per trajectory.""" key_c, key_a, key_s = jax.random.split(key, 3) centers = jax.random.uniform( key_c, (n_traj, N_GAUSSIANS, DIM), minval=0.2, maxval=0.8 ) amplitudes = jax.random.uniform( key_a, (n_traj, N_GAUSSIANS), minval=0.5, maxval=2.0 ) sigmas = jax.random.uniform(key_s, (n_traj, N_GAUSSIANS), minval=0.08, maxval=0.18) return centers, amplitudes, sigmas def solve_trajectories(variables, centers, amplitudes, sigmas): """Every trajectory, solved together by one ``jax.vmap`` over the SOURCE. Follows ``uq_batched_transport_1d.py``'s idiom for batching a solve over a family of models: build ONE reference scheme, then map a function that drops a different model into a ``copy.copy`` of it -- the same FE space, the same warmed caches, only the physics swapped. Mapping over a single pytree argument means ``in_axes=0`` and nothing else to say. The shallow copy is what keeps this simple: ``_before_solve``'s boundary-geometry cache is warmed once on the reference and inherited by every copy, so no stage ever recomputes ``boundary_dof_positions`` inside the trace (which is a ``TracerArrayConversionError``). Args: variables: The FE space, shared by every trajectory. centers, amplitudes, sigmas: Batched source parameters, leading axis ``n_traj``. Returns: ``(n_traj, N_STEPS + 1, n_dof, 1)`` DOFs, ``u^0`` included. """ n_traj = centers.shape[0] dofsl_init = jnp.zeros_like(variables.dofsl) reference = make_scheme( GaussianMixtureSource(centers[0], amplitudes[0], sigmas[0]), variables ) reference._before_solve(dofsl_init, 0.0) # (nt, verbose, no_tqdm, keep_history) -- no_tqdm=True because a progress # bar per trajectory is noise when the whole batch is one vmapped call. run = TimeDiscreteFEscheme._make_solve_fn(N_STEPS, False, True, True) def solve_one(source): # ⚠ `copy.copy`, not a fresh scheme: the same space and the same warm # caches, with only the weak form's source replaced -- exactly what # `uq_batched_transport_1d.py` does with `scheme.pde = pde`. scheme = copy.copy(reference) scheme.spatial_weak_form_factory = _ConstantFactory(weak_form_for(source)) return run(scheme, dofsl_init, 0.0) _, history_tail = jax.vmap(solve_one)( GaussianMixtureSource(centers, amplitudes, sigmas) ) initial = jnp.broadcast_to(dofsl_init, (n_traj, 1, *dofsl_init.shape)) return jnp.concatenate([initial, history_tail], axis=1) def source_fields(node_points, centers, amplitudes, sigmas): """Every source sampled on the node grid, shape ``(n_traj, n_dof)``.""" return jax.vmap( lambda c, a, s: jax.vmap(GaussianMixtureSource(c, a, s))(node_points) )(centers, amplitudes, sigmas) def node_grid_positions(variables) -> jnp.ndarray: """The FE node positions, in DOF order, shape ``(n_dof, dim)``. Straight from the library rather than re-derived: for nodal DOFs ``VariablesFE.boundary_dof_positions`` is a plain ``np.unravel_index`` over the node grid (it takes any indices, not only boundary ones), and it is the definition of the convention this script depends on. """ n_dof = variables.dofsl.shape[0] return jnp.asarray(variables.boundary_dof_positions(np.arange(n_dof))) def to_grid(dofs: jnp.ndarray) -> jnp.ndarray: """``(..., n_dof, 1)`` DOFs -> ``(..., N_NODES, N_NODES)`` fields.""" return dofs[..., 0].reshape(*dofs.shape[:-2], N_NODES, N_NODES) def relative_error(prediction, target) -> float: """Relative L2 error over the whole set.""" return float(jnp.linalg.norm(prediction - target) / jnp.linalg.norm(target)) # ── Build the dataset ──────────────────────────────────────────────────────── print("=" * 74) print("U-NET AS A LEARNED TIME-STEPPER (u^n, f) -> u^{n+1}") print("=" * 74) variables = make_variables() n_dof = variables.dofsl.shape[0] assert n_dof == N_NODES**2, f"expected {N_NODES**2} DOFs, got {n_dof}" node_points = node_grid_positions(variables) # `to_grid` reshapes DOFs straight into the U-Net's grid, which is only valid # if DOF k really is node (k // N, k % N) of a C-ordered N x N grid. Assert it # rather than trust it: a silent transpose here would train the network on a # scrambled field and still converge to something plausible-looking. _axis = np.linspace(0.0, 1.0, N_NODES) _expected = np.stack(np.meshgrid(_axis, _axis, indexing="ij"), axis=-1).reshape(-1, DIM) _mismatch = float(np.max(np.abs(np.asarray(node_points) - _expected))) assert _mismatch < 1e-12, f"DOF order is not the node grid order ({_mismatch:.2e})" print( f" DOF order matches a C-ordered {N_NODES}x{N_NODES} node grid ({_mismatch:.0e})" ) key = jax.random.PRNGKey(0) key, key_train, key_test, key_net, key_fit = jax.random.split(key, 5) started = time.time() c_tr, a_tr, s_tr = draw_sources(key_train, N_TRAIN_TRAJ) c_te, a_te, s_te = draw_sources(key_test, N_TEST_TRAJ) traj_train = solve_trajectories(variables, c_tr, a_tr, s_tr) traj_test = solve_trajectories(variables, c_te, a_te, s_te) jax.block_until_ready(traj_train) print( f" {N_TRAIN_TRAJ + N_TEST_TRAJ} trajectories x {N_STEPS} implicit steps " f"on {N_CELLS}^2 Q1 cells, one vmap each: {time.time() - started:.1f}s" ) f_train = to_grid(source_fields(node_points, c_tr, a_tr, s_tr)[..., None]) f_test = to_grid(source_fields(node_points, c_te, a_te, s_te)[..., None]) u_train, u_test = to_grid(traj_train), to_grid(traj_test) # Every consecutive pair of every trajectory is one training example. # # ⚠ The network learns the INCREMENT, not the next state. One implicit step # barely moves the solution (u^{n+1} - u^n is ~40x smaller than u^n here), so a # network asked for u^{n+1} spends its whole capacity copying its own input and # is beaten by the do-nothing predictor -- measured, before this was changed: # 1.85e-01 against the identity's 9.94e-02. Predicting u^{n+1} = u^n + s * net # puts the copy in the architecture, where it is free and exact, and leaves the # network only the part that is actually dynamics. scale_f = float(jnp.std(f_train)) scale_du = float(jnp.std(u_train[:, 1:] - u_train[:, :-1])) def make_pairs(u_traj, f_field): """``(u^n, f) -> (u^{n+1} - u^n) / scale_du`` over every step.""" n_traj = u_traj.shape[0] current = u_traj[:, :-1].reshape(-1, N_NODES, N_NODES, 1) increment = (u_traj[:, 1:] - u_traj[:, :-1]).reshape(-1, N_NODES, N_NODES, 1) source = jnp.broadcast_to( (f_field / scale_f)[:, None, :, :, None], (n_traj, N_STEPS, N_NODES, N_NODES, 1), ).reshape(-1, N_NODES, N_NODES, 1) return jnp.concatenate([current, source], axis=-1), increment / scale_du def step_with(projector, state, source_channel): """One learned step: ``u^{n+1} = u^n + scale_du * net(u^n, f)``.""" prediction = projector.evaluate(jnp.concatenate([state, source_channel], axis=-1)) return state + scale_du * prediction inputs_train, targets_train = make_pairs(u_train, f_train) inputs_test, targets_test = make_pairs(u_test, f_test) print( f" {inputs_train.shape[0]} training pairs, {inputs_test.shape[0]} test pairs, " f"grid {N_NODES}x{N_NODES}, u in [{float(u_train.min()):.3f}, " f"{float(u_train.max()):.3f}]" ) # ── What the network is asked to learn, before any training ────────────────── def plot_dataset(f_field, u_traj, filename: str, n_shown: int = 4) -> None: """One row per trajectory: its source, then the state at three times. Worth looking at before reading any error: it shows the VARIETY the network has to cover (each row is a different Gaussian mixture) and how little one step moves the solution compared with how much the whole horizon does -- which is the reason the network predicts the increment rather than the state. The colour scale is shared across each row's three states, so the growth between columns is real and not a rescaling. Args: f_field: ``(n, N, N)`` sources. u_traj: ``(n, N_STEPS + 1, N, N)`` trajectories. filename: Where to save the figure. n_shown: How many trajectories to draw. """ steps = [1, N_STEPS // 2, N_STEPS] fig, axes = plt.subplots(n_shown, 4, figsize=(15, 3.5 * n_shown)) for row in range(n_shown): source = np.asarray(f_field[row]) image = axes[row, 0].imshow( source.T, origin="lower", extent=(0, 1, 0, 1), cmap="viridis" ) axes[row, 0].set_title(f"source $f$ (trajectory {row})") axes[row, 0].set_ylabel("$y$") fig.colorbar(image, ax=axes[row, 0]) vmax = float(np.max(np.asarray(u_traj[row, steps[-1]]))) for column, step in enumerate(steps, start=1): field = np.asarray(u_traj[row, step]) image = axes[row, column].imshow( field.T, origin="lower", extent=(0, 1, 0, 1), cmap="magma", vmin=0.0, vmax=vmax, ) axes[row, column].set_title(f"$u$ at $t = {step * DT:.2f}$ ({step} steps)") fig.colorbar(image, ax=axes[row, column]) for ax in axes[row]: ax.set_xlabel("$x$") fig.suptitle( "The solutions being learned: FEM trajectories driven by random " "Gaussian-mixture sources\n" "(one row per trajectory, colour shared across each row's three states)" ) fig.tight_layout() fig.savefig(filename, dpi=200) print(f" dataset figure written to {filename}") plot_dataset(f_train, u_train, "unet_time_stepper_dataset.png") # ── Train ──────────────────────────────────────────────────────────────────── grid = GridData(DIM, HypercubeND([(0.0, 1.0)] * DIM), (N_NODES,) * DIM) unet = UNet( grid, 2, 1, key_net, levels=LEVELS, base_channels=BASE_CHANNELS, use_coordinates=True, ) print(f" U-Net: {unet.ndof()} parameters, {LEVELS} levels\n") projector = NOProjector( unet, (inputs_train, targets_train), learning_rate=LEARNING_RATE ) _, projector = projector.project( key_fit, unet, N_EPOCHS, batch_size=BATCH_SIZE, tqdm_desc="time-stepper" ) # ── One step: TRAIN and TEST, against the control that costs nothing ───────── def one_step_errors(inputs, targets) -> dict: """One-step errors on ``u^{n+1}``, plus the do-nothing control. Reported on ``u^{n+1}`` itself and never on the scaled increment, so that the network and the identity predictor are measured on the same thing. Args: inputs: ``(n, N, N, 2)`` stacked ``[u^n, f / scale_f]``. targets: ``(n, N, N, 1)`` scaled increments. Returns: ``unet`` and ``identity`` errors on ``u^{n+1}``, and ``increment``, the error on the part the network actually predicts. """ current = inputs[..., :1] next_true = current + scale_du * targets next_unet = current + scale_du * projector.evaluate(inputs) return dict( unet=relative_error(next_unet, next_true), identity=relative_error(current, next_true), increment=relative_error(next_unet - current, next_true - current), ) # Train AND test, because their GAP is the only thing that separates a network # short on capacity from one that has memorised its training set -- and the two # look alike if only the test error is read. Same convention as # ``unet_source_to_solution.py``. It matters here: 487k parameters against 768 # training pairs is a heavily over-parameterised regime, so the gap is the # first number to look at before tuning anything else. errors_train = one_step_errors(inputs_train, targets_train) errors_test = one_step_errors(inputs_test, targets_test) print(f"\n final loss : {float(projector.best_loss['total']):.3e}") print(f" {'one-step relative error':34s} {'train':>10s} {'test':>10s}") for label, field in [ ("on u^n+1, U-Net", "unet"), ("on u^n+1, identity (control)", "identity"), ("on the INCREMENT the net predicts", "increment"), ]: print(f" {label:34s} {errors_train[field]:10.3e} {errors_test[field]:10.3e}") print( f" {'test / train on the increment':34s} " f"{errors_test['increment'] / errors_train['increment']:10.2f}" ) # ── Rollout: the question a one-step error cannot answer ───────────────────── print("\n rolled out from u^0 = 0, against the FEM trajectory:") state = jnp.zeros((N_TEST_TRAJ, N_NODES, N_NODES, 1)) source_channel = (f_test / scale_f)[..., None] rollout = [state] for _ in range(N_STEPS): state = step_with(projector, state, source_channel) rollout.append(state) rollout = jnp.stack(rollout, axis=1)[..., 0] for step in (1, N_STEPS // 4, N_STEPS // 2, N_STEPS): print( f" after {step:3d} steps (t = {step * DT:.2f}): " f"relative error {relative_error(rollout[:, step], u_test[:, step]):.3e}" ) # ── Plot ───────────────────────────────────────────────────────────────────── sample = 0 steps_shown = [N_STEPS // 4, N_STEPS // 2, N_STEPS] fig, axes = plt.subplots(3, len(steps_shown), figsize=(4.2 * len(steps_shown), 11)) for column, step in enumerate(steps_shown): reference = np.asarray(u_test[sample, step]) predicted = np.asarray(rollout[sample, step]) vmax = float(np.max(reference)) for row, (field, title, cmap) in enumerate( [ (reference, "FEM (reference)", "magma"), (predicted, "U-Net rollout", "magma"), (predicted - reference, "difference", "coolwarm"), ] ): ax = axes[row, column] limits = ( dict(vmin=0.0, vmax=vmax) if row < 2 else dict(vmin=-0.1 * vmax, vmax=0.1 * vmax) ) image = ax.imshow( field.T, origin="lower", extent=(0, 1, 0, 1), cmap=cmap, **limits ) ax.set_title(f"{title}\n$t = {step * DT:.2f}$ ({step} steps)") fig.colorbar(image, ax=ax) fig.suptitle( f"U-Net time-stepper vs. the implicit FEM solver it learned from " f"(rolled out from $u^0 = 0$, {N_NODES}x{N_NODES} Q1 nodes)" ) fig.tight_layout() fig.savefig("unet_time_stepper_from_fem.png", dpi=200) print("\nfigure written to unet_time_stepper_from_fem.png") plt.show()