"""Heterogeneous Helmholtz, pushed harder: k_bg=100 background, k=150 inside a small central circle -- the same physical setup as the heterogeneous dataset of ``benchmarks/benchmarks_jax/helmholtz_shifted_laplacian`` (k_bg=20, k_inclusion=40), scaled up to where a direct solve is no longer an option at all, only the CSL+multigrid-preconditioned solve is run here. In the library's terms (what this script is now written with): the weak form ``HelmholtzWeakForm(dim, kappa, f, kappa_im=shift)`` -- ``kappa`` and the shift floats or functions of ``x`` --, the absorbing condition ``Sommerfeld(k)``, and ``shifted_laplacian_solver(scheme_factory, n_cells, shift, ...)`` which builds the two-level hierarchy, the Jacobi-smoothed V-cycle and the preconditioned BiCGSTAB from ONE factory of the problem. Measured at k = 100 on 160 x 160 (see ``solvers/helmholtz.py``): the whole cost is the dense LU of the coarse level -- 9 s to factorise, 100 ms per V-cycle, i.e. a 6 s solve of 35 iterations; a sparse coarse factorisation is the one lever left, not done. Delta u + k(x)^2 u = f in (0,1)^2, u = (u_Re, u_Im) d_n u - i*k_bg*u = 0 on every side Same shift, applied pointwise, no new tuning: ``kappa_re(x) = k(x)^2``, ``kappa_im(x) = 2*k(x)`` (see the docstring of ``benchmarks/benchmarks_jax/helmholtz_shifted_laplacian``'s ``benchmark_utils/helmholtz.py`` for why this additive, O(k) shift is the one that converges in this framework, and why it generalizes pointwise for a heterogeneous medium with no separate re-derivation). **Mesh: pollution-free at the HARDER local wavenumber.** ``N_FINEST`` comes from the same ``h ~ k^-1.5`` scaling used throughout this family, calibrated to k=40/h=1/40 and evaluated at the highest local wavenumber, k=150: ``h(150) = (1/40)*40^1.5*150^-1.5 ~ 1/290``. That gives ``n_dof=169362`` -- 6x the smaller heterogeneous script's problem size. **Why only the preconditioned solve.** A dense direct solve at this size would need ``n_dof^2 * 8 bytes ~ 230 GB`` for the Jacobian alone -- not an option on a normal machine (this one has ~17 GB RAM), which is exactly the point: this is the regime CSL+multigrid exists for. Correctness at smaller, directly-verifiable scales (both the homogeneous scripts and the k_bg=20/ k_inclusion=40 heterogeneous one) already established that this shift and this multigrid recipe solve the right problem; nothing about moving to a larger, still-heterogeneous mesh changes that construction, only its cost. **Result, measured here.** CSL+multigrid converges to residual 6.4e-9 in 16 BiCGSTAB iterations -- between the homogeneous k=100 case's 37 (the homogeneous dataset of the benchmark) and the smaller heterogeneous case's 10, consistent with most of this domain sitting at k_bg=100 (harder than k_bg=20, easier than a uniform k=150 would be) with only a small k=150 region. The solution field visibly shows the shorter wavelength inside the inclusion (more fringes per unit length there than in the k_bg=100 background) -- the expected refraction signature of a higher-wavenumber region, not an artifact. Needs no optional extra (structured Cartesian mesh, no GMSH). Saves ``solve_helmholtz_shifted_laplacian_precond_heterogeneous_hard.png`` next to this file. """ # %% import time from pathlib import Path import matplotlib matplotlib.use("Agg") # headless: never block on a display 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 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.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.solvers import ( shifted_laplacian_solver, ) from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE 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.helmholtz_weak_form import ( HelmholtzWeakForm, ) from scimba_jax.physical_models.weak_boundary_conditions import Sommerfeld _HERE = Path(__file__).resolve().parent DIM = 2 POLY_ORDER = 1 QUAD_ORDER = 4 K_BG = 100.0 K_INCLUSION = 150.0 CENTER = jnp.array([0.5, 0.5]) RADIUS = 0.12 EPS_COEFF = 2.0 # eps(x) = EPS_COEFF * k(x), the same ratio verified throughout # Pollution-free at the HARDER local wavenumber (k=150), calibrated to the # verified k=40/h=1/40 reference: h ~ (1/40)*40**1.5*k**-1.5. N_FINEST = 290 N_LEVELS = 2 # 290 -> 145: kh_coarse = 150/145 ~ 1.03, still resolved X0 = jnp.array([[0.5, 1.0 / 32.0]]) # source near the bottom edge, away from # the inclusion: the wave must pass through it to reach the far side SIGMA_SOURCE = 0.02 BICGSTAB_TOL = 1e-8 MAX_ITER_LINEAR = 100 JACOBI_OMEGA = 0.5 def k_of_x(x): """Background ``K_BG``, ``K_INCLUSION`` inside a small centred circle.""" r2 = jnp.sum((x - CENTER) ** 2) return jnp.where(r2 < RADIUS**2, K_INCLUSION, K_BG) def f_source(x): """A real point source, ``[f_Re, f_Im] = [delta_sigma(x - x0), 0]``.""" bump = jnp.sum(jnp.exp(-jnp.sum((x - X0) ** 2, axis=1) / (2.0 * SIGMA_SOURCE**2))) bump = bump / (2.0 * jnp.pi * SIGMA_SOURCE**2) return jnp.array([bump, 0.0]) # %% The problem in the library's terms. ONE factory: the real problem at ``shift = 0``, # the preconditioner's operator at ``shift = eps`` -- ``HelmholtzWeakForm`` with a # ``kappa_im``, the absorbing condition ``Sommerfeld(k)``, and the shifted-Laplacian # solver built from that factory. Module-level functions, not lambdas: they are part of # the scheme's compile key, and a fresh lambda per call recompiled the solve for every # scheme built (measured: 1.3 s per scheme). def identity(x): return x def local_q1(c, i, m): return local_lagrange_basis(c, i, m, order=POLY_ORDER, out_dim=2) def make_variables(n_cells) -> VariablesFE: n = n_cells[0] if isinstance(n_cells, tuple) else n_cells mesh = Mesh( dim=DIM, n_cells=(n, n), ref_quad=UnitSquareTensorized(dim=DIM, order=QUAD_ORDER), mapping=Mapping(mappings=[InvertibleFunction(identity, identity)]), ) basis = AnalyticBasis( nb_basis=(POLY_ORDER + 1) ** DIM, out_dim=2, mesh=mesh, basis_type="vec", local_basis=local_q1, ) return VariablesFE(basis=basis, nb_variables=2) def kappa_of_x(x): """``k(x)^2``.""" return k_of_x(x) ** 2 def eps_of_x(x): """The shift, pointwise at the local wavenumber: ``EPS_COEFF * k(x)``.""" return EPS_COEFF * k_of_x(x) def scheme_factory(n_cells, shift=0.0) -> EllipticFEscheme: """``Delta u + (k(x)^2 + i shift(x)) u = f_source``, absorbing at ``k_bg`` on every side.""" model = AbstractPhysicalWeakModel(dim=DIM) model.add_weak_form( "main", HelmholtzWeakForm(dim=DIM, kappa=kappa_of_x, f=f_source, kappa_im=shift) ) for side in ("west", "east", "south", "north"): model.add_boundary_condition(side, Sommerfeld(K_BG)) return EllipticFEscheme(model, make_variables(n_cells)) n_dof = 2 * (N_FINEST + 1) ** 2 print( f"\nk_bg = {K_BG}, k_inclusion = {K_INCLUSION}, radius = {RADIUS}, " f"eps_coeff = {EPS_COEFF}, N_FINEST = {N_FINEST}, n_dof = {n_dof}\n" ) precond_start = time.perf_counter() csl = shifted_laplacian_solver( scheme_factory, (N_FINEST, N_FINEST), shift=eps_of_x, tol=BICGSTAB_TOL, max_iter=MAX_ITER_LINEAR, n_levels=N_LEVELS, omega=JACOBI_OMEGA, ) precond_setup_time = time.perf_counter() - precond_start print(f"CSL+MG setup: {precond_setup_time:.2f} s ({N_LEVELS}-level hierarchy)") start = time.perf_counter() scheme = scheme_factory((N_FINEST, N_FINEST)) solved, report = EllipticFEscheme.solve( scheme, solver=csl, return_report=True, ) jax.block_until_ready(solved.variables.dofsl) elapsed = time.perf_counter() - start print( f"CSL+MG preconditioner: {int(report.n_linear)} BiCGSTAB iters, " f"residual {float(report.residual):.3e}, converged {bool(report.converged)}, " f"{elapsed:.2f} s (incl. JIT)" ) # %% Plot: the wavenumber field and the solution. def sample_complex_solution(scheme, n_side=10): mesh = scheme.variables.mesh grid = jnp.stack( jnp.meshgrid( jnp.linspace(0.0, 1.0, n_side), jnp.linspace(0.0, 1.0, n_side), indexing="ij", ), axis=-1, ).reshape(-1, 2) dofs = scheme.variables.dofsl connectivity = scheme.variables.connectivity def one_cell(cell): points = mesh._unit_hypercube_to_cell(cell, grid) theta = dofs[connectivity[cell]] values = jax.vmap( lambda p: jnp.einsum( "iv,iv->v", theta, scheme.variables.trial_basis(cell, p[None, :])[0] ) )(points) return points, values points, values = jax.vmap(one_cell)(jnp.arange(mesh.n_cells_total)) corner = np.arange(n_side - 1) i, j = np.meshgrid(corner, corner, indexing="ij") bottom_left = (i * n_side + j).ravel() local = np.concatenate( [ np.stack([bottom_left, bottom_left + n_side, bottom_left + 1], axis=1), np.stack( [bottom_left + n_side, bottom_left + n_side + 1, bottom_left + 1], axis=1, ), ] ) offsets = (np.arange(mesh.n_cells_total) * n_side**2)[:, None, None] triangles = (local[None, :, :] + offsets).reshape(-1, 3) return ( np.asarray(points).reshape(-1, 2), triangles, np.asarray(values).reshape(-1, 2), ) figure = plt.figure(figsize=(15, 5)) ax_k = figure.add_subplot(1, 4, 1) grid_1d = np.linspace(0.0, 1.0, 200) xx, yy = np.meshgrid(grid_1d, grid_1d, indexing="ij") kk = np.asarray(jax.vmap(jax.vmap(k_of_x))(jnp.stack([xx, yy], axis=-1))) drawing_k = ax_k.pcolormesh(xx, yy, kk, cmap="turbo", shading="auto") figure.colorbar(drawing_k, ax=ax_k, fraction=0.046) ax_k.set_title("k(x): wavenumber field") ax_k.set_aspect("equal") ax_k.set_xlabel("x") ax_k.set_ylabel("y") points, triangles, values = sample_complex_solution(solved, n_side=6) u_re, u_im, amplitude = values[:, 0], values[:, 1], np.hypot(values[:, 0], values[:, 1]) for col, (field, title) in enumerate( zip((u_re, u_im, amplitude), ("Re(u)", "Im(u)", "|u|")) ): ax = figure.add_subplot(1, 4, 2 + col) drawing = ax.tricontourf( points[:, 0], points[:, 1], triangles, field, levels=60, cmap="turbo" ) figure.colorbar(drawing, ax=ax, fraction=0.046) circle = plt.Circle( (float(CENTER[0]), float(CENTER[1])), RADIUS, fill=False, color="white", lw=1.0 ) ax.add_patch(circle) ax.set_title(f"{title} (CSL+MG-preconditioned)") ax.set_aspect("equal") ax.set_xlabel("x") ax.set_ylabel("y") figure.suptitle( f"Heterogeneous Helmholtz, k_bg={K_BG}/k_inclusion={K_INCLUSION}, unit square " f"({N_FINEST}x{N_FINEST} cells), CSL+multigrid only: " f"{int(report.n_linear)} BiCGSTAB iterations to {float(report.residual):.1e}" ) figure.tight_layout() figure.savefig( _HERE / "solve_helmholtz_shifted_laplacian_precond_heterogeneous_hard.png", dpi=130 ) print( f"\nSaved " f"{_HERE / 'solve_helmholtz_shifted_laplacian_precond_heterogeneous_hard.png'}" ) plt.close("all") # the backend is Agg: the figure is the PNG, `show` would only warn