"""Solve the Laplacian DG scheme via Newton-Raphson.""" # import os # os.environ["JAX_PLATFORM_NAME"] = "cpu" 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 ( # , plot_convergence h1_seminorm_error, 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 out_dim = 1 mapping_id = InvertibleFunction(lambda x: x, lambda y: y) pde = LaplacianWeakForm( dim=physical_dim, f=lambda x: jnp.sin(jnp.pi * x[0]), ) def u_exact_fn(x): return jnp.sin(jnp.pi * x[0]) / jnp.pi**2 # (200,) # ── Convergence curves ──────────────────────────────────────────────────────── n_cells_list = [20] flux_configs = [ ("SIPG", lambda h, p: SIPGFlux(sigma=p * (p + 1), h=h)), ] fig_conv, axes = plt.subplots( len(flux_configs), 3, figsize=(15, 4 * len(flux_configs)), squeeze=False ) for row, (flux_name, flux_factory) in enumerate(flux_configs): ax_l2, ax_h1, ax_linf = axes[row] for deg in [1]: # , 2]: basis_order = deg def make_solver(b=basis_order, ff=flux_factory, p=deg): def solver(n_cells): # t_total = time.perf_counter() # t_mesh = time.perf_counter() 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=b + 1, out_dim=out_dim, mesh=m, local_basis=lambda coords, i, mesh, _b=b: local_taylor_basis( coords, i, mesh, order=_b, out_dim=out_dim ), basis_type="scalar", ) variables = VariablesDG(basis=taylor_basis, nb_variables=out_dim) def dirichlet_bc(__x): return jnp.zeros(out_dim) assembler = EllipticDGscheme( AbstractPhysicalWeakModel.from_weak_form( pde, dirichlet=dirichlet_bc ), variables, ff(1 / n_cells, p), ) # print(f" [n={n_cells}] mesh+basis setup: {time.perf_counter() - t_mesh:.3f}s") return EllipticDGscheme.solve(assembler) return solver print(f"\nConvergence study {flux_name} p={deg}:") for n in n_cells_list: assembler = make_solver()(n) # print(f" [n={n}] solver (total from plot_convergence): {time.perf_counter() - t_solver:.3f}s") # t_err = time.perf_counter() e2 = l2_error(assembler, u_exact_fn) # e2.block_until_ready() # print(f" [n={n}] l2_error: {time.perf_counter() - t_err:.3f}s") # t_err = time.perf_counter() eh1 = h1_seminorm_error(assembler, u_exact_fn) # eh1.block_until_ready() # print(f" [n={n}] h1_error: {time.perf_counter() - t_err:.3f}s") # t_err = time.perf_counter() einf = linf_error(assembler, u_exact_fn) # einf.block_until_ready() # print(f" [n={n}] linf_error: {time.perf_counter() - t_err:.3f}s") print(f" [n={n}] errors: L2={e2:.3e} H1={eh1:.3e} Linf={einf:.3e}")