"""Phase space: a 2-D domain times a velocity DIRECTION on the sphere. The four-dimensional case the tensorisation was for. Position ``x`` runs over a square, velocity direction ``omega`` over ``S^2``, and the product mesh has 4-D cells whose nodes carry 5 coordinates -- two of space and three of the sphere's ambient space. Nothing about the product knows that; it reads the intrinsic dimension of each factor and concatenates. **Why the sphere is six patches.** A latitude-longitude mesh has a coordinate singularity at each pole: cells degenerate, the Jacobian drops rank, and the quadrature there is meaningless. Each face of a cube instead maps to its part of the sphere by ``(a, b) -> normalise(tan a, tan b, 1)``, a diffeomorphism with a bounded Jacobian everywhere. Six patches, no pole. **Why the area element changes.** A sphere cell's Jacobian is ``(3, 2)`` and has no determinant; the area it stretches a reference cell to is ``sqrt(det(J^T J))``, the Gram determinant. On a product the Jacobian is block diagonal, so the Gram determinant FACTORS -- which is why a correct product of correct factors needs no new geometry, and is checked here to 1e-14. **The control.** The projected field is separable and both of its integrals are in closed form, so the whole 4-D quadrature is checked against a number owing nothing to the code: * over the square, a product of error functions; * over the sphere, ``2 pi sigma^2 (1 - exp(-2 / sigma^2))`` for the chordal Gaussian below -- exactly, since ``1 - cos`` integrates against ``sin``. """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from matplotlib.colors import Normalize from mpl_toolkits.mplot3d.art3d import Poly3DCollection from scimba_jax.linear_approximation.basis.analytic_bases import ( local_lagrange_basis, local_lagrange_basis_by_logical, ) from scimba_jax.linear_approximation.basis.general_bases import ( AnalyticBasis, basis_values, ) from scimba_jax.linear_approximation.meshes.manifold_mesh import cubed_sphere from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.meshes.tensor_mesh import by_variables, tensor_mesh from scimba_jax.linear_approximation.meshes.unstructured_mesh import ( unit_cell_quad_points, ) 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 SPACE_CELLS = 3 SPHERE_CELLS = 3 GEOMETRY_ORDER = 2 BASIS_ORDER = 2 QUAD_ORDER = 5 # Centred on the middle of the square and on the north pole of the sphere, both # wide enough for the mesh: the spatial cells are 1/3 across and the spherical # ones subtend about 0.5 rad. POSITION_CENTRE = (0.5, 0.5) SIGMA_POSITION = 0.22 VELOCITY_CENTRE = (0.0, 0.0, 1.0) SIGMA_VELOCITY = 0.7 def phase_space(): """``[0,1]^2`` times ``S^2``: 4-D cells, 5 coordinates.""" identity = Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]) position = Mesh( dim=2, n_cells=(SPACE_CELLS, SPACE_CELLS), ref_quad=UnitSquareTensorized(dim=2, order=QUAD_ORDER), mapping=identity, ) velocity = cubed_sphere( n=SPHERE_CELLS, order=GEOMETRY_ORDER, ref_quad=UnitSquareTensorized(dim=2, order=QUAD_ORDER), ) # `names` is what lets the physics below be written `f(x, omega)` instead # of as slices of a five-vector -- the PINNs' `model_type`, applied to a # mesh. return ( position, velocity, tensor_mesh(position, velocity, names=("x", "omega")), ) def field(x, omega): """``exp(-|x - x0|^2 / 2 sx^2) exp(-(1 - omega.omega0) / sv^2)``. Written of ``x`` and ``omega``, not of a concatenated vector: that is the whole point of naming the product's variables, and it is what makes this readable as physics. ⚠ The velocity factor uses the CHORDAL distance ``1 - omega . omega0``, not the angle ``arccos(omega . omega0)``. The two agree to leading order, but the angle has a kink at the antipode -- it is not differentiable there -- and the projection would see a discontinuous derivative sitting on the far side of the sphere. The chordal form is a polynomial in the ambient coordinates, smooth everywhere, and that is also what gives it a closed-form integral. Args: x: ``(..., 2)`` positions. omega: ``(..., 3)`` directions, on the unit sphere. Returns: ``(..., 1)``. """ centred = sum((x[..., d] - POSITION_CENTRE[d]) ** 2 for d in range(2)) alignment = sum(omega[..., d] * VELOCITY_CENTRE[d] for d in range(3)) return jnp.exp( -centred / (2.0 * SIGMA_POSITION**2) - (1.0 - alignment) / SIGMA_VELOCITY**2 )[..., None] def exact_integral(): """The closed form the 4-D quadrature is checked against. Separable, so it is the product of two one-dimensional facts: a Gaussian over a bounded square gives error functions, and the chordal Gaussian over the sphere integrates exactly because ``1 - cos(theta)`` meets ``sin(theta)`` -- ``2 pi sigma^2 (1 - exp(-2 / sigma^2))``. """ from math import erf, exp, pi, sqrt def over_axis(centre): scale = SIGMA_POSITION * sqrt(0.5 * pi) return scale * ( erf((1.0 - centre) / (SIGMA_POSITION * sqrt(2.0))) + erf(centre / (SIGMA_POSITION * sqrt(2.0))) ) over_square = over_axis(POSITION_CENTRE[0]) * over_axis(POSITION_CENTRE[1]) over_sphere = 2.0 * pi * SIGMA_VELOCITY**2 * (1.0 - exp(-2.0 / SIGMA_VELOCITY**2)) return over_square * over_sphere def make_basis(mesh, order=BASIS_ORDER): """Lagrange of degree ``order`` on a 4-D cell -- ``(order+1)^4`` functions.""" return AnalyticBasis( nb_basis=(order + 1) ** mesh.dim, out_dim=1, mesh=mesh, basis_type="scalar", local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=order, out_dim=1 ), local_basis_by_logical=lambda y, i, m: local_lagrange_basis_by_logical( y, i, m, order=order, out_dim=1 ), ) def evaluate(basis, variables, cells, reference): """The expansion on a batch of cells, at given unit-cell points. ⚠ Unit-cell points, always: a manifold cell has NO inverse map -- an ambient point need not even lie on the surface -- so the physical-point path is not merely slow here, it raises. """ mesh = basis.mesh def per_cell(cell, unit_points): physical = mesh._unit_hypercube_to_cell(cell, unit_points) values = jax.vmap(lambda x_hat, x: basis_values(basis, cell, x, x_hat)[0])( unit_points, physical ) return physical, jnp.einsum("iv,qiv->q", variables.dofsl[cell], values) physical, values = jax.vmap(per_cell)(cells, reference) return np.asarray(physical), np.asarray(values) def measure(mesh): """Quadrature of ``1``: the mesh's own idea of its measure.""" weights = jax.vmap(lambda c: mesh._local_weights_points(c)[0])( jnp.arange(mesh.n_cells_total) ) return float(jnp.sum(weights)) def integral_and_error(mesh, basis, variables): """``(integral of the projection, relative L2 error)``, on the quadrature.""" x_hat = unit_cell_quad_points(mesh) cells = jnp.arange(mesh.n_cells_total) reference = jnp.broadcast_to(x_hat, (mesh.n_cells_total,) + x_hat.shape) points, got = evaluate(basis, variables, cells, reference) weights = jax.vmap(lambda c: mesh._local_weights_points(c)[0])(cells) want = by_variables(mesh, field)(jnp.asarray(points))[..., 0] error = float(jnp.sum(weights * (jnp.asarray(got) - want) ** 2)) reference_norm = float(jnp.sum(weights * want**2)) return float(jnp.sum(weights * jnp.asarray(got))), float( np.sqrt(error / reference_norm) ) # ── Slices: 4-D cannot be drawn, but its factors can ──────────────────────── def slice_reference(fixed, moving_grid, fixed_first): """Reference points of the product with one factor held at ``fixed``.""" held = np.broadcast_to(np.asarray(fixed), (moving_grid.shape[0], 2)) parts = (held, moving_grid) if fixed_first else (moving_grid, held) return jnp.asarray(np.concatenate(parts, axis=1)) def grid_2d(samples): line = np.linspace(0.0, 1.0, samples) return np.stack(np.meshgrid(line, line, indexing="ij"), -1).reshape(-1, 2) def at_fixed_position(basis, variables, n_velocity, space_cell, ref_x, samples=5): """The velocity distribution at one point of space, on the sphere. A is the slowest factor, so the cells sharing space cell ``a`` are ``a * n_b + [0 .. n_b)``: the whole sphere at that ``x``, reached by fixing the first two REFERENCE coordinates. No search, no interpolation. """ moving = grid_2d(samples) cells = jnp.arange(n_velocity) + space_cell * n_velocity reference = slice_reference(ref_x, moving, fixed_first=True) physical, values = evaluate( basis, variables, cells, jnp.broadcast_to(reference, (n_velocity,) + reference.shape), ) # Ambient coordinates of the sphere are the last three. return physical[..., 2:].reshape(n_velocity, samples, samples, 3), values.reshape( n_velocity, samples, samples ) def at_fixed_velocity( basis, variables, n_velocity, velocity_cell, ref_omega, samples=5 ): """The spatial distribution for one velocity direction, on the square.""" moving = grid_2d(samples) n_space = basis.mesh.n_cells_total // n_velocity cells = jnp.arange(n_space) * n_velocity + velocity_cell reference = slice_reference(ref_omega, moving, fixed_first=False) physical, values = evaluate( basis, variables, cells, jnp.broadcast_to(reference, (n_space,) + reference.shape), ) return physical[..., :2].reshape(n_space, samples, samples, 2), values.reshape( n_space, samples, samples ) def draw_on_sphere(axes, patches, values, title): """Each cell of the sphere as a shaded quad grid, in its own place.""" norm = Normalize(vmin=float(values.min()), vmax=float(values.max())) colours = plt.get_cmap("turbo") axes.figure.colorbar( plt.cm.ScalarMappable(norm=norm, cmap=colours), ax=axes, fraction=0.03, pad=0.02, ) quads, facecolours = [], [] for patch, patch_values in zip(patches, values): rows, columns = patch.shape[0] - 1, patch.shape[1] - 1 for i in range(rows): for j in range(columns): quads.append( [patch[i, j], patch[i + 1, j], patch[i + 1, j + 1], patch[i, j + 1]] ) facecolours.append( colours(norm(patch_values[i : i + 2, j : j + 2].mean())) ) axes.add_collection3d( Poly3DCollection(quads, facecolors=facecolours, linewidths=0, edgecolors="none") ) axes.set_xlim(-1, 1) axes.set_ylim(-1, 1) axes.set_zlim(-1, 1) axes.set_box_aspect((1, 1, 1)) axes.set_title(title) axes.set_xticks([]) axes.set_yticks([]) axes.set_zticks([]) def draw_on_square(axes, patches, values, title): """Each spatial cell as its own coloured patch, WITH its scale. ⚠ The colourbar is the point of these panels, not decoration. Each slice is normalised on its own maximum, so the three look identical -- which is exactly the claim being made, that the shape does not depend on the direction. Without the scale beside it that picture would also be consistent with three fields of the same amplitude, which is the opposite of what is happening: the amplitudes differ by a factor of twenty. """ low, high = float(values.min()), float(values.max()) for patch, patch_values in zip(patches, values): mesh = axes.pcolormesh( patch[..., 0], patch[..., 1], patch_values, cmap="turbo", vmin=low, vmax=high, shading="gouraud", ) axes.figure.colorbar(mesh, ax=axes, fraction=0.046, pad=0.03) axes.set_aspect("equal") axes.set_title(title) axes.set_xticks([]) axes.set_yticks([]) def main(): started = time.perf_counter() position, velocity, mesh = phase_space() print( f"espace des phases : {mesh.n_cells_total} mailles de dimension {mesh.dim}, " f"coordonnees en dimension {mesh.ambient_dim}" ) print(f" construction {time.perf_counter() - started:6.2f} s") # The sharp geometric check: on a block-diagonal Jacobian the Gram # determinant factors, so the product's measure IS the product of the # factors' -- to machine precision, with no discretisation excuse. product_measure = measure(mesh) factored = measure(position) * measure(velocity) print( f" mesure {product_measure:.12f} contre le produit des facteurs " f"{factored:.12f} ecart {abs(product_measure - factored):.2e}" ) print( f" aire de la sphere {measure(velocity):.10f} contre 4 pi " f"{4 * np.pi:.10f} ecart relatif " f"{abs(measure(velocity) - 4 * np.pi) / (4 * np.pi):.2e}" ) basis = make_basis(mesh) variables = VariablesDG(basis=basis, nb_variables=1) started = time.perf_counter() variables.project_jit(by_variables(mesh, field)) jax.block_until_ready(variables.dofsl) n_dofs = mesh.n_cells_total * (BASIS_ORDER + 1) ** mesh.dim print( f" projection Q{BASIS_ORDER} {time.perf_counter() - started:6.2f} s" f" ({n_dofs} ddl)" ) integral, error = integral_and_error(mesh, basis, variables) exact = exact_integral() print(f" erreur L2 relative {error:9.3e}") print( f" integrale {integral:.9f} contre la forme close {exact:.9f} " f"ecart relatif {abs(integral - exact) / exact:.2e}" ) # The centre of the square sits in the middle cell, at its own centre. middle = (SPACE_CELLS // 2) * SPACE_CELLS + SPACE_CELLS // 2 sphere_patches, sphere_values = at_fixed_position( basis, variables, velocity.n_cells_total, middle, (0.5, 0.5) ) # Three velocity cells at increasing angle from the pole, to look at the # SAME spatial Gaussian through three different directions. centres = np.asarray( jax.vmap( lambda c: velocity._unit_hypercube_to_cell(c, jnp.full((1, 2), 0.5))[0] )(jnp.arange(velocity.n_cells_total)) ) alignment = centres @ np.asarray(VELOCITY_CENTRE) ordered = np.argsort(-alignment) chosen = [ int(ordered[0]), int(ordered[len(ordered) // 3]), int(ordered[2 * len(ordered) // 3]), ] figure = plt.figure(figsize=(16, 4.4)) draw_on_sphere( figure.add_subplot(141, projection="3d"), sphere_patches, sphere_values, "vitesses en x = (0.5, 0.5)", ) # ⚠ The separability check, and the sharpest thing on this page. The field # is `g(x) h(omega)`, so every spatial slice is the SAME Gaussian scaled by # `h` -- their ratio must be CONSTANT in x. It is not automatic: the basis # is a tensor product but the projection solves one 4-D mass matrix per # cell, so any error coupling the two factors (a wrong node order, a # reference axis swapped between them) would tilt the ratio across the # square while leaving each slice looking perfectly Gaussian. slices = [] for cell in chosen: patches, values = at_fixed_velocity( basis, variables, velocity.n_cells_total, cell, (0.5, 0.5) ) slices.append((patches, values, float(alignment[cell]))) print("\n coupes spatiales a differentes directions :") _, base_values, _ = slices[0] for index, (patches, values, cosine) in enumerate(slices): peak = float(values.max()) axes = figure.add_subplot(1, 4, 2 + index) draw_on_square( axes, patches, values, f"omega . omega0 = {cosine:+.2f} (max {peak:.3f})" ) if index: # Ratio to the first slice, where the field is large enough for the # quotient to mean something. large = base_values > 0.05 * base_values.max() ratio = values[large] / base_values[large] print( f" cos {cosine:+.2f} : max {peak:.4f}, " f"rapport a la 1re coupe {ratio.mean():.6f} " f"+- {ratio.std():.2e} (constant si separable)" ) figure.suptitle( "espace des phases 4D par tensorisation : carre 2D x sphere (6 patchs)" ) plt.tight_layout() plt.show() if __name__ == "__main__": main()