Skip to content
Open
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
17 changes: 9 additions & 8 deletions docs/developer/UW3_Developers_MathematicalObjects.md
Original file line number Diff line number Diff line change
Expand Up @@ -449,21 +449,22 @@ atoms = momentum.atoms(sympy.Function) # Finds V_0, V_1

## Expression Unwrapping and Constants

The JIT compiler performs a **two-phase unwrap** on expressions:
The JIT compiler lowers each expression onto the graph of named quantities
(`_jit_graph`, #823; design note `design/jit-shared-graph-codegen.md`):

1. **Phase 1 — Constants extraction**: `UWexpression` atoms that resolve to
pure numbers (no spatial/field dependencies) are replaced with
`_JITConstant` symbols that render as `constants[i]` in C code.
1. **Constants**: `UWexpression` atoms that resolve to pure numbers (no
spatial/field dependencies) are leaves, written as `constants[i]` in C code.

2. **Phase 2 — Full unwrap**: Remaining `UWexpression` atoms are expanded
to their numerical values and baked into the C code.
2. **Named quantities**: every other `UWexpression` atom becomes one C
temporary, computed once per kernel call from the leaves and earlier
temporaries; field variables are read from `petsc_a[]` and `petsc_u[]`.

```python
# Expression with nested UWexpressions
complex_expr = alpha * (temperature - T0) * velocity

# Phase 1: alpha and T0 are constants → constants[0], constants[1]
# Phase 2: temperature and velocity are field variables → petsc_a[], petsc_u[]
# alpha and T0 are constants → constants[0], constants[1]
# temperature and velocity are field variables → petsc_a[], petsc_u[]

# Generated C code:
# result = constants[0] * (petsc_a[0] - constants[1]) * petsc_u[0];
Expand Down
6 changes: 6 additions & 0 deletions docs/developer/design/jacobian-consistent-tangent.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,12 @@

**Status**: Implemented — PR #258 (2026-07-02): opt-in `solver.consistent_jacobian`, default off (Picard tangent unchanged).

**Since #823 (tier 2)** the Newton source is no longer the `symbolic_keep_constants`
expansion described below. `_jacobian_unwrap` replaces each non-constant atom by a
*node* of the shared graph (`_jit_graph`), whose partial derivatives SymPy's chain rule
uses, so the derivative sees the same `∂η/∂(grad v)` without expanding the law into a
tree. The tangent is the same function; see `jit-shared-graph-codegen.md`.

## The bug

The SNES Jacobian assembly in `src/underworld3/cython/petsc_generic_snes_solvers.pyx`
Expand Down
741 changes: 741 additions & 0 deletions docs/developer/design/jit-shared-graph-codegen.md

Large diffs are not rendered by default.

15 changes: 15 additions & 0 deletions docs/developer/guides/plasticity-solvers.md
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,21 @@ G_newton = uw.function.derive_by_array_wrt_field(F1_unwrapped, L)
# a nonzero difference == the Newton form is present
```

The solver forms the same tangent without expanding the flux: each named quantity is a
node whose partial derivatives the chain rule composes (#823), so its compiled kernel
computes each quantity once. Its blocks (`stokes._uu_G3`, ...) hold those nodes;
`unwrap_expression` expands them.

## Rule the JIT out

If a solve misbehaves in a way the physics does not explain, solve it again with the
JIT's other route: `stokes.jit_route = "expanded"` compiles the solver the way it was
compiled before #823 (every named quantity expanded into one expression), keeps the
solver's state, and rebuilds its kernels at the next solve. The two routes agree to
round-off. If the behaviour follows the route, report it as a JIT defect; if it
appears on both, the model is the place to look. `stokes.jit_route = None` returns to
the default.

Differentiate with respect to fields the way the solvers do, with
`uw.function.derive_by_array_wrt_field` (or `uw.function.diff_wrt_field` for one
entry), never plain `sympy.diff` or `sympy.derive_by_array`, in this check and in a
Expand Down
1 change: 1 addition & 0 deletions docs/developer/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,7 @@ design/declined-coord-units-proposal
design/nonlinear-solver-homotopy-warmstart
design/fault-zone-hybrid-architecture
design/eulerian-supg-transport
design/jit-shared-graph-codegen
```

```{toctree}
Expand Down
30 changes: 18 additions & 12 deletions docs/developer/subsystems/expressions-functions.md
Original file line number Diff line number Diff line change
Expand Up @@ -99,19 +99,20 @@ double result = constants[0] * velocity_gradient; // update via PetscDSSetConst

### What Happens Automatically

1. **Constant detection** — Before JIT compilation, `_extract_constants()`
scans all expression trees for `UWexpression` atoms whose fully-unwrapped
value is a pure number. This works at any nesting depth (user expression →
constitutive model parameter → solver template).
1. **Constant detection** — The JIT lowers every callback onto the graph of
named quantities (`_jit_graph`, #823). A `UWexpression` whose fully-unwrapped
value is a pure number is a *leaf* of that graph, at any nesting depth (user
expression → constitutive model parameter → solver template); every constant
leaf the kernels read gets a `constants[i]` slot, ordered by name, then
creation order.

2. **Structural hashing** — The JIT cache key is computed from the
*structural* form of expressions (constants replaced with placeholders).
Changing a constant value produces the same hash → cache hit → no
recompilation.
2. **Structural hashing** — The JIT cache key is computed from the generated C,
in which each constant is written as its slot. Changing a constant value
produces the same hash → cache hit → no recompilation.

3. **Two-phase unwrap** — During code generation:
- Phase 1: constant UWexpressions → `_JITConstant` symbols (render as `constants[i]`)
- Phase 2: remaining UWexpressions → numerical values (baked into C code)
3. **Lowering** — During code generation, each non-constant `UWexpression`
becomes one C temporary, computed from field values, coordinates, the
`constants[i]` slots and other temporaries; a number is written as a literal.

4. **Runtime update** — Before every `snes.solve()`, the solver calls
`_update_constants()` which packs current values from the manifest
Expand Down Expand Up @@ -147,9 +148,14 @@ underlying values. Two modes are used internally:

| Mode | Purpose | Used by |
|------|---------|---------|
| `nondimensional` | Numeric values for JIT/evaluate | `_createext()`, `evaluate()` |
| `nondimensional` | Numeric values for constants and evaluate | `_pack_constants()`, `evaluate()` |
| `dimensional` | Display values with units | `print()`, notebooks |

The JIT does not unwrap whole kernels: it lowers each named quantity to a node of its
own (`_jit_graph`, #823). A solver's Newton Jacobian blocks (`_uu_G3`, ...) hold those
nodes; both `unwrap_expression` and `unwrap_for_evaluate` expand a node to its body, so
code that evaluates a block sees the expanded expression.

```python
# Nested expressions
alpha = uw.expression("alpha", 3e-5)
Expand Down
45 changes: 34 additions & 11 deletions docs/developer/subsystems/jit-cache.md
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,28 @@ Setting `UW_JIT_CACHE=0` (or `false` / `no`) disables the on-disk cache
entirely — the in-memory dict still works, but nothing is persisted across
processes.

## What is compiled

There are two routes. The default, `"graph"`, lowers each callback onto the shared
graph of named quantities
(`src/underworld3/utilities/_jit_graph.py`, design note
`docs/developer/design/jit-shared-graph-codegen.md`): each non-constant `UWexpression`
becomes one C temporary, evaluated once per kernel call (a matrix-valued atom is
expanded in place), and the Newton tangent is
formed through those quantities by the chain rule. The temporaries are ordered and merged
by a hash of the C each computes, with every leaf written as the C the kernel reads
(`petsc_u[3]`, `petsc_x[0]`, `constants[2]`). The generated source, and so the key below,
is therefore a function of the mathematics and the data layout only: the same under any
`PYTHONHASHSEED`, in any process, whatever the script created first.

The `"expanded"` route is the JIT before #823 tier 2, kept as a fallback and a
reference: every named quantity is expanded into one expression tree, differentiated
and printed whole. Select it for the process with `uw.use_jit_route("expanded")` or
`UW_JIT_ROUTE=expanded`, or for one solver with `solver.jit_route = "expanded"`. The
two routes generate different C for any law with a named non-constant quantity, so
they never share a cache entry; for a law without one they generate the same C and
share it, which is correct, because the same C is the same function.

## Cache key

The key is the SHA-256 of the **canonical** generated C source plus an
Expand Down Expand Up @@ -104,32 +126,33 @@ The next solve will repopulate.

When `mpi.size > 1`:

- Every rank computes the C-source hash independently. The hashes are
`comm.allgather`'d and compared — a mismatch raises immediately rather
than letting ranks diverge. Non-determinism in `generate_c_source` (e.g.
set/dict iteration order leaking into emitted C) would land here.
- Rank 0 writes the cache entry; other ranks rely on the disk-cache hit
path on subsequent calls.
- Every rank computes the C-source hash independently, and the hashes are
`comm.allgather`'d and compared (`_agree_source_across_ranks`). If they
differ, every rank adopts rank 0's source, rehashes and warns: the ranks
compile one module. Canonical emission makes the source rank-independent by
construction; the check stays as a guard (#752).
- Whether to compile is decided collectively: if any rank lacks the module,
rank 0 compiles and publishes it while the others wait at a barrier, then
load it from the disk cache. With the disk cache off, every rank compiles.
- The `flock` on the per-hash lockfile serialises cross-shell concurrent
writes (e.g. a `mpirun -np 4` and a `mpirun -np 2` started seconds apart
on the same machine).

A future refinement would have rank 0 compile while other ranks wait on a
barrier and read the resulting `.so` directly; today every rank still
performs the cold compile but only rank 0 publishes the result.

## Environment variables

| Variable | Effect |
|----------------------|---------------------------------------------------------|
| `UW_JIT_CACHE` | Set to `0`/`false`/`no` to disable disk cache |
| `UW_JIT_CACHE_DIR` | Override the cache directory location |
| `XDG_CACHE_HOME` | Used when `UW_JIT_CACHE_DIR` is unset |
| `UW_JIT_ROUTE` | `graph` (default) or `expanded`, the JIT route; `uw.use_jit_route()` and `solver.jit_route` override it |

## Code references

- `src/underworld3/utilities/_jitextension.py` — `getext`, `generate_c_source`,
`compile_and_load`, `_abi_salt`, `_extract_constants`.
`compile_and_load`, `_abi_salt`, `_extract_constants`, `_leaf_spellings`.
- `src/underworld3/utilities/_jit_graph.py` — the lowering: nodes, their
derivatives, canonical emission.
- `src/underworld3/utilities/_jit_cache.py` — disk cache: `get_cache_dir`,
`load_module`, `store_module`, `_file_lock`.
- `uw` (shell driver) — `.env-fingerprint` write + cache wipe inside
Expand Down
85 changes: 85 additions & 0 deletions scripts/sessions/jit_graph/assembly_split.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
"""How much of an assembly is the pointwise callbacks (#823, tier 2).

Times the residual and Jacobian assembly of a fixture at a saved state, then swaps
its constitutive law for a constant viscosity on the same mesh, fields and boundary
conditions and times the same assemblies again. The constant law's callbacks cost a
few nanoseconds, so its time is the assembly machinery (element loop, tabulation,
quadrature, insertion); the difference is the fixture's callbacks. Divided by the
number of quadrature points, it is the callbacks' cost per point. The route is chosen
by the build: run it in a build of this branch
(graph) and in a build of development (tree).
"""
import importlib.util
import os
import sys
import time

import numpy as np
import underworld3 as uw

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import fixtures # noqa: E402

params = uw.Params(
uw_fixture=uw.Param("notch", description="fixture"),
uw_state=uw.Param("~/+Simulations/jit_graph/tier2/notch_newton_tree.npz",
description="route_ab .npz whose X is the state"),
uw_repeat=uw.Param(20, description="assemblies timed"),
)
route = ("graph" if importlib.util.find_spec("underworld3.utilities._jit_graph")
else "tree") # the route is the build's: this branch, or development
fixture = str(params.uw_fixture)
stokes, _ = fixtures.build(fixture)
kwargs = fixtures.prepare_solve(stokes, fixture)
stokes.petsc_options["snes_max_it"] = 0
X_state = np.load(os.path.expanduser(str(params.uw_state)))["X"]


def timed_assembly(label, solve_kwargs):
stokes.solve(**solve_kwargs) # (re)builds the kernels; no iteration
snes = stokes.snes
X = snes.getSolution()
X.array[:] = X_state
F = X.duplicate()
J, P = snes.getJacobian()[0], snes.getJacobian()[1]
snes.computeFunction(X, F)
snes.computeJacobian(X, J, P)
n = int(params.uw_repeat)
# the minimum of single assemblies: the least disturbed by other load
t_res = t_jac = float("inf")
for _ in range(n):
t = time.perf_counter()
snes.computeFunction(X, F)
t_res = min(t_res, time.perf_counter() - t)
for _ in range(n):
t = time.perf_counter()
snes.computeJacobian(X, J, P)
t_jac = min(t_jac, time.perf_counter() - t)
print(f"[{fixture} {route} {label}] residual {1e3 * t_res:.2f} ms, "
f"Jacobian {1e3 * t_jac:.2f} ms (minimum of {n}); Jacobian and "
f"preconditioner {'one matrix' if J.handle == P.handle else 'two matrices'}")
return t_res, t_jac


law = timed_assembly("law", kwargs)

# quadrature points of the velocity block
dm = stokes.mesh.dm
c0, c1 = dm.getHeightStratum(0)
nq = len(stokes.dm.getField(0)[0].getQuadrature().getData()[1])
points = (c1 - c0) * nq

cm = uw.constitutive_models.ViscousFlowModel
stokes.constitutive_model = cm
stokes.constitutive_model.Parameters.shear_viscosity_0 = 1.0
stokes.is_setup = False
# the constant law has no stress history: a viscoelastic fixture's history store and its
# timestep go with the old law
stokes.Unknowns.DFDt = None
floor = timed_assembly("constant viscosity",
{k: v for k, v in kwargs.items() if k != "timestep"})

print(f"[{fixture} {route}] {c1 - c0} cells x {nq} quadrature points = {points}")
for name, a, b in (("residual", law[0], floor[0]), ("Jacobian", law[1], floor[1])):
print(f"[{fixture} {route}] {name}: callbacks {1e3 * (a - b):.2f} ms of {1e3 * a:.2f} ms"
f" ({100 * (a - b) / a:.0f}%), {1e9 * (a - b) / max(points, 1):.0f} ns per point per assembly")
Loading
Loading