r"""2D wave equation as a genuinely hyperbolic first-order system, with TimeDiscreteCollocationScheme. .. math:: \partial_t v + \nabla \cdot \mathbf{w} &= 0, \\ \partial_t \mathbf{w} + c^2 \nabla v &= 0, on :math:`(x, y) \in (0, 1)^2`, homogeneous Dirichlet on :math:`v` only (:math:`v = 0` on :math:`\partial\Omega`; :math:`\mathbf{w}` is left unconstrained there, see below), with a Gaussian pulse initial condition -- the same 2D test case as ``examples_jax/pinns/time_dependent_pdes/hyperbolic_systems/wave_nd.py``: .. math:: u_0(x, y) = \exp\!\bigl(-75\,((x-\tfrac12)^2+(y-\tfrac12)^2)\bigr), \qquad \partial_t u_0 = 0, translated into this file's :math:`(v, \mathbf{w})` unknowns as :math:`v_0 = \partial_t u_0 = 0`, :math:`\mathbf{w}_0 = -c^2 \nabla u_0` (see below for why). **Where this system comes from, and the sign that makes it hyperbolic.** Introducing :math:`v = \partial_t u` and :math:`\mathbf{w} = -c^2 \nabla u` for the second-order wave equation :math:`\partial_{tt} u = c^2 \Delta u` gives exactly the system above: :math:`\partial_t v = \partial_{tt} u = c^2 \Delta u = -\nabla\cdot\mathbf{w}`, and :math:`\partial_t \mathbf{w} = -c^2 \nabla \partial_t u = -c^2 \nabla v`. Plane-wave analysis (:math:`v, \mathbf{w} \propto e^{i(\mathbf{k}\cdot\mathbf{x}-\omega t)}`) confirms it is genuinely hyperbolic with the correct wave speed -- :math:`\omega = \pm c|\mathbf{k}|`, real -- *only* with this relative sign between the two equations (both ``+``, as written above); flipping the sign of the second equation flips the dispersion relation to :math:`\omega^2 = -c^2|\mathbf{k}|^2`, i.e. purely imaginary frequencies, not a wave at all. Unlike an earlier version of this file (built around :math:`\partial_t u = v`, :math:`\partial_t v = c^2 \Delta u`), the spatial operator here is a genuine first-order flux (:math:`\nabla\cdot`/:math:`\nabla`, not a Laplacian) acting on a 3-vector state :math:`(v, w_1, w_2)` -- the standard first-order acoustic-wave reduction, not merely "first order in time". **Not the same as this codebase's own ``WaveResidual``/``WaveND``** (:mod:`scimba_jax.physical_models.temporal_pde.wave_equations`): those write :math:`u_{tt} - \Delta u = f` directly, with :math:`u_{tt}` a *second* time derivative of a single scalar unknown (via :meth:`~scimba_jax.nonlinear_approximation.model_class.funcparam_matrix. ParamScalarFunction.d_tt`) -- built for continuous-time-and-space PINN collocation (the ``wave_nd.py`` file this example's test case is taken from), not for the first-order-in-time, Butcher-tableau time-marching :class:`~scimba_jax.linear_approximation.collocation. time_dependent_collocation_scheme.TimeDiscreteCollocationScheme` uses here (``dU/dt + A(U) = f``, one derivative). ``WaveResidual`` also hardcodes :math:`c = 1` (no coefficient survives a plain :math:`u_{tt}`); the ``c`` appearing squared in :math:`A` below is a genuine, freely chosen parameter. Structurally this mirrors the Gray-Scott collocation example this file once was, minus what turned out to be specific to a stiff nonlinear reaction: 1. **A coupled system on one collocation space.** ``v``/``w1``/``w2`` live on one ``output_dim=3`` (``basis_type="vec"``) :class:`~scimba_jax.linear_approximation.basis.kernel_basis.KernelBasis` -- the same ``ParamVecFunction(dim=...)`` collocation code path the regression test in ``tests/test_jax/linear_approximation/ test_time_dependent_collocation.py`` exercises. 2. **The whole operator is written directly in the residual**, no bilinear/linear split (collocation is strong-form, test-function-free): :meth:`WaveSystemResidual.construct_residual` returns :math:`A(v, \mathbf{w}) = (\nabla\cdot\mathbf{w},\ c^2\nabla v)` such that :math:`\partial_t U + A(U) = 0`, built from :meth:`~scimba_jax.nonlinear_approximation.model_class.funcparam_matrix. ParamFieldFunction.divergence` and :meth:`~scimba_jax.nonlinear_approximation.model_class.funcparam_matrix. ParamScalarFunction.gradient` rather than a Laplacian -- linear, so (like the earlier version of this file, unlike Gray-Scott's ``GrayScottResidual``) there is no nonlinear term and nothing to guard against under vmap. 3. **Boundary condition: Dirichlet on ``v`` only, not the whole state.** ``wave_nd.py``'s actual boundary condition is periodic (via a periodic feature embedding in its network -- not available here); homogeneous Dirichlet on ``v`` is used instead, a good approximation only because the Gaussian pulse (:math:`\sigma \approx 0.08`, centered at the domain's middle) stays well inside :math:`(0,1)^2` over the short horizon :math:`T_\mathrm{final} = 0.2` used here -- the same "does not reach the boundary" argument the Gray-Scott file made for its own BC mismatch. Constraining only ``v`` (not ``w`` too) is also the *mathematically* appropriate choice, not just a convenient one: :math:`v = 0` on the boundary for all :math:`t`, together with :math:`u(\cdot, 0) \approx 0` there, is exactly the Dirichlet condition :math:`u = 0` on :math:`\partial\Omega` (since :math:`\partial_t u = 0` there for all :math:`t` and :math:`u` starts at :math:`\approx 0` implies :math:`u` *stays* :math:`\approx 0` there) -- whereas :math:`\mathbf{w} = -c^2\nabla u` generically has a nonzero *normal* derivative at a Dirichlet boundary (only its *tangential* component is forced to 0, automatically, by :math:`u \equiv 0` along the boundary curve): constraining the whole of ``w`` to 0 is an over-constraint an earlier version of this file made (out of convenience -- see below), which produced a measurable, spurious edge artifact even in the initial-condition projection alone, before any dynamics. :class:`~scimba_jax.linear_approximation.collocation. time_dependent_collocation_scheme._CollocationModel` only ever wires up a :class:`~scimba_jax.physical_models.boundary_residuals. DirichletResidual` on the *whole* state, so constraining ``v`` alone needed a small framework addition: :class:`TimeDiscreteCollocationScheme` now accepts a ``weights`` dict, forwarded to :class:`~scimba_jax.linear_approximation.collocation. collocation_elliptic.EllipticCollocationScheme` (which already supported it, just not threaded through here) -- see :func:`solve` below, ``weights={"interior": [1.0, 1.0, 1.0], "boundary": [1.0, 0.0, 0.0]}`` zeroes ``w``'s contribution to the boundary residual entirely (not merely down-weights it), the collocation way to impose Dirichlet on a subset of a coupled state's components. **Time integrator: Crank-Nicolson.** The system is linear and non-stiff, so nothing demands an L-stable SDIRK (Gray-Scott's Pareschi-Russo). Crank-Nicolson (2-stage, 2nd order, A-stable, non-dissipative for a linear problem -- the natural fit for a wave that should not lose amplitude to numerical damping) is used instead; per CLAUDE.md ("le solve Galerkin est deja un Newton... un schema lineaire est le cas ou une iteration suffit"), each stage costs about as much as an explicit stage's identity solve would -- an implicit tableau is essentially free here. **Reference: an independent finite-difference (leapfrog) solve, imposing the *same* boundary condition as collocation -- Dirichlet on ``v`` (via ``u``) only.** (Unlike an earlier, separable-standing-wave version of this file, a localized pulse on a bounded domain has no closed-form solution.) Central differences in time (the standard explicit leapfrog scheme for :math:`u_{tt} = c^2\Delta u`) and in space, with zero-padding, give :math:`u` on a fine grid, from which :math:`(v, \mathbf{w})` are read off by finite differences at :math:`T_\mathrm{final}`. Zero-padding imposes :math:`u = 0` on the boundary, matching collocation's Dirichlet-on-``v`` condition exactly (see above: the two are the same condition, on ``u`` versus its equivalent statement in terms of ``v``) -- and, as a direct consequence, :math:`v = \partial_t u` comes out exactly 0 at the boundary here too (:math:`u` is held at 0 at every step, so its central-difference time derivative is :math:`(0-0)/(2\,dt) = 0`), with no special-casing needed. :math:`\mathbf{w} = -c^2\nabla u` is deliberately *not* forced to 0 here either -- it is read off as an ordinary central difference of the neighbouring (generally nonzero) ``u`` values, the finite-difference analogue of collocation's own boundary points being unconstrained on ``w`` and left to the interior residual. An earlier version of this file had both sides impose the stronger :math:`v = \mathbf{w} = 0`; this file's own :func:`fd_reference` at the time even zeroed ``w1_fd``/``w2_fd`` on the boundary by hand to match -- that hand-zeroing has been removed now that collocation itself no longer constrains ``w`` there, restoring :func:`fd_reference` to a plain FD solve of ``u`` with no explicit handling of ``v``/``w`` on the boundary beyond what falls out of ``u = 0``. Resolution is fixed at ``N_POINTS_PER_DIM = 20`` (no sweep here -- see the Gray-Scott/earlier-version file's docstrings for that kind of study). **``sigma_factor`` needs to be much narrower here than in this directory's other examples, because the target is narrow.** The default ``sigma_factor =3.0`` used elsewhere gives a kernel width (``sigma_factor * spacing``, with ``spacing = 1/(n_points_per_dim - 1)``) far *wider* than this problem's Gaussian pulse (standard deviation :math:`1/\sqrt{150} \approx 0.082`) at any ``n_points_per_dim`` in the range this file is meant to run at (a handful of points to a few dozen per side): fitting a narrow target with wide, heavily-overlapping kernels forces large positive/negative coefficient cancellations to reproduce the sharp peak, and that cancellation is least constrained (fewest neighbouring collocation points) right at the domain edges -- producing Gibbs-like overshoot concentrated exactly at :math:`x=0,1` (in ``w2``) and :math:`y=0,1` (in ``w1``), where the true field is negligible. This is a pure RBF-fitting effect, present already right after :meth:`~scimba_jax.linear_approximation.collocation. time_dependent_collocation_scheme.TimeDiscreteCollocationScheme.initialize` (before any time-stepping, confirmed by checking ``dofsl_init`` directly) -- not a sign error in the residual (which was checked independently via the plane-wave/dispersion argument above) and not the RK/Newton solve. ``SIGMA_FACTOR = 1.2`` below (checked to remain finite/stable through the full first-order, differential-operator residual -- unlike a Laplacian residual, this system's Jacobian did not become ill-conditioned at a narrow ``sigma_factor``, see ``1d_heat_equation.py``'s docstring for the contrast) removes the edge artifact almost entirely and roughly halves every field's relative L2 error against the FD reference, compared to ``sigma_factor=3.0``. """ 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_2d import Square2D from scimba_jax.linear_approximation.basis.kernel_basis import KernelBasis from scimba_jax.linear_approximation.basis.kernel_function import GaussianKernel from scimba_jax.linear_approximation.collocation.time_discrete_collocation_scheme import ( TimeDiscreteCollocationScheme, ) from scimba_jax.linear_approximation.variables.collocation_variables import ( CollocationVariables, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamFieldFunction, ParamVecFunction, ) from scimba_jax.physical_models.abstract_residuals import InteriorResidual from scimba_jax.time_discrete.butcher_tableau import build_crank_nicolson_tableau DIM = 2 # ── Problem parameters (same 2D test case as wave_nd.py) ───────────────────── C = 1.0 # wave speed (WaveResidual/wave_nd.py's implicit c=1) A_LO, A_HI = 0.0, 1.0 # domain (0, 1)^2 X_MID, Y_MID = 0.5, 0.5 # Gaussian pulse center def u0_scalar(xy): """Gaussian bump, exactly wave_nd.py's f_init (its u0, not its u0/dt=0 pair).""" x, y = xy[0], xy[1] r2 = (x - X_MID) ** 2 + (y - Y_MID) ** 2 return jnp.exp(-75.0 * r2) def initial_condition(xy, mu=None): """(v0, w0) = (0, -c^2 grad(u0)), from u0's zero initial time-derivative.""" grad_u0 = jax.grad(u0_scalar)(xy) v0 = jnp.zeros((1,)) w0 = -(C**2) * grad_u0 return jnp.concatenate([v0, w0]) # ── Residual: dU/dt + A(U) = 0, first-order hyperbolic wave system ─────────── class WaveSystemResidual(InteriorResidual): r"""``A(v, w) = (div(w), c^2 grad(v))`` such that :math:`\partial_t U + A(U) = 0`. Written directly (no bilinear/linear split: see the module docstring), the same shape as :class:`~scimba_jax.physical_models.elliptic_pde. general_elliptic.GeneralEllipticResidual` -- linear in ``(v, w)``, so there is no nonlinear term and nothing to guard against under vmap. """ def __init__(self, domain, model_type: str = "x_mu"): super().__init__( domain=domain, size=3, model_type=model_type, f_rhs=lambda x, mu: jnp.zeros(3), ) def construct_residual(self, *vars: ParamVecFunction) -> ParamVecFunction: v, w1, w2 = vars[0].components() w_field = ParamFieldFunction.cat([w1, w2], main_var="x") res_v = w_field.divergence("x") res_w = (C**2) * v.gradient("x") return ParamVecFunction.cat([res_v, res_w]) # ── Kernel collocation space and solve ──────────────────────────────────────── def make_variables(n_points_per_dim: int, sigma_factor: float = 1.2): """Build a Gaussian-kernel ``out_dim=3`` collocation space on a grid over (0,1)^2. Same construction as the Gray-Scott file: the grid (interior + boundary) doubles as kernel centers and interior collocation points, and its boundary points become the boundary collocation points, holding the (approximate, see module docstring) ``v = w = 0`` Dirichlet datum. """ spacing = (A_HI - A_LO) / (n_points_per_dim - 1) sigma = sigma_factor * spacing xs = jnp.linspace(A_LO, A_HI, n_points_per_dim) xx, yy = jnp.meshgrid(xs, xs, indexing="ij") xy_centers = jnp.stack([xx.ravel(), yy.ravel()], axis=-1) on_boundary = ((xx == A_LO) | (xx == A_HI) | (yy == A_LO) | (yy == A_HI)).ravel() xy_centers_bc = xy_centers[on_boundary] domain = Square2D([(A_LO, A_HI), (A_LO, A_HI)], is_main_domain=True) basis = KernelBasis( dim=DIM, output_dim=3, kernel_function=GaussianKernel(sigma=sigma), centers=xy_centers, basis_type="vec", ) variables = CollocationVariables(basis=basis, nb_variables=3) return variables, domain, xy_centers, xy_centers_bc def solve(n_points_per_dim: int, dt: float, nt: int, sigma_factor: float = 1.2): variables, domain, xy_centers, xy_centers_bc = make_variables( n_points_per_dim, sigma_factor ) spatial_residual_factory = lambda t: WaveSystemResidual(domain=domain) # noqa: E731 scheme = TimeDiscreteCollocationScheme( spatial_residual_factory=spatial_residual_factory, variables=variables, collocation_points=xy_centers, bc_collocation_points=xy_centers_bc, main_domain=domain, butcher_tableau=build_crank_nicolson_tableau(), dt=dt, dirichlet_factory=lambda t: (lambda x, n, mu: jnp.zeros(3)), # Only v is actually constrained at the boundary (weight 1); w's two # components get weight 0 in the "boundary" block, i.e. dropped from # that residual entirely (see the module docstring) -- w is left to # whatever the interior residual (which is also evaluated at every # boundary point, see EllipticCollocationScheme._assembly_scheme_pure) # determines there, rather than additionally pinned to 0. weights={"interior": [1.0, 1.0, 1.0], "boundary": [1.0, 0.0, 0.0]}, max_iter=10, tol=1e-10, ) dofsl_init = scheme.initialize(initial_condition) dofsl_final, history = scheme.solve(dofsl_init, t0=0.0, nt=nt) return variables, dofsl_final, history # ── Reference: leapfrog finite-difference solve of u_tt = c^2 Delta(u) ─────── # Independent of the collocation solve: standard explicit central differences # in time and space, homogeneous Dirichlet (zero-padded Laplacian/gradient -- # boundary nodes are never updated and stay at 0), same BC as the collocation # solve above. N_FD = 160 _dx_fd = (A_HI - A_LO) / (N_FD - 1) _x_fd = jnp.linspace(A_LO, A_HI, N_FD) _xx_fd, _yy_fd = jnp.meshgrid(_x_fd, _x_fd, indexing="ij") _xy_fd = jnp.stack([_xx_fd.reshape(-1), _yy_fd.reshape(-1)], axis=-1) def _lap_dirichlet(u): u_pad = jnp.pad(u, 1) return ( u_pad[2:, 1:-1] + u_pad[:-2, 1:-1] + u_pad[1:-1, 2:] + u_pad[1:-1, :-2] - 4 * u ) / _dx_fd**2 def _grad_dirichlet(u): u_pad = jnp.pad(u, 1) dudx = (u_pad[2:, 1:-1] - u_pad[:-2, 1:-1]) / (2 * _dx_fd) dudy = (u_pad[1:-1, 2:] - u_pad[1:-1, :-2]) / (2 * _dx_fd) return dudx, dudy def _zero_boundary(u): return u.at[0, :].set(0.0).at[-1, :].set(0.0).at[:, 0].set(0.0).at[:, -1].set(0.0) @jax.jit def _leapfrog_step_fd(u_prev, u_curr, dt_fd): u_next = 2 * u_curr - u_prev + (C * dt_fd) ** 2 * _lap_dirichlet(u_curr) return _zero_boundary(u_next) def fd_reference(t_final: float, dt_fd: float = 1e-3): """Independent (v, w1, w2) at ``t_final``, from a leapfrog FD solve of u.""" n_steps = int(round(t_final / dt_fd)) u0 = jax.vmap(u0_scalar)(_xy_fd).reshape(N_FD, N_FD) u0 = _zero_boundary(u0) # First leapfrog step, from u_t(0) = 0 (standard ghost-point formula). u1 = _zero_boundary(u0 + 0.5 * (C * dt_fd) ** 2 * _lap_dirichlet(u0)) u_prev, u_curr = u0, u1 for _ in range(n_steps - 1): u_prev, u_curr = u_curr, _leapfrog_step_fd(u_prev, u_curr, dt_fd) # One extra step past t_final, so v = du/dt at t_final can be read off by # the same central difference the leapfrog scheme itself is built on. u_after = _leapfrog_step_fd(u_prev, u_curr, dt_fd) v_fd = (u_after - u_prev) / (2 * dt_fd) dudx, dudy = _grad_dirichlet(u_curr) w1_fd, w2_fd = -(C**2) * dudx, -(C**2) * dudy return np.array(v_fd), np.array(w1_fd), np.array(w2_fd) if __name__ == "__main__": N_POINTS_PER_DIM = 20 SIGMA_FACTOR = 1.2 DT = 1e-2 NT = 20 T_FINAL = DT * NT print( f"Solving the wave system with kernel collocation: " f"n_points_per_dim={N_POINTS_PER_DIM}, sigma_factor={SIGMA_FACTOR}, " f"dt={DT:.1e}, nt={NT} (T_final={T_FINAL:.4f}) -- Crank-Nicolson, Dirichlet" ) t0 = time.time() variables, dofsl_final, _history = solve(N_POINTS_PER_DIM, DT, NT, SIGMA_FACTOR) # block_until_ready: the solve is dispatched asynchronously, so a plain # time.time() delta here would just measure dispatch, not compute. jax.block_until_ready(dofsl_final) print(f"Collocation solve done in {time.time() - t0:.1f}s") print( "Solving independent leapfrog/FD reference " f"on a {N_FD}x{N_FD} Dirichlet grid..." ) t0 = time.time() v_ref, w1_ref, w2_ref = fd_reference(T_FINAL) print(f"FD reference done in {time.time() - t0:.1f}s") variables.dofsl = dofsl_final vw_colloc = np.asarray(jax.vmap(lambda xy: variables.evaluate(xy))(_xy_fd)) v_colloc = vw_colloc[:, 0].reshape(N_FD, N_FD) w1_colloc = vw_colloc[:, 1].reshape(N_FD, N_FD) w2_colloc = vw_colloc[:, 2].reshape(N_FD, N_FD) fields = [ ("v", v_colloc, v_ref), ("w_1", w1_colloc, w1_ref), ("w_2", w2_colloc, w2_ref), ] errs = {} for name, colloc, ref in fields: err = np.abs(colloc - ref) rel_l2 = np.linalg.norm(err) / np.linalg.norm(ref) errs[name] = (err, rel_l2) print( f"{name}: rel. L2 error vs. FD reference = {rel_l2:.3e}, " f"max abs = {err.max():.3e}" ) # ── Plot: collocation vs. FD reference vs. |difference|, one row per field ─ X, Y = np.array(_x_fd), np.array(_x_fd) fig, axes = plt.subplots(3, 3, figsize=(15, 13.5)) for row, (name, colloc, ref) in enumerate(fields): err, rel_l2 = errs[name] panels = [ (colloc, f"${name}$ collocation"), (ref, f"${name}$ FD reference"), (err, rf"$|{name}^\mathrm{{colloc}} - {name}^\mathrm{{FD}}|$"), ] for col, (grid, title) in enumerate(panels): ax = axes[row, col] # grid[i, j] holds the value at (xs[i], xs[j]) (axis 0 = x-index, # from the "ij"-indexed meshgrid above), but pcolormesh(X, Y, C) # expects C's first axis to index Y and second to index X -- so # without the transpose, every panel renders x and y swapped # (e.g. a feature that is really a function of x alone would show # up varying along y instead). im = ax.pcolormesh(X, Y, grid.T, cmap="turbo", shading="auto") fig.colorbar(im, ax=ax, pad=0.02) ax.set_title(title) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") fig.suptitle( rf"Wave system 2D, kernel collocation (Crank-Nicolson, " rf"{N_POINTS_PER_DIM}x{N_POINTS_PER_DIM} points, Dirichlet) " rf"vs. FD reference at $t={T_FINAL:.4f}$" ) fig.tight_layout() plt.show()