Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 44 additions & 13 deletions src/underworld3/utilities/_jitextension.py
Original file line number Diff line number Diff line change
Expand Up @@ -300,13 +300,20 @@ def getext(
import time

time_s = time.time()
primary_field_list = tuple(primary_field_list)

raw_fns_residual = tuple(fns_residual)
raw_fns_bcs = tuple(fns_bcs)
raw_fns_jacobian = tuple(fns_jacobian)
raw_fns_bd_residual = tuple(fns_bd_residual)
raw_fns_bd_jacobian = tuple(fns_bd_jacobian)

raw_fns = (
tuple(fns_residual)
+ tuple(fns_bcs)
+ tuple(fns_jacobian)
+ tuple(fns_bd_residual)
+ tuple(fns_bd_jacobian)
raw_fns_residual
+ raw_fns_bcs
+ raw_fns_jacobian
+ raw_fns_bd_residual
+ raw_fns_bd_jacobian
)

# Extract constant UWexpressions that will go through constants[] array
Expand All @@ -315,23 +322,41 @@ def getext(
# Build structurally-expanded functions for cache hashing.
# Constants are replaced with placeholder symbols (value-independent),
# so changing a constant value won't cause a cache miss.
expanded_fns = []
for fn in raw_fns:
def _structural_expand(fn):
# Phase 1: Substitute constants with _JITConstant placeholders
if constants_subs_map and fn is not None:
try:
fn_structural = fn.xreplace(constants_subs_map) if hasattr(fn, 'xreplace') else fn
fn_structural = fn.xreplace(constants_subs_map) if hasattr(fn, "xreplace") else fn
except Exception:
fn_structural = fn
else:
fn_structural = fn

# Phase 2: Unwrap remaining (non-constant) expressions
expanded_fns.append(
underworld3.function.expressions.unwrap(fn_structural, keep_constants=False, return_self=False)
return underworld3.function.expressions.unwrap(
fn_structural, keep_constants=False, return_self=False
)

fns = tuple(expanded_fns)
expanded_fns_residual = tuple(_structural_expand(fn) for fn in raw_fns_residual)
expanded_fns_bcs = tuple(_structural_expand(fn) for fn in raw_fns_bcs)
expanded_fns_jacobian = tuple(_structural_expand(fn) for fn in raw_fns_jacobian)
expanded_fns_bd_residual = tuple(_structural_expand(fn) for fn in raw_fns_bd_residual)
expanded_fns_bd_jacobian = tuple(_structural_expand(fn) for fn in raw_fns_bd_jacobian)

fns = (
expanded_fns_residual
+ expanded_fns_bcs
+ expanded_fns_jacobian
+ expanded_fns_bd_residual
+ expanded_fns_bd_jacobian
)
fns_signature = (
expanded_fns_residual,
expanded_fns_bcs,
expanded_fns_jacobian,
expanded_fns_bd_residual,
expanded_fns_bd_jacobian,
)

if debug and underworld3.mpi.rank == 0:
print(f"Expanded functions for compilation:")
Expand All @@ -353,8 +378,14 @@ def getext(
# unique modules.
jitname += "_" + str(len(_ext_dict.keys()))

else: # Else name from fns hash — uses structural form (constants as placeholders)
jitname = abs(hash((mesh, fns, tuple(mesh.vars.keys()))))
else: # Else name from a structured hash — function role/signature must be preserved.
primary_field_signature = tuple(
(getattr(field, "field_id", None), getattr(field, "clean_name", None))
for field in primary_field_list
)
jitname = abs(
hash((mesh, fns_signature, tuple(mesh.vars.keys()), primary_field_signature))
)

# Create the module if not in dictionary
if jitname not in _ext_dict.keys() or not cache:
Expand Down
43 changes: 43 additions & 0 deletions tests/test_0502_boundary_integrals.py
Original file line number Diff line number Diff line change
Expand Up @@ -291,3 +291,46 @@ def test_bd_integral_annulus_internal_normal_tangential():
value = bd_int.evaluate()

assert abs(value) < 0.05, f"Expected ~0, got {value}"


def _build_spherical_shell_for_integrals():
from underworld3.meshing import SphericalShell

mesh_spherical = SphericalShell(
radiusOuter=1.0,
radiusInner=0.5,
cellSize=1.0 / 4.0,
degree=1,
qdegree=2,
)
uw.discretisation.MeshVariable("P_spherical_int", mesh_spherical, 1, degree=1, continuous=True)
return mesh_spherical


def test_spherical_bd_then_integral_does_not_poison_volume_path():
"""Boundary and volume integrals must not collide in the JIT cache."""

mesh_spherical = _build_spherical_shell_for_integrals()

boundary_before = float(uw.maths.BdIntegral(mesh_spherical, fn=1.0, boundary="Lower").evaluate())
volume = float(uw.maths.Integral(mesh_spherical, fn=1.0).evaluate())
boundary_after = float(uw.maths.BdIntegral(mesh_spherical, fn=1.0, boundary="Lower").evaluate())

assert boundary_before > 0.0
assert volume > 0.0
assert abs(boundary_after - boundary_before) < 1.0e-10


def test_spherical_integral_then_bd_does_not_poison_boundary_path():
"""Volume and boundary integrals must remain order-independent on spherical meshes."""

mesh_reference = _build_spherical_shell_for_integrals()
boundary_reference = float(uw.maths.BdIntegral(mesh_reference, fn=1.0, boundary="Lower").evaluate())

mesh_spherical = _build_spherical_shell_for_integrals()
volume = float(uw.maths.Integral(mesh_spherical, fn=1.0).evaluate())
boundary_after = float(uw.maths.BdIntegral(mesh_spherical, fn=1.0, boundary="Lower").evaluate())

assert volume > 0.0
assert boundary_reference > 0.0
assert abs(boundary_after - boundary_reference) < 1.0e-10
Loading