"""Solve the Laplacian DG scheme via Newton-Raphson.""" import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.linear_approximation.basis.analytic_bases import local_taylor_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.error_analysis import l2_error, linf_error from scimba_jax.linear_approximation.galerkin.dg.elliptic_dg_scheme import ( EllipticDGscheme, ) from scimba_jax.linear_approximation.galerkin.dg.flux import SIPGFlux from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_dg import VariablesDG from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) physical_dim = 1 quad_order = 3 polynomial_degree = 2 basis_order = polynomial_degree nb_basis = basis_order + 1 out_dim = 1 n_cells = 20 # smaller for jacrev cost mapping_id = InvertibleFunction(lambda x: x, lambda y: y) m = Mesh( dim=physical_dim, n_cells=[n_cells], ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), ) taylor_basis = AnalyticBasis( nb_basis=nb_basis, out_dim=out_dim, mesh=m, local_basis=lambda coords, i, mesh: local_taylor_basis( coords, i, mesh, order=basis_order, out_dim=out_dim ), basis_type="scalar", ) variables = VariablesDG(basis=taylor_basis, nb_variables=out_dim) variables.dofsl = jnp.ones_like(variables.dofsl) pde = LaplacianWeakForm( dim=physical_dim, f=lambda x: jnp.ones_like(x), ) def u_exact_fn(x): return jnp.array([-0.5 * x[0] ** 2 + 0.5 * x[0]]) # + 1.0 # (200,) # sigma = 2.0 * n_cells # * 10 sigma = polynomial_degree * (polynomial_degree + 1) def dirichlet_bc(x): return jnp.zeros(out_dim) assembler = EllipticDGscheme( AbstractPhysicalWeakModel.from_weak_form(pde, dirichlet=dirichlet_bc), variables, SIPGFlux(sigma=sigma, h=1 / n_cells), ) # Function: dofsl -> residual vector def residual(dofsl): return EllipticDGscheme._assembly_scheme_pure(assembler, dofsl) ## first check : compute residual and jacobian shapes and norms # Compute residual at current dofsl res = jax.jit(residual)(variables.dofsl) print("Residual shape:", res.shape) assert res.shape == (m.n_cells[0] * nb_basis * out_dim,), ( f"Unexpected shape: {res.shape}" ) # Compute Jacobian of residual w.r.t. dofsl J = jax.jit(jax.jacrev(residual))(variables.dofsl) print("Jacobian shape:", J.shape) assert J.shape == ( m.n_cells[0] * nb_basis * out_dim, m.n_cells[0], nb_basis, out_dim, ), f"Unexpected shape: {J.shape}" assembler = EllipticDGscheme.solve(assembler, max_iter=3) dofsl_sol = assembler.variables.dofsl print("Solution dofsl shape:", dofsl_sol.shape) # ── Plot DG solution vs exact ──────────────────────────────────────────────── x_plot = jnp.linspace(0.0, 1.0, 200)[:, jnp.newaxis] # (200, 1) u_dg = variables.evaluate(x_plot)[:, 0] # (200,) l2_error = l2_error(assembler, u_exact_fn) print(f"L2 error ||u - u_h||_L2 = {l2_error:.6e}") linf_error = linf_error(assembler, u_exact_fn) print(f"Linf error ||u - u_h||_Linf = {linf_error:.6e}") u_exact = jax.vmap(u_exact_fn)(x_plot)[:, 0] plt.figure() plt.plot(x_plot[:, 0], u_exact, label="Exact") plt.plot(x_plot[:, 0], u_dg, "--", label="DG") plt.xlabel("x") plt.ylabel("u") plt.legend() plt.tight_layout() plt.show()