"""Complex Helmholtz on an unstructured disk, ``N_RHS`` random Gaussian-mixture sources solved together with a single ``jax.vmap`` call. Delta u + kappa u = f in the unit disk, u = (u_Re, u_Im) d_n u - i*k*u = 0 on the whole boundary Same physics, mesh, basis and boundary condition as ``solve_2d_helmholtz_robin_unstructured_disk.py`` -- this file only changes the source term and how many problems get solved. ``f`` is no longer a single narrow bump at two fixed points, but ``N_RHS`` independent draws of a mixture of ``N_GAUSSIANS`` real Gaussians (random centers, amplitudes, widths), and every draw is solved in ONE compiled program via ``jax.vmap`` instead of a Python loop over ``N_RHS`` separate calls. **Why not ``jax.vmap(EllipticFEscheme.solve)``.** ``Galerkin.solve`` ends with ``scheme_pytree._store_dofs(dofsl_sol)``, i.e. ``self.variables.dofsl = ...``: an in-place Python mutation, which is exactly what ``solve_2d_magneto_static.py``'s own vmap benchmark flags as incompatible with tracing. The fix used there generalizes directly: bypass ``solve()`` and call the class's *pure* building blocks instead -- ``Galerkin.factorise``/``Galerkin._make_back_solve_fn`` (in ``galerkin.py``), which never touch ``variables`` and just return the solved DOFs as a plain array. That pair is also the better tool here on its own merits, "use the existing infra" aside: the assembled Jacobian ``K`` comes only from ``bilinear_form``, so it is the SAME operator for every draw (only ``linear_form``, hence the residual/RHS, depends on ``f``). Factorising ``K`` once, outside the vmap, and vmapping only the cheap "assemble RHS + back substitute" step keeps a single dense ``(n_dof, n_dof)`` LU factorisation from being replicated ``N_RHS`` times -- which a naive ``jnp.linalg.solve`` inside the vmapped function would otherwise risk, since nothing would tell XLA the matrix is shared. Needs the optional ``mesh`` extra (pygmsh/meshio/gmsh). Saves ``helmholtz_robin_unstructured_disk_batched_sources.png`` next to this file and shows it. """ # %% import time from pathlib import Path 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.meshes.unstructured_mesh import UnstructuredMesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.macro_mesh import macro_mesh_from_curve 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 ComplexRobin from scimba_jax.utils.scimba_pytree import ScimbaPytree _HERE = Path(__file__).resolve().parent DIM = 2 ORDER = 3 # geometric AND FE degree: isoparametric, like the O-grid disk example MESH_SIZE = 0.1 QUAD_ORDER = 2 * ORDER + 2 K = 32.0 # wavenumber KAPPA = K**2 # ── Random Gaussian-mixture sources ───────────────────────────────────────── N_RHS = 64 # number of independent source draws, solved in one jax.vmap call N_GAUSSIANS = 3 # Gaussians per mixture (fixed, for a uniform batch shape) CENTER_RADIUS = 0.85 # centers drawn uniformly in a disk of this radius, away # from the boundary so the bumps stay resolved by MESH_SIZE/QUAD_ORDER AMPLITUDE_MIN, AMPLITUDE_MAX = 0.5, 1.5 SIGMA_MIN, SIGMA_MAX = 0.02, 0.05 N_DISPLAY = 3 # how many of the N_RHS draws to plot, picked at random SEED = 0 # %% A mixture of real Gaussians, held as a ScimbaPytree so it stays a genuine # pytree child of HelmholtzWeakForm/EllipticFEscheme -- exactly like # DiffusionWeakForm.alpha in fem/solve/time_dependent/ # 1d_heat_equation_parametric.py. A bare closure over centers/amplitudes/sigmas # would work for a single solve, but ScimbaPytree's auto-classification # (_is_dynamic_value) only looks for jnp.ndarray LEAVES: a plain Python # callable is an opaque leaf with none, so it would be classified static # (aux_data) -- and a tracer parked there by jax.vmap/jax.jit raises. Making # the source itself a ScimbaPytree gives jax.tree_util.tree_leaves something # real to find, so centers/amplitudes/sigmas correctly become dynamic children # of the whole scheme, and one draw per batch element is exactly what # jax.vmap needs. class GaussianMixtureSource(ScimbaPytree): """Real source ``[f_Re, f_Im] = [sum_i A_i * gauss_sigma_i(x - c_i), 0]``. Args: centers: Gaussian centers, shape ``(n_gaussians, dim)``. amplitudes: Per-Gaussian amplitude, shape ``(n_gaussians,)``. sigmas: Per-Gaussian width, shape ``(n_gaussians,)``. """ def __init__(self, centers, amplitudes, sigmas): self.centers = jnp.asarray(centers) self.amplitudes = jnp.asarray(amplitudes) self.sigmas = jnp.asarray(sigmas) def __call__(self, x): dist2 = jnp.sum((x - self.centers) ** 2, axis=1) bumps = self.amplitudes * jnp.exp(-dist2 / (2.0 * self.sigmas**2)) bumps = bumps / (2.0 * jnp.pi * self.sigmas**2) return jnp.array([jnp.sum(bumps), 0.0]) # %% The absorbing condition is the library's ComplexRobin (weak_boundary_conditions), # the complex Robin of HelmholtzWeakForm's sign convention: bilinear -beta u v, linear # -g v, both signs opposite to the real Robin (which pairs with -Delta u = f). def g_zero(x): """Homogeneous ABC data.""" return jnp.zeros(2) # %% Mesh: same traditional unstructured mesh of the disk, built ONCE and # shared by every draw -- it never changes, only the source does. def disk_boundary(t): return (np.cos(2.0 * np.pi * t), np.sin(2.0 * np.pi * t)) macro = macro_mesh_from_curve(disk_boundary, order=ORDER, mesh_size=MESH_SIZE) mesh = UnstructuredMesh.from_macro_mesh( macro, UnitSquareTensorized(dim=DIM, order=QUAD_ORDER) ) print(f"unstructured disk mesh: {mesh.n_cells_total} cells, degree {mesh.order}") basis = AnalyticBasis( nb_basis=(ORDER + 1) ** DIM, out_dim=2, mesh=mesh, basis_type="vec", local_basis=lambda y, i, m: local_lagrange_basis(y, i, m, order=ORDER, out_dim=2), local_basis_by_logical=lambda y, i, m: local_lagrange_basis_by_logical( y, i, m, order=ORDER, out_dim=2 ), ) variables = VariablesFE(basis=basis, nb_variables=2, dof_map=UnstructuredLagrangeDofMap) print(f"k = {K}, kappa = k^2 = {KAPPA}, {variables.n_nodes_total} nodes") def build_scheme(centers, amplitudes, sigmas) -> EllipticFEscheme: """One scheme for one Gaussian-mixture draw, sharing ``variables``/mesh.""" source = GaussianMixtureSource(centers, amplitudes, sigmas) model = AbstractPhysicalWeakModel(dim=DIM) model.add_weak_form("main", HelmholtzWeakForm(dim=DIM, kappa=KAPPA, f=source)) model.add_boundary_condition( "boundary", ComplexRobin(beta_re=0.0, beta_im=-K, g=g_zero) ) return EllipticFEscheme(model, variables) # %% Factorise the (source-independent) Jacobian once. ComplexRobin/ # HelmholtzWeakForm are both linear in u and f enters only linear_form, so the # assembled Jacobian K is the SAME operator whatever the source -- computing # and LU-factorising it here, outside jax.vmap, means the vmapped step below # only has to assemble a new right-hand side and run one cheap back # substitution per draw, instead of re-factorising a dense (n_dof, n_dof) # matrix N_RHS times (or, with a naive jnp.linalg.solve, risk XLA replicating # it across the batch). _reference_scheme = build_scheme( centers=jnp.zeros((N_GAUSSIANS, DIM)), amplitudes=jnp.zeros((N_GAUSSIANS,)), sigmas=jnp.ones((N_GAUSSIANS,)), ) _dofsl_init = _reference_scheme._initial_dofs() print("factorising the (source-independent) Jacobian ...") _factorisation = EllipticFEscheme.factorise(_reference_scheme, _dofsl_init) _back_solve = EllipticFEscheme._make_back_solve_fn() def solve_one(centers, amplitudes, sigmas) -> jnp.ndarray: """One Helmholtz solve for one Gaussian-mixture draw. Pure (no mutation of ``variables``), unlike ``EllipticFEscheme.solve`` -- see the module docstring -- so :func:`run_batch` can ``jax.vmap`` it over a whole batch of draws. Returns: The solved DOFs, shape like ``variables.dofsl``. """ scheme = build_scheme(centers, amplitudes, sigmas) return _back_solve(scheme, _dofsl_init, _factorisation.lu, _factorisation.pivots) def run_batch(centers_batch, amplitudes_batch, sigmas_batch) -> jnp.ndarray: """Solves every draw in a SINGLE compiled program via ``jax.vmap``.""" return jax.vmap(solve_one)(centers_batch, amplitudes_batch, sigmas_batch) # %% Draw N_RHS random Gaussian mixtures: centers uniform in a disk, positive # amplitudes, widths resolved by the mesh/quadrature. key = jax.random.PRNGKey(SEED) key_r, key_theta, key_amp, key_sigma, key_pick = jax.random.split(key, 5) radius = CENTER_RADIUS * jnp.sqrt(jax.random.uniform(key_r, (N_RHS, N_GAUSSIANS))) angle = 2.0 * jnp.pi * jax.random.uniform(key_theta, (N_RHS, N_GAUSSIANS)) centers_batch = jnp.stack([radius * jnp.cos(angle), radius * jnp.sin(angle)], axis=-1) amplitudes_batch = jax.random.uniform( key_amp, (N_RHS, N_GAUSSIANS), minval=AMPLITUDE_MIN, maxval=AMPLITUDE_MAX ) sigmas_batch = jax.random.uniform( key_sigma, (N_RHS, N_GAUSSIANS), minval=SIGMA_MIN, maxval=SIGMA_MAX ) print(f"solving {N_RHS} Helmholtz problems in one jax.vmap call ...") t0 = time.perf_counter() dofsl_batch = run_batch(centers_batch, amplitudes_batch, sigmas_batch) jax.block_until_ready(dofsl_batch) elapsed = time.perf_counter() - t0 print(f" done in {elapsed:.2f}s ({1e3 * elapsed / N_RHS:.2f} ms/draw)") # %% Sampling helpers. Cell-by-cell with no point search, like the base # example's ``sample_complex_solution`` -- required on a non-convex/curved # domain, where a rectangular grid would put points outside the disk. Split # in two here: the sampling grid/triangulation depends only on the (shared) # mesh, and is built once; evaluating a specific DOF field on it is cheap and # repeated per displayed draw. def build_sampling_grid(mesh, n_side=10): 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) points = jax.vmap(lambda cell: mesh._unit_hypercube_to_cell(cell, grid))( jnp.arange(mesh.n_cells_total) ) # (n_cells, n_side**2, dim) 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 points, triangles def sample_dof_field(dofs, points): connectivity = variables.connectivity def one_cell(cell, cell_points): theta = dofs[connectivity[cell]] return jax.vmap( lambda p: jnp.einsum( "iv,iv->v", theta, variables.trial_basis(cell, p[None, :])[0] ) )(cell_points) return jax.vmap(one_cell)(jnp.arange(mesh.n_cells_total), points) # %% Plot Re(u), Im(u), |u| -- and the source itself -- for N_DISPLAY draws # picked at random among the N_RHS generated mixtures. points, triangles = build_sampling_grid(mesh, n_side=10) points_flat = points.reshape(-1, DIM) points_np = np.asarray(points_flat) idx_display = np.asarray( jax.random.choice(key_pick, N_RHS, shape=(N_DISPLAY,), replace=False) ) figure, axes = plt.subplots(N_DISPLAY, 4, figsize=(18, 4.6 * N_DISPLAY), squeeze=False) for row, idx in enumerate(idx_display): idx = int(idx) source = GaussianMixtureSource( centers_batch[idx], amplitudes_batch[idx], sigmas_batch[idx] ) f_re = np.asarray(jax.vmap(source)(points_flat)[:, 0]) values = np.asarray(sample_dof_field(dofsl_batch[idx], points)).reshape(-1, 2) u_re, u_im = values[:, 0], values[:, 1] amplitude = np.hypot(u_re, u_im) for col, (field, label) in enumerate( [(f_re, "f_Re (source)"), (u_re, "Re(u)"), (u_im, "Im(u)"), (amplitude, "|u|")] ): ax = axes[row, col] drawing = ax.tricontourf( points_np[:, 0], points_np[:, 1], triangles, field, levels=40, cmap="turbo" ) figure.colorbar(drawing, ax=ax, fraction=0.046) ax.set_title(f"draw #{idx}: {label}" if col == 0 else label) ax.set_aspect("equal") ax.set_xlabel("x") ax.set_ylabel("y") figure.suptitle( f"Complex Helmholtz, k={K}, unstructured disk ({mesh.n_cells_total} cells): " f"{N_RHS} random Gaussian-mixture sources solved in one jax.vmap call, " f"{N_DISPLAY} random draws shown" ) figure.tight_layout() figure.savefig(_HERE / "helmholtz_robin_unstructured_disk_batched_sources.png", dpi=130) print(f"\nSaved {_HERE / 'helmholtz_robin_unstructured_disk_batched_sources.png'}") plt.show()