"""Time-dependent CG-FEM: rotational anisotropic diffusion on the disk. dT/dt = div(B grad T) on the unit disk, T = 0.1 on the boundary B = b b^T, b = (-y/r, x/r), r = sqrt(x^2 + y^2) T(x, y, 0) = 0.1 + 10 * exp(-((x - 0.6)^2 + y^2) / 0.02) Case 5.3 of https://inria.hal.science/hal-01020955v1/document. ``b`` is the unit vector field circling the origin, so ``B`` has eigenvalues 1 (along the circles) and 0 (radial): heat moves only *along* field lines. It is a classic stress test for a mesh not aligned with the field -- any radial spreading of the initial blob is discretisation error, not physics. ``b`` is singular at the origin; ``R_MIN`` floors ``r`` there, which changes nothing elsewhere. Isoparametric Q2 on two disks of comparable size -- an unstructured Gmsh mesh (``h = 0.15``, 207 cells) and a block-structured O-grid (``n = 6``, 180 cells) -- and two integrators on each, to the same ``T_FINAL``: * **Pareschi-Russo** (SDIRK, ``gamma = 1 - sqrt(2)/2``), ``dt = 5e-3``. The problem is linear, so each implicit stage is ONE linear solve of ``M + gamma dt K``: a ``LinearKrylov`` (no outer Newton loop), preconditioned by the exact Jacobi diagonal of that operator. SDIRK at a fixed ``dt`` means every stage of every step solves the same operator, so the preconditioner is built once, from :func:`stage_scheme`, and is exact for the whole run. * **RK4**, the reference: explicit, on the same consistent mass (each stage a Jacobi-preconditioned CG on ``M``), at ``dt = 0.8 * 2.7853 / lambda_max(M^-1 K)``, with ``2.7853`` RK4's real-axis stability bound and ``lambda_max`` from 30 matrix-free power iterations (an inner CG solve on ``M`` at each). The smaller of the two meshes' ``dt`` is shared. ⚠ The explicit ``dt`` must come from the operator actually solved. A hand-picked ``dt = 5e-4`` was unstable: with the consistent mass, ``T`` reached about ``1e58`` while every stage solve converged (residual ~1e-11). Measured on a finer Gmsh disk (``h = 0.1``, 424 cells): estimate ``dt ~ 1.2e-4``, a separate bisection stable at ``1.35e-4`` and divergent at ``1.4e-4``. ⚠ Why the consistent mass for the reference. ``mass_lumping=True`` makes an explicit stage one division per DOF, but it is a different spatial discretisation, with its own (smaller) ``lambda_max``: against a lumped RK4 reference, the gap of the implicit schemes stopped shrinking under refinement, on both meshes -- that plateau was the lumped/consistent mass difference. The figure has one row per (mesh, scheme), at ``t = 0``, ``0.3`` and ``1.5``; the L2 gap between the two schemes is printed per mesh. The preconditioner x time-scheme comparison on finer meshes is the benchopt benchmark ``benchmarks/benchmarks_jax/anisotropic_diffusion`` (dataset ``disk_2d``). """ import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.linear_approximation.basis.analytic_bases import ( local_lagrange_basis, local_lagrange_basis_by_logical, ) from scimba_jax.linear_approximation.basis.dof_map import UnstructuredLagrangeDofMap from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme import ( EllipticFEscheme, ) from scimba_jax.linear_approximation.galerkin.fem.time_discrete_fe_scheme import ( TimeDiscreteFEscheme, stage_scheme, ) from scimba_jax.linear_approximation.meshes.unstructured_mesh import UnstructuredMesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.solvers.krylov import cg_lax from scimba_jax.linear_approximation.solvers.newton import LinearKrylov from scimba_jax.linear_approximation.solvers.preconditioners import coloured_jacobi from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.macro_mesh import macro_mesh_disk, macro_mesh_ogrid_disk from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.diffusion_advection_reaction_weak_form import ( # noqa: E501 EllipticWeakForm, ) from scimba_jax.plots.plots_galerkin import sample_solution from scimba_jax.time_discrete.butcher_tableau import ( build_pareschi_russo_tableau, build_rk4_tableau, ) ORDER = 2 DISK_MESH_SIZE = 0.15 # unstructured Gmsh disk: 207 cells O_GRID_N = 6 # O-grid disk, 5 n^2 = 180 cells R_MIN = 1e-3 T_BND = 0.1 T_FINAL = 1.5 T_MID = 0.3 # the middle snapshot DT_IMP = 5e-3 # The stage solves: tol well below the O(0.1)-O(10) scale of T, far from the # 1e-14 regime where unpreconditioned CG was measured to lose A-orthogonality. LINEAR_SOLVE_TOL = 1e-10 LINEAR_SOLVE_MAX_ITER = 300 # RK4's real-axis stability bound: the root of |1+z+z^2/2+z^3/6+z^4/24| = 1. RK4_STABILITY_BOUND = 2.7853 DT_SAFETY_FACTOR = 0.8 POWER_ITERATIONS = 30 MASS_SOLVE_TOL = 1e-8 # the inner CG on M of the power iteration MASS_SOLVE_MAX_ITER = 500 # ── The problem: module-level callables (they sit in pytree aux_data) ──────── def diffusion_tensor(x): r = jnp.maximum(jnp.sqrt(x[0] ** 2 + x[1] ** 2), R_MIN) b = jnp.array([-x[1], x[0]]) / r return jnp.outer(b, b) def zero_vector(x): return jnp.zeros(2) def zero_scalar(x): return jnp.zeros(()) def boundary_value(x): return jnp.full((1,), T_BND) def u0(x): return T_BND + 10.0 * jnp.exp(-((x[0] - 0.6) ** 2 + x[1] ** 2) / 0.02) def lagrange(y, i, m): return local_lagrange_basis(y, i, m, order=ORDER, out_dim=1) def lagrange_by_logical(y, i, m): return local_lagrange_basis_by_logical(y, i, m, order=ORDER, out_dim=1) def spatial_weak_form(): return EllipticWeakForm( dim=2, A=diffusion_tensor, b=zero_vector, c=zero_scalar, f=zero_scalar ) def build_variables(macro): mesh = UnstructuredMesh( nodes=macro.nodes, cells=macro.cells, ref_quad=UnitSquareTensorized(dim=2, order=2 * ORDER + 2), order=ORDER, ) basis = AnalyticBasis( nb_basis=(ORDER + 1) ** 2, out_dim=1, mesh=mesh, local_basis_by_logical=lagrange_by_logical, local_basis=lagrange, basis_type="scalar", ) return VariablesFE(basis=basis, nb_variables=1, dof_map=UnstructuredLagrangeDofMap) # ── Stage solver and explicit time step ────────────────────────────────────── def stage_solver(variables, a_ii, dt): """Jacobi-preconditioned ``LinearKrylov`` for ``M + a_ii dt K`` (``a_ii = 0``: M).""" probe = stage_scheme( spatial_weak_form(), variables, a_ii, dt, dirichlet=boundary_value ) return LinearKrylov( preconditioner=coloured_jacobi(probe), max_iter_linear=LINEAR_SOLVE_MAX_ITER, tol=LINEAR_SOLVE_TOL, ) @jax.jit def _power_iterate(k_scheme, m_scheme, dofsl_zero, v0): # The schemes are ARGUMENTS: closed over, they would be baked in as constants. jvp = EllipticFEscheme._make_jvp_base_fn() def step(v, _): w, _, _ = cg_lax( lambda z: jvp(m_scheme, dofsl_zero, z), jvp(k_scheme, dofsl_zero, v), tol=MASS_SOLVE_TOL, max_iter=MASS_SOLVE_MAX_ITER, ) return w / jnp.linalg.norm(w), None return jax.lax.scan(step, v0, None, length=POWER_ITERATIONS)[0] def estimate_stable_dt(variables): """``0.8 * 2.7853 / lambda_max(M^-1 K)``, M the consistent mass: ``(dt, lambda_max)``. ``K`` and ``M`` carry the same Dirichlet rows (identity at the boundary): that spurious eigenvalue 1 is far below the interior ``lambda_max``. """ k_scheme = EllipticFEscheme( AbstractPhysicalWeakModel.from_weak_form( spatial_weak_form(), dirichlet=boundary_value ), variables, ) m_scheme = stage_scheme( spatial_weak_form(), variables, 0.0, 1.0, dirichlet=boundary_value ) dofsl_zero = jnp.zeros_like(variables.dofsl) v0 = jnp.asarray(np.random.default_rng(0).normal(size=dofsl_zero.size)) v = _power_iterate(k_scheme, m_scheme, dofsl_zero, v0 / jnp.linalg.norm(v0)) jvp = EllipticFEscheme._make_jvp_base_fn() lambda_max = float( jnp.dot(v, jvp(k_scheme, dofsl_zero, v)) / jnp.dot(v, jvp(m_scheme, dofsl_zero, v)) ) return DT_SAFETY_FACTOR * RK4_STABILITY_BOUND / lambda_max, lambda_max # ── One run ────────────────────────────────────────────────────────────────── def solve(variables, label, tableau, dt): """Integrate to ``T_FINAL``; ``(scheme, dofsl at 0, T_MID, T_FINAL)``. Two legs without history: at a CFL-bound ``dt`` the stacked trajectory, not the solve, is what runs out of memory. """ nt = round(T_FINAL / dt) dt = T_FINAL / nt # land exactly on T_FINAL a_ii = float(tableau.A_imp[-1, -1]) # SDIRK: the same at every stage scheme = TimeDiscreteFEscheme( spatial_weak_form_factory=spatial_weak_form(), variables=variables, butcher_tableau=tableau, dt=dt, dirichlet=boundary_value, solver=stage_solver(variables, a_ii, dt), ) dofsl_init = scheme.initialize(u0) n_mid = round(T_MID / dt) dofsl_mid, _ = scheme.solve(dofsl_init, t0=0.0, nt=n_mid, keep_history=False) dofsl_final, _ = scheme.solve( dofsl_mid, t0=n_mid * dt, nt=nt - n_mid, keep_history=False ) jax.block_until_ready(dofsl_final) print( f" {label}: {nt} steps of dt = {dt:.3e}, " f"min/max T(t = {T_FINAL}) = {float(dofsl_final.min()):.4f} / " f"{float(dofsl_final.max()):.4f}" ) return scheme, (dofsl_init, dofsl_mid, dofsl_final) def l2_gap(variables, first, second): """``||u_first - u_second||_{L2}`` on one space.""" variables.dofsl = first - second weights, points = variables.mesh.evaluate_mesh_weights_points() values = variables.evaluate_quad(points) return float(jnp.sqrt(jnp.sum(weights[:, :, None] * values**2))) if __name__ == "__main__": meshes = { "unstructured disk": macro_mesh_disk( radius=1.0, order=ORDER, mesh_size=DISK_MESH_SIZE ), "O-grid disk": macro_mesh_ogrid_disk(radius=1.0, n=O_GRID_N, order=ORDER), } spaces = {label: build_variables(macro) for label, macro in meshes.items()} estimates = {label: estimate_stable_dt(v) for label, v in spaces.items()} for label, (dt_stable, lambda_max) in estimates.items(): print(f"{label}: lambda_max(M^-1 K) = {lambda_max:.4e} -> dt = {dt_stable:.4e}") dt_exp = min(dt_stable for dt_stable, _ in estimates.values()) schemes = {"Pareschi-Russo": (build_pareschi_russo_tableau(), DT_IMP)} schemes["RK4"] = (build_rk4_tableau(), dt_exp) results = {} for mesh_label, variables in spaces.items(): print(f"{mesh_label}, {variables.mesh.n_cells_total} Q{ORDER} cells") for label, (tableau, dt) in schemes.items(): results[mesh_label, label] = solve(variables, label, tableau, dt) gap = l2_gap( variables, results[mesh_label, "Pareschi-Russo"][1][-1], results[mesh_label, "RK4"][1][-1], ) print(f" ||T_PR - T_RK4||_L2 at t = {T_FINAL}: {gap:.3e}") figure, axes = plt.subplots(len(results), 3, figsize=(16, 5 * len(results))) for row, ((mesh_label, label), (scheme, snapshots)) in enumerate(results.items()): for col, (dofsl, t) in enumerate(zip(snapshots, (0.0, T_MID, T_FINAL))): ax = axes[row, col] scheme.variables.dofsl = dofsl points, triangles, values = sample_solution(scheme, n_side=3) drawing = ax.tricontourf( points[:, 0], points[:, 1], triangles, values, levels=128, cmap="turbo" ) ax.triplot( points[:, 0], points[:, 1], triangles, lw=0.25, color="k", alpha=0.5 ) figure.colorbar(drawing, ax=ax, fraction=0.046) ax.set_title(f"{mesh_label}, {label}, t={t:.2f}") ax.set_aspect("equal") ax.set_xlabel("x") ax.set_ylabel("y") figure.tight_layout() plt.show()