r"""Transport of a Gaussian pulse in 2D with a PINN, using a relu4 output activation. .. math:: \partial_t \rho + v \cdot \nabla \rho & = 0 \quad \text{in } \Omega \times (0, T) \\ \rho & = g \quad \text{on } \partial \Omega \times (0, T) \\ \rho & = \rho_0 \quad \text{on } \Omega \times \{0\} The exact solution is a Gaussian pulse translated at constant velocity `v`, which stays non-negative everywhere: the last-layer activation of the MLP is set to ``"relu4"`` (:math:`\mathrm{ReLU}(x)^4`) to build that property into the network itself. Boundary and initial conditions are enforced weakly. """ import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_time import ( UniformTimeSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.temporal_pde.advection import AdvectionND from scimba_jax.plots.plots_nd import plot_abstract_approx_space X0, Y0 = 0.3, 0.3 VX, VY = 0.9, 0.6 VELOCITY = jnp.array([VX, VY]) SIGMA = 0.08 T_MIN, T_MAX = 0.0, 0.5 N_COLLOC = 2500 N_BC_COLLOC = 2000 N_IC_COLLOC = 2000 N_EPOCHS = 100 HIDDEN = [16] * 3 key = jax.random.PRNGKey(0) def exact_sol_pointwise(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: x1, x2 = x[0:1], x[1:2] r2 = (x1 - X0 - VX * t) ** 2 + (x2 - Y0 - VY * t) ** 2 return jnp.exp(-r2 / (2 * SIGMA**2)) def exact_sol(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Batched exact solution, for plotting.""" x1, x2 = x[:, 0:1], x[:, 1:2] r2 = (x1 - X0 - VX * t) ** 2 + (x2 - Y0 - VY * t) ** 2 return jnp.exp(-r2 / (2 * SIGMA**2)) def f_init(x: jnp.ndarray) -> jnp.ndarray: return exact_sol_pointwise(jnp.zeros(1), x) def dirichlet_bc(t: jnp.ndarray, x: jnp.ndarray, n: jnp.ndarray) -> jnp.ndarray: return exact_sol_pointwise(t, x) domain_x = Square2D([(0.0, 1.0), (0.0, 1.0)], is_main_domain=True) domain_t = (T_MIN, T_MAX) sampler = TensorizedSampler( [UniformTimeSampler(domain_t), DomainSampler(domain_x)], model_type="t_x", bc=True, ic=True, ) model = AdvectionND( main_domain=domain_x, time_domain=domain_t, velocity=VELOCITY, bc="weak", f_bc_rhs=dirichlet_bc, ic="weak", f_ic_rhs=f_init, ) key, subkey = jax.random.split(key) nn = MLP( in_size=3, out_size=1, hidden_sizes=HIDDEN, activation="tanh", activation_output="relu4", key=subkey, ) space = ApproximationSpace({"x": 2, "t": 1}, [(nn, "scalar", None)], model_type="t_x") pinn = Projector(model, space, sampler) t0 = timeit.default_timer() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC, N_BC_COLLOC, N_IC_COLLOC) print("best loss: ", pinn.best_loss) print("time for %d epochs: " % N_EPOCHS, timeit.default_timer() - t0) plot_abstract_approx_space( pinn.space, domain_x, time_domain=domain_t, time_values=(T_MIN, (T_MIN + T_MAX) / 2, T_MAX), loss=pinn.losses, residual=pinn.model, error=exact_sol, title="transport of a 2D Gaussian pulse (relu4 output activation)", ) plt.show()