From 5baae402fb3afeba6a9e4ccf06041a4d39c69154 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Wed, 7 Oct 2026 22:32:51 +1100 Subject: [PATCH 01/14] Design note: generate JIT kernels from the shared expression graph (#823, tier 2) Tier 2 of #823, measured against tier 1 (#830) as the competitor. The note sets out where the JIT should be (named quantities as the unit of compilation, SymPy acting locally, one walk for C and manifest, canonical and readable C), the mechanism (each named quantity an applied function of its leaves whose fdiff gives the partial per argument slot, so SymPy's chain rule serves every derivative call site), the scorecard against tier 1, the risks with the test that closes each, the benchmark plan and the staging. scripts/sessions/jit_graph holds the prototype that the measurements come from: kernel_graph.py (the lowering and canonical emission), graph_vs_library.py (residual, Picard and Newton kernels compiled both ways and compared), fixtures.py, library_setup_profile.py and graph_size_probe.py. Underworld development team with AI support from Claude Code --- .../design/jit-shared-graph-codegen.md | 575 ++++++++++++++++++ docs/developer/index.md | 1 + scripts/sessions/jit_graph/fixtures.py | 105 ++++ .../sessions/jit_graph/graph_size_probe.py | 66 ++ .../sessions/jit_graph/graph_vs_library.py | 261 ++++++++ scripts/sessions/jit_graph/kernel_graph.py | 278 +++++++++ .../jit_graph/library_setup_profile.py | 86 +++ 7 files changed, 1372 insertions(+) create mode 100644 docs/developer/design/jit-shared-graph-codegen.md create mode 100644 scripts/sessions/jit_graph/fixtures.py create mode 100644 scripts/sessions/jit_graph/graph_size_probe.py create mode 100644 scripts/sessions/jit_graph/graph_vs_library.py create mode 100644 scripts/sessions/jit_graph/kernel_graph.py create mode 100644 scripts/sessions/jit_graph/library_setup_profile.py diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md new file mode 100644 index 000000000..df109a34b --- /dev/null +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -0,0 +1,575 @@ +# Generating JIT kernels from the shared expression graph + +**Status**: Proposed, 2026-10-07. Tier 2 of [#823](https://github.com/underworldcode/underworld3/issues/823). +Tier 1 ([#830](https://github.com/underworldcode/underworld3/pull/830), +`bugfix/jit-setup-cost-823`) is the fix for today and lands first: the memoised +unwrap, realness for field values, coordinates and constant slots, derivatives with +respect to fields through `uw.function.diff_wrt_field`, and the power-mean sharpness as +an atom of its own. Tier 2 is the design we would have chosen from the start, and tier +1 is the competitor it is measured against: no slower, significantly more robust and +easier to maintain. + +## Summary + +A constitutive law in Underworld3 is a graph: named sub-expressions (`UWexpression` +atoms) that refer to one another, with constants and mesh-variable symbols at the +leaves. The JIT expands that graph into a tree before it differentiates and prints it. +On the Spiegelman notch the viscosity is a graph of 137 distinct nodes and a tree of +12,191, and every stage after the expansion — the derivative, the constant scans, the C +printer, the C compiler and the kernel at run time — works on the tree. + +We propose that the JIT keeps the graph. Each named sub-expression is lowered once, to +a C temporary, and the Jacobian is formed by the chain rule through the named +sub-expressions instead of by differentiating the expanded tree. A prototype builds the +Newton `uu_G3` block of the notch in half a second against 9 s for the library's route +on tier 1, emits 6.6 KB of C against 3.2 MB, and evaluates it in an eighth of the time. +The compiled kernels — residual, Picard and Newton — agree with the library's to +round-off, and the generated source depends neither on the hash seed nor on what the +script created before the law. + +Tier 1 has removed most of the setup time the tree used to cost: on the notch the +library's full Newton setup went from 181 s to 16 s. What tier 1 cannot remove is +proportional to the tree — the size of the generated C, the compiler's time and memory, +and the kernel's run time at every quadrature point of every assembly — and the tree +grows with every named layer a law gains. + +## Where we want to be + +The JIT should compile the model the user wrote, not a copy of it with every name +erased. Five properties follow: + +1. **Named quantities are the unit of compilation.** Each is lowered once, evaluated + once per quadrature point, and differentiated once per argument; the tangent is + assembled from those derivatives by the chain rule. +2. **SymPy's algebra acts locally.** Simplification, assumption queries and + differentiation run on one named quantity's body at a time, never on the whole + expanded model, so their cost and their surprises stay the size of a body. +3. **One walk produces the C and the constants manifest**, so the two cannot disagree + and no guard has to check that they do. +4. **The generated C is a pure function of the mathematics and the data layout.** No + class is mutated to make it print, and no name, counter, hash seed or rank enters it. +5. **The generated C is readable.** One line per named quantity, so a kernel can be + checked by eye and a wrong one can be found. + +Today's JIT has none of these. Tier 1 makes it fast enough and correct on the cases we +have found, with patches placed where each failure surfaced. + +## The current pipeline expands every named sub-expression + +A solver hands the JIT a `JITCallbackSet` of residual and Jacobian expressions. +`getext()` turns them into one C function per callback. For a Newton tangent: + +1. The solver collects its residual sources `F0` and `F1`, which still contain the + named sub-expressions as symbols. +2. `_jacobian_source` passes each source through `_jacobian_unwrap`, which expands + every non-constant atom recursively (`unwrap_expression(mode="symbolic_keep_constants")`) + and then adds $10^{-36}$ to the base of every half-integer power whose base has free + symbols, constant atoms included (the sqrt guard, so that the tangent is finite at a + state of rest). +3. The solver differentiates the expanded sources into the `G0`–`G3` blocks with + explicit loops of `sympy.diff`. +4. `getext()` builds the constants manifest by testing every atom with + `_is_truly_constant`, which itself unwraps the atom completely. For each callback it + then reveals nested constants (a second complete unwrap), replaces the manifested + constants by `_JITConstant` placeholders, unwraps what remains, scans the result + for unconvertible symbols, and prints each output entry with SymPy's C99 printer. +5. The C source is hashed; the hash is the cache key. On a miss, rank 0 compiles the + module and the other ranks load it. + +The Picard tangent skips step 2. The atoms stay symbols when the solver differentiates, +so the derivative treats the effective viscosity as a constant, and `getext()` expands +them only when it prints. The residual is never guarded. A constitutive model's own +`flux_jacobian`, when it supplies one, also skips step 2 and is differentiated as given. + +## The expanded tree is the cost, and it grows with every named layer + +The notch viscosity (ViscoPlastic, Drucker–Prager yield with residual strength, +viscosity floor, the campaign law in `notch_model.py`) on tier 1 (`eb167214`; #830 generates byte-identical C), counted by +`scripts/sessions/jit_graph/graph_size_probe.py` without expanding anything: + +| | graph | expanded tree | +|---|---:|---:| +| viscosity | 137 distinct nodes; 6 named non-constant atoms, 12 constants | 12,191 nodes | +| `F1`, all four entries | the same 137 nodes | 73,548 nodes | + +The tree is large because each named atom is copied into every place that refers to +it, and the copies nest: the yield stress appears inside the plastic viscosity, which +appears inside the effective viscosity, which appears in every entry of `F1` and, after +differentiation, in every entry of `G3` — $d^4$ entries, 16 in two dimensions and 81 in +three. Each layer of named sub-expression a law gains multiplies the tree. + +Issue #823 profiled the notch setup on the soft-min law of PR #794. Before tier 1, the +keep-constants unwrap of the viscosity alone took 563 s, because each fixed-point pass +re-traversed the growing tree and re-tested every atom for constancy, and a declared +`plastic_rate_strengthening` took the setup past 40 minutes. Tier 1 makes the unwrap a +single memoised pass, removes the repeated scans, keeps fields real through +differentiation, and holds the power-mean sharpness as an atom of its own; the notch's +full Newton setup falls from 181 s to 16 s. What remains is the work proportional to the +tree — differentiating it, printing it, compiling it and running it. #547 recorded a +1.45 MB header and 1.3 GB of `gcc -O3` memory on a 16-phase collision model. At run time +the expanded kernel evaluates the viscosity's tree once for each output entry that +contains it, unless the compiler recognises the repetition, and it runs at every +quadrature point of every residual and Jacobian assembly. + +## The proposal: lower each named sub-expression once + +### A node is an applied function of the leaves its value depends on + +During lowering, each non-constant atom becomes a **node**: an applied undefined +function whose arguments are the leaves its value depends on, and whose body is its +`.sym` with each child atom replaced by the child's node. The leaves are the symbols the +kernel reads: + +| leaf | C | +|---|---| +| field component and gradient (unknowns) | `petsc_u[i]`, `petsc_u_x[i]` | +| field component and gradient (auxiliary) | `petsc_a[i]`, `petsc_a_x[i]` | +| coordinate, boundary normal (base scalars) | `petsc_x[i]`, `petsc_n[i]` | +| constant atom | `constants[i]` | + +A node is identified by its body: two atoms whose lowered bodies are equal share one +node. SymPy's equality sees the difference between two constants that share a display +name — the two $\eta$ of a two-material model — so bodies that use them are different +bodies and different nodes. The node's class name is a hash of its body written with +display identities only (a field's function name, a coordinate's system and index, a +constant's name), and its arguments are sorted by those identities, ties broken by +relative creation order as the manifest breaks them. Neither contains a creation +counter. The class also carries the serial number of the compile that made it in its +SymPy identity (`UndefinedFunction(..., _ctx=serial)`), so classes from two compiles +never compare equal and SymPy's cache cannot return a node whose body belongs to an +earlier version of the law. + +Inside a body, a sub-expression that repeats without a name is split into an anonymous +node by common sub-expression elimination on that body alone, which is cheap because the +body is small. + +### Derivatives follow from SymPy's own chain rule + +A node's `fdiff(i)` returns its partial derivative with respect to argument slot $i$, +and SymPy's `Function._eval_derivative` assembles + +$$ +\frac{\partial n(a_1,\dots,a_k)}{\partial v} \;=\; \sum_{i=1}^{k} n_{,i}(a_1,\dots,a_k)\,\frac{\partial a_i}{\partial v}. +$$ + +So `sympy.diff`, `derive_by_array`, `Matrix.diff`, second derivatives and derivatives +with respect to a constant (a sensitivity) work through nodes with no change at the call +site. Tier 1 routes every Jacobian derivative in the solvers through +`uw.function.diff_wrt_field` / `uw.function.derive_by_array_wrt_field`, which differentiate with respect to a +field through a stand-in symbol that keeps the field's realness: the field is +`xreplace`d by the stand-in, the expression is differentiated, and the field is +substituted back. Nodes pass through it: the stand-in reaches their arguments, the +derivative follows their `fdiff`, and the substitution rebuilds them. On both prototype +fixtures, every derivative of the lowered `F1` with respect to the velocity gradient and +the pressure is identical through the stand-in and through `sympy.diff`. Tier 2 changes +none of those call sites. + +The partial derivative is a node too, built on demand. Five details carry its +correctness: + +- **It is keyed by argument slot, never by the leaf object.** Before calling `fdiff`, + SymPy replaces the differentiation variable by a `Dummy`, so the arguments `fdiff` + sees are not the leaves the body was written in. A prototype that looked the leaf up + by object returned zero. +- **It is a partial derivative.** The body is differentiated with every leaf replaced + by an independent dummy symbol. Differentiated directly, a field that is itself a + function of the coordinates would be differentiated a second time through a + coordinate argument. +- **Its arguments are the leaves its own body uses.** A constant that differentiates + away does not remain an argument, so it does not reach the manifest. +- **A derivative that is zero, a number or a single leaf is returned as itself**, not + as a node. Cancellation between levels of the graph is then seen when the lower level + is a number ($A = x - y$, $B = A + y$ gives $\partial B/\partial y = -1 + 1 = 0$). + Deeper symbolic cancellation is not seen: such an entry stays an expression that + evaluates to zero. No solver tests a Jacobian block for zero, so the cost is the + assembly of that entry, not a wrong answer. +- **A node cannot be a `Symbol` subclass.** `Derivative` returns zero before it calls + `_eval_derivative` when the variable is an applied function that does not occur + visibly in the expression — and the strain-rate symbols $L_{ij}$ are applied + functions. Measured: a `Symbol` subclass whose `_eval_derivative` returns the right + answer is differentiated to zero by `sympy.diff`. An applied function shows its + dependencies as arguments, so the variable is visible. + +Body partials are taken with respect to real dummy symbols, the same remedy tier 1's +stand-in applies at the call sites, for the same reason: with a real field $P$, +`sqrt((C + μP − τ)**2)` becomes `Abs(C + μP − τ)`, and SymPy 1.14 differentiating that with +respect to $P(x, y)$ through its own assumption-free dummy leaves `Derivative(P, P)` +unevaluated, which the C printer refuses. + +### Picard, Newton and continuation keep their meaning + +The tangent mode is decided, as now, by whether an atom is visible to the derivative. +Each atom has two node variants, and they are different nodes: + +- **plain**: the body as written. The residual and the Picard tangent use it; the + residual is never guarded. +- **guarded**: the body with the sqrt guard applied. `_jacobian_unwrap` replaces each + non-constant atom of a Newton source by its guarded node instead of expanding it, so + the derivative passes through it. + +In the Picard tangent the atom stays a `UWexpression` symbol while the solver +differentiates, so the coefficient is frozen, and the JIT lowers it to its plain node +when it emits the kernel. The continuation blend contains both, and each lowers to its +own temporaries. + +The guard moves from the expanded tree to each node body, under the same rule. The two +placements could differ only where SymPy merges powers across a sub-expression +boundary: a bare `sqrt(g)**(-2/3)` becomes `g**(-1/3)`, which is no longer a +half-integer power. A real law raises a quotient — $(\dot\varepsilon_{II}/\dot\varepsilon_0)^{1/n-1}$ — +which SymPy does not merge, and on a power law over a named invariant, with $n$ as a +constant atom or as the number 3, both routes give a finite Newton tangent at a state of +rest, equal entry for entry. + +### One lowering per setup, read back by `getext()` + +`_setup_pointwise_functions` creates one lowering context, and `_jacobian_unwrap` builds +its nodes in it. Each node class holds its body and its context, so `getext()` reads +every body from the nodes it is handed and needs no new argument. `getext()` lowers the +remaining atoms — those of the residual and the Picard blocks — in the same way; a +caller that is not a solver (`Integral`, `CellWiseIntegral`, `BdIntegral`) gets a +context of its own. + +`getext()` no longer runs `unwrap_expression` over a whole kernel. Its present phases — +reveal the constants, substitute them, unwrap the rest — become: lower atoms to nodes, +collect the manifest from the leaves, emit. + +### A node expands to its body on request + +Code outside the JIT that evaluates a Jacobian block must still see a plain expression: +`test_1066` builds its finite-difference oracle by passing `_uu_G3` through +`unwrap_expression(mode="nondimensional")` and `lambdify`. The unwrappers therefore +treat a node application as one more atom, whose one-level expansion is its body with +its arguments substituted. There are two: `unwrap_expression`, whose walkers +(`_unwrap_expression_once` and tier 1's complete expansion) look at `free_symbols`, which +never contains an application, and `unwrap_for_evaluate`, which `uw.function.evaluate` +uses. Both must visit node applications. `pure_sympy_evaluator` classifies any function +whose `__module__` is `None` as mesh-variable data; a node's class has `__module__` +`None`, so evaluation unwraps first. `getext()` must have stopped calling +`unwrap_expression` on whole kernels before the unwrappers learn to expand nodes, or a +kernel expands silently back into the tree. + +A block unwrapped completely is the same function as today's tree. It is not always the +same expression: where SymPy merged a power across a sub-expression boundary, the tree +escaped the guard and the expanded nodes keep it. + +### Emission writes one temporary per distinct computation + +For each callback the JIT gives every node application reachable from its outputs a +**canonical key**: a hash of its body with each leaf written as the C the kernel reads +(`petsc_u[3]`, `petsc_x[0]`, `constants[2]`) and each child replaced by the child's key. +Applications with equal keys compute the same C and share one temporary. The +temporaries are written so that each follows the ones it uses, ties broken by key: + +```c +const double t0 = ...; /* one line per distinct computation, evaluated once */ +const double t1 = ...; +out[0] = ...; /* outputs in terms of t0, t1, ... and the leaves */ +``` + +The generated source is therefore a function of the mathematics and of the kernel's data +layout alone. Python class names, creation counters, object identities and set +iteration order never reach it. Each body passes through the existing validation for +unconvertible symbols and integration-point derivatives. + +### The constants manifest is built from the leaves + +The manifest is the set of constant atoms among the leaves of the emitted kernels, +tested with the same predicate (`_is_truly_constant`, memoised once per setup) and +ordered by the same key (name, then creation order). It contains every slot today's +manifest contains. It can contain more: where the tree cancels a constant across a name +boundary ($A = c\,x$, then $A/c$), the graph keeps it, and the constant keeps its slot. +An extra slot is harmless — it is packed and read — but each one must be traced to such +a cancellation. The prototype finds none on its two fixtures, and the benchmark adds a +fixture whose Newton source differs from its residual (`set_jacobian_F1_source`). + +### The cache key and the parallel agreement keep their mechanism + +The key remains the hash of the generated source. Every cached module recompiles once +after the change, because the source text changes; the ABI salt already includes the +Underworld version. + +The prototype's generated C is byte-identical under `PYTHONHASHSEED` 0, 1 and 2, after +a preamble that creates extra objects before the law (shifting every creation counter), +and when the law is declared twice in one process, as a re-run notebook cell does. +Whether canonical emission also removes the cross-rank disagreement of #752 is a +hypothesis we test (np ≥ 3, counting how often `_agree_source_across_ranks` repairs); the +repair stays. + +## The prototype agrees with the library to round-off + +`scripts/sessions/jit_graph/kernel_graph.py` is the lowering; +`scripts/sessions/jit_graph/graph_vs_library.py` builds three kernels of a fixture twice +and compiles each into a C function of the same leaves: + +- the residual flux `F1`, lowered with plain nodes; the library route unwraps it as + `getext()` does (keep-constants); +- the Picard `uu_G3`: `F1` differentiated with the atoms frozen, then lowered with plain + nodes or unwrapped; +- the Newton `uu_G3`: the library's own `_jacobian_source`, or guarded nodes. + +Derivatives in both routes go through tier 1's `diff_wrt_field` in the solver's explicit +loops. Two fixtures: + +- **box**: `ViscoPlasticFlowModel` with a Drucker–Prager yield stress $C + \mu p$, a + temperature-dependent viscosity $\eta_0 e^{-\theta T}$, a yield-stress floor and a + viscosity floor, on a unit square. Built in the script. +- **notch**: the campaign law (`~/+Simulations/spiegelman_hardcase/drivers/notch_model.py`), + with the power-mean soft minimum. + +Measured on tier 1 at `eb167214` (#830 generates byte-identical C) with Apple clang `-O3`, while another session's test +suite held six of the machine's eight cores; the run times are indicative. One kernel +call is timed at one state, with the cost of an empty kernel's `ctypes` call (145–165 ns) +subtracted. Library time first, graph time second: + +| | box | notch | +|---|---|---| +| residual `F1`: unwrap or lower | 0.08 s / 0.04 s | 0.07 s / 0.04 s | +| Newton source or lowering | 0.06 s / 0.03 s | 0.05 s / 0.04 s | +| Newton differentiate | 0.36 s / 0.15 s | 3.5 s / 0.35 s | +| Newton print or emit | 0.36 s / 0.04 s | 5.0 s / 0.08 s | +| C source: residual | 5.3 KB / 0.8 KB | 75 KB / 1.4 KB | +| C source: Picard `G3` | 8.5 KB / 0.8 KB | 125 KB / 1.4 KB | +| C source: Newton `G3` | 91 KB / 3.3 KB (20 temporaries) | 3.2 MB / 6.6 KB (65 temporaries) | +| `cc -O3`, Newton `G3` | 0.41 s / 0.58 s | 0.80 s / 0.40 s | +| one call: residual | 31 ns / 22 ns | 195 ns / 83 ns | +| one call: Picard `G3` | 51 ns / 32 ns | 355 ns / 91 ns | +| one call: Newton `G3` | 196 ns / 23 ns | 984 ns / 125 ns | + +| agreement | box | notch | +|---|---|---| +| states | 2,000 | 3,000 | +| largest difference, relative to the block's largest entry: residual, Picard, Newton | $6.2\times10^{-16}$, $6.0\times10^{-15}$, $6.1\times10^{-15}$ | $9.2\times10^{-16}$, $4.4\times10^{-16}$, $4.4\times10^{-16}$ | +| bit-identical entries: residual, Picard, Newton | 93%, 88%, 65% | 79%, 64%, 34% | +| NaN in one route, or in both | none | none | +| entries the library computes as exactly 0 and the graph as round-off | 31 of 17,409 (Newton) | none | +| constants manifest | identical, 7 slots | identical, 12 slots | +| structural zeros | identical | identical | +| state of rest, all three kernels | finite, elementwise equal | finite, elementwise equal | +| generated C, all three kernels, both routes | byte-identical under `PYTHONHASHSEED` 0, 1, 2, after a preamble of extra objects, and with the law declared twice | — | + +The states draw the velocity gradient over eight decades and every other field from a +normal distribution, except fields that are bounded by construction: the box's +temperature and the notch's material fraction lie in $[0, 1]$, and the notch's rate cap +is non-negative. Entries that differ by more than $10^{-12}$ of themselves are round-off +zeros, at most $10^{-16}$ of the block's largest entry, which cancel in both routes. + +Before its last two commits, tier 1's route spent 43 s on the notch's Newton source and +72–131 s differentiating it. The source's time went into four `im(base)` constructions +inside `Pow.__new__` on whole expanded bases: the notch fixture uses the power-mean soft +minimum (opt-in; the default is the square-root form), whose exponent +$s = 1/(\delta + 0.001)$ has a sum in its denominator, and SymPy then evaluated the sign +of the base's imaginary part on every rebuild of the tree. Holding $s$ as a constant +atom of its own removed it. The graph never met that cost, because the bases of its +powers are small. What tier 1 cannot reduce is proportional to the tree: megabytes of C +and a kernel several times slower. + +Outside those bounds the routes can disagree. With a material fraction above one, the +notch law's linear blend of yield stresses is negative, the soft maximum with a zero +floor is then exactly zero, and the law divides by it. The graph computes the yield +stress once, the soft maximum cancels exactly, and the tangent is NaN; the tree has +distributed constants differently in its two copies of the yield stress, so the +cancellation leaves rounding noise and the tangent is finite but meaningless. The law is +singular there. The notch cannot reach it — its material fraction is a P0 field holding +exactly 0 or 1 — but a model with a projected or higher-degree material field could. + +## Against tier 1: no slower, more robust, less code + +Tier 1 is the competitor. Each criterion is measured against it, not against the code +before it. + +**No slower.** On every fixture so far, every stage is as fast or faster: lowering, +differentiating and emitting the notch's Newton block take 0.5 s against 9 s, the C is +6.6 KB against 3.2 MB, and the kernel runs in about an eighth of the time. On a +constant-viscosity Stokes the two routes emit byte-identical C, because a constant law +has no node to lower. On a small power law the graph spends 0.02 s more differentiating +(0.09 s against 0.07 s), and the run time of its small kernels is equal within the +noise of a loaded machine. The claim needs the benchmark plan's measurement: an idle +machine, the minimum of repeated runs, the library's whole setup end to end, on Linux +as well as macOS. + +**More robust.** Each failure class below needed a patch in tier 1, placed where it +surfaced; in the graph it cannot arise, or arises only in one body: + +| failure class | tier 1 | graph | shown | +|---|---|---|---| +| SymPy's automatic algebra on an expanded base is slow (`im()` inside `Pow.__new__`, 40 s on the notch) | a constant atom added to the power-mean law; realness for unit-carrying parameters | bases are bodies, not trees | graph lowering took 0.04 s on every tier 1 commit, including those where the library took 40 s | +| realness turns `sqrt(x**2)` into `Abs`, whose derivative with respect to a field SymPy leaves unevaluated | a stand-in derivative at 49 call sites | body partials are taken against real dummies | the graph differentiated the Drucker–Prager floor law cleanly with plain `sympy.diff` where the library failed | +| manifest and C built by two walks (#302) | two consistency guards and `_reveal_constants` | one walk | by construction; manifests identical on all fixtures | +| generated C that differs between ranks (#752, open) | rank 0's source adopted | canonical emission | byte-identical under hash seeds, preambles and re-declaration; ranks not yet tested | +| field symbols given their C names by mutating their classes, in an order that matters | `ccode_patch_fns`, the coordinate recovery block | an explicit map from leaf to C | by construction | +| generated C too large to read or to compile (#547) | opt-in CSE, lower optimisation flags | one line per named quantity | 6.6 KB against 3.2 MB on the notch | + +The graph brings one failure class tier 1 does not have: it cannot cancel a quantity +against its own reciprocal across a name. + +**Less code, less global state.** By line ranges in `_jitextension.py`, the graph +replaces about 530 lines — `_reveal_constants`, `_extract_constants` and +`_collect_constant_atoms`, the two consistency guards, the global `_ccode` patching and +the coordinate recovery, the expanded-tree lowering, the opt-in CSE path, the +identity-walk scans that exist for large trees, and the dead `prepare_for_cache_key` and +`_createext` — with about 360: one lowering module, a leaf-to-C map, a manifest built +from the leaves, and node expansion in two unwrappers. Several tier 1 measures become +belt-and-braces rather than load-bearing: the stand-in derivative at every call site, +realness for unit-carrying parameters, and the patch to SymPy's private +`BaseScalar._prop_handler` table, the last of which we would want to remove. These +counts are estimates until the change exists; the PR that makes it shows them. + +## Results agree to round-off, not bit for bit + +The graph changes the order of evaluation. A temporary is rounded once, where the tree +may have folded constants across a sub-expression boundary: $2\,(\tau/(2\dot\varepsilon))$ +is the tree's $\tau/\dot\varepsilon$, but in the graph the inner quotient is a temporary +and the factor of two stays. Kernel outputs therefore differ in the last few bits, and +converged solutions differ at the solver tolerance's round-off. + +### The graph does not cancel across a name + +SymPy cancels on construction: $x \cdot x^{-1}$ is $1$ as soon as both factors meet in +one product. In the tree they meet wherever the law multiplies a named quantity by its +own reciprocal through another name — $A = 1/x$, then $B = A\,x\,T$ is the tree's $T$. +In the graph $A$ is a temporary, $B$ is $t_A\,x\,T$, and at $x = 0$ that is +$\infty \cdot 0 = \text{NaN}$ where the tree gave the limit. The law is not singular +there; the cancellation was simply lost. In a Newton source the sqrt guard keeps the +usual reciprocal, $1/\dot\varepsilon_{II}$, finite; the residual and the Picard tangent +are not guarded. The acceptance criteria therefore include every residual and Picard +kernel at a state of rest, and the constitutive models are read for a quantity +multiplied by a name that holds its reciprocal. + +### Acceptance + +- every tier A and tier B test passes with its current tolerance — none is loosened; +- on every fixture, the assembled residual and Jacobian agree with the expanded route + to $10^{-13}$ relative at a cold state, a converged state and a perturbed state, and + every residual, Picard and Newton kernel is finite and equal to the tree's at a state + of rest; +- a Newton solve of the notch converges under both routes with the same SNES reasons, + and nonlinear and linear iteration counts on the fixtures are unchanged, or each + change is explained; +- the constants manifest contains today's, with every extra slot traced to a + cancellation, and every Jacobian entry that is zero today is zero or evaluates to zero. + +## What does not change + +- The solver code: the explicit Jacobian loops, the PETSc block layout + (`petsc-jacobian-layout.md`), `JITCallbackSet` and the `getext()` signature. +- The `constants[]` contract: a manifested constant ramps without recompiling. +- The cache protocol: memory, disk, rank-0 compile behind a collective decision. +- The residual is never guarded, and the Picard tangent is frozen exactly as now. +- The slot dictionaries `getext()` returns are keyed by the callback objects the solver + passed in, so the solvers' lookups (`ext_dict.jac[self._uu_G3]`) are unaffected. +- `describe()`, the transcript and the model fingerprints record the named, unexpanded + forms (`str(template.sym)`) and the packed constants; none reads the generated C. +- A model's `flux_jacobian` is still differentiated as given. + +## Risks, and the test that closes each + +| risk | how it would fail | closed by | +|---|---|---| +| a dependency missing from a node's arguments | its derivative is silently zero: a Picard-like tangent | arguments from a complete leaf walk, base scalars included; assembled Jacobian against the expanded route on every fixture; a finite-difference oracle with an $h$ sweep (`test_1066`) | +| a slot derivative taken as a total derivative | a field differentiated twice through a coordinate | body partials against independent dummies; a unit test with a field and a coordinate in one node | +| derivative keyed by leaf instead of slot | zero derivative | slot keying; a unit test through `sympy.diff`, where the `Dummy` substitution happens | +| two constants with one display name merged | one material gets the other's coefficient | nodes identified by body equality, which distinguishes them; `test_0103`'s two-material case | +| a node class from an earlier compile reused | a changed `.sym` is ignored | the compile serial in each class's identity; a test that changes a body and rebuilds in one process without clearing SymPy's cache | +| source that depends on the hash seed, a creation counter or a re-declaration | every parallel run repairs; a new process or a re-run notebook cell misses the cache | canonical emission; `test_0105` extended with a Newton fixture, a preamble of extra objects and a re-declared law | +| a cancellation across a name lost | NaN at a state where the tree gave a finite limit, in an unguarded residual or Picard kernel | every residual and Picard kernel compared with the tree at a state of rest on every fixture | +| `test_0022` (tier 1) asserts the expanded, guarded tree of `_jacobian_unwrap`, and checks a block for `Derivative` with `has()`, which does not see node bodies | one test fails by construction, the other cannot fail | both rewritten against expanded nodes in the change that alters `_jacobian_unwrap` | +| an entry that is zero today becomes a non-zero expression | assembly of an entry that evaluates to zero | numbers and leaves inlined; zero patterns compared on every fixture | +| guard placement changes cold-start behaviour | NaN at a state of rest | `test_1067`, extended to a power-law viscosity (the prototype finds both routes finite and equal) | +| `getext()` or another walker expands nodes back into the tree | the cost returns, silently | `getext()` lowers atoms itself; a size check on the emitted source of the notch fixture | +| an atom whose `.sym` is a matrix | it cannot be a scalar temporary | such atoms are expanded in place, as now | +| verbose-output assertions in `test_0004` | a test fails on wording, not on a defect | rewrite those assertions against the kernel contract | +| overhead on small kernels | constant-viscosity problems get slower to set up | measured on the small fixtures; budget: no slower than today | +| a law singular at a state a model reaches (a zero denominator) | NaN where the tree gave rounding noise | the Newton solve of the acceptance criteria; such a law is the model's to fix | + +## Benchmark plan + +During development both routes run in one process on identical inputs, selected by a +private switch, so every comparison is A against B on one build. The switch and the +expanded route are removed before the change merges. The fixtures are built in +`scripts/sessions/jit_graph/` so that the measurements can be repeated from the +repository; the notch is the one exception, and its driver is named. + +**Fixtures** + +1. The notch, Newton: default soft-min, power-mean soft-min, and a declared + `plastic_rate_strengthening`; refinement 1 and 3. +2. Poisson, linear and with $k(u)$; Stokes, viscous and power-law; viscoplastic hard-min + under Picard and Newton, with the sqrt and power-mean smoothers and the continuation + blend — the nine cases tier 1 checked for identical source. +3. A Newton source that differs from the residual (`set_jacobian_F1_source`). +4. SolCx with Nitsche free-slip, whose boundary terms carry the full stress. +5. Transversely isotropic power-law Stokes with rotated free-slip (the #752 fixture). +6. Visco-elasto-plastic Stokes with a stress history (`scripts/sessions/profile_jit_phases.py`). +7. A three-dimensional viscoplastic Stokes, where `G3` has 81 entries. +8. Small kernels: constant-viscosity Stokes, Poisson, `Integral` and `BdIntegral`. + +**Measurements** + +- setup time by phase (`uw.timing`): lowering, differentiation, constants, emission, + compile — with `UW_NO_USAGE_METRICS=1`, since the import-time usage report makes an + HTTP request on a thread that runs during setup; +- source size, `.so` size, and peak compiler memory; +- kernel run time: residual and Jacobian assembly (`SNESFunctionEval`, + `SNESJacobianEval`) repeated on a fixed state, on macOS and on Linux, where gcc's + default `-fmath-errno` keeps it from merging repeated `pow`, `exp` and `log` calls; +- agreement: compiled kernel outputs at random states (fraction bit-identical, largest + difference relative to the block); assembled residual and Jacobian at the three + states; converged solutions; SNES reasons and iteration counts; manifest and zero + patterns; +- determinism: generated source under `PYTHONHASHSEED` 0, 1 and 2 on every fixture; +- parallel: source-hash agreement in ten runs each at np = 3 and np = 4 (within the + eight-core budget), counting repairs, and `tests/parallel/ptest_jit_cache.py` run by + hand at np = 4 — it matches no CI glob, so CI has never run it; +- cache: ramping a constant leaves the key unchanged, a second process hits the disk + cache, and so does a second process whose script creates other objects first. + +## Staging + +1. Tier 1 merges. +2. `getext()` emits every callback from the graph and stops unwrapping whole kernels. + Residual and Picard kernels gain the temporaries; Newton sources still arrive + expanded and print as before. Nodes exist only inside `getext()`. +3. `_jacobian_unwrap` builds nodes instead of expanding, so nodes reach the solver's + blocks; in the same change the two unwrappers learn to expand them and `test_0022`'s + two tree-shaped tests are rewritten. +4. The expanded route and the switch are removed, together with the two functions that + mirror it and have no callers (`prepare_for_cache_key`, `_createext`), and + `jit-cache.md`, `expressions-functions.md` and `jacobian-consistent-tangent.md` are + updated. `jit-cache.md` already describes the cross-rank check as an abort and every + rank as compiling; both stopped being true before this change. + +Each step is benchmarked against the one before it and reviewed adversarially before +the next begins. + +## CSE on the tree, call-site chain rules and Symbol nodes were rejected + +- **Common sub-expression elimination on the expanded tree** (`UW_JIT_CSE=1`, opt-in + today). It shrinks the C, but only after the tree has been built, differentiated and + scanned, and SymPy's `cse` on a tree of $10^5$ nodes is itself slow. +- **A chain-rule function called at each Jacobian site.** Tier 1 has since put one + function at every site, for realness, and a chain rule could live there. It would + still leave the derivative's correctness to the call: any derivative taken another + way (a test's oracle, an adjoint, a user's `sympy.diff`) would see frozen nodes. With + `fdiff` on the node, every route gives the same answer. +- **A `Symbol` subclass with its own `_eval_derivative`.** SymPy returns zero before + calling it. + +## Decisions + +- **2026-10-07: tier 2 may change the generated C.** Tier 1's speed work keeps the + generated C byte-identical, because bit-identical kernels are what the JIT cache rests + on. Tier 2 redesigns the JIT, and that constraint does not bind it: every kernel + containing a named non-constant quantity is evaluated in a new order and agrees with + tier 1 to round-off, not bit for bit. A law with no such quantity emits byte-identical + C. After the change the C is canonical, bit-identical across hash seeds, preambles, + re-declarations and processes by construction. + +## Questions for the maintainer + +1. One PR for steps 2–4, or one PR per step. +2. Whether the generated C should carry each temporary's display name as a comment. + It makes kernels readable; it also puts display names into the cache key, so a + `rename()` recompiles. +3. Whether the adjoint's own unwrap (`_peel_except`, on the adjoint branches) adopts the + nodes in this change or after it. +4. Whether `UW_JIT_CSE` is retired once this lands. diff --git a/docs/developer/index.md b/docs/developer/index.md index 3424240c2..1a7793f87 100644 --- a/docs/developer/index.md +++ b/docs/developer/index.md @@ -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} diff --git a/scripts/sessions/jit_graph/fixtures.py b/scripts/sessions/jit_graph/fixtures.py new file mode 100644 index 000000000..e77c51fc7 --- /dev/null +++ b/scripts/sessions/jit_graph/fixtures.py @@ -0,0 +1,105 @@ +"""Stokes fixtures for the JIT measurements of #823 (tier 1 and tier 2). + +Each builder returns ``(stokes, admissible)``: the solver with its constitutive law set +and the Newton tangent selected, and the bounds of any field that is bounded by +construction (used when random states are drawn). +""" +import os +import sys + +import sympy +import underworld3 as uw + + +def _box_stokes(): + mesh = uw.meshing.UnstructuredSimplexBox(cellSize=0.25) + v = uw.discretisation.MeshVariable("V", mesh, 2, degree=2) + p = uw.discretisation.MeshVariable("P", mesh, 1, degree=1) + stokes = uw.systems.Stokes(mesh, velocityField=v, pressureField=p) + stokes.bodyforce = sympy.Matrix([[0.0, -1.0]]) + return mesh, p, stokes + + +def build_box(): + """ViscoPlastic: Drucker-Prager yield C + mu p, temperature-dependent viscosity, + yield-stress and viscosity floors.""" + mesh, p, stokes = _box_stokes() + T = uw.discretisation.MeshVariable("T", mesh, 1, degree=1) + stokes.constitutive_model = uw.constitutive_models.ViscoPlasticFlowModel + P = stokes.constitutive_model.Parameters + P.shear_viscosity_0 = uw.expression(r"\eta_0", 1.0) * sympy.exp( + -uw.expression(r"\theta", 3.0) * T.sym[0]) + P.yield_stress = uw.expression(r"C", 0.5) + uw.expression(r"\mu", 0.6) * p.sym[0] + P.yield_stress_min = uw.expression(r"\tau_{\min}", 0.01) + P.shear_viscosity_min = uw.expression(r"\eta_{\min}", 1.0e-3) + return stokes, {"T": (0.0, 1.0)} + + +def build_powerlaw(numeric_n=False): + """A power-law viscosity on a NAMED strain-rate invariant, no viscosity cap: + eta = eta_0 (edot_II / edot_ref)^(1/n - 1), n = 3, either as a constant atom or as + the number 3.""" + mesh, p, stokes = _box_stokes() + stokes.constitutive_model = uw.constitutive_models.ViscousFlowModel + edot = uw.expression(r"\dot\varepsilon_{II}", stokes.Unknowns.Einv2, "strain-rate invariant") + n = sympy.Integer(3) if numeric_n else uw.expression(r"n", 3, "stress exponent") + stokes.constitutive_model.Parameters.shear_viscosity_0 = ( + uw.expression(r"\eta_0", 1.0) * (edot / uw.expression(r"\dot\varepsilon_0", 1.0)) ** (1 / n - 1)) + return stokes, {} + + +def build_linear(): + """Constant-viscosity Stokes: the small kernels on which any overhead shows.""" + mesh, p, stokes = _box_stokes() + stokes.constitutive_model = uw.constitutive_models.ViscousFlowModel + stokes.constitutive_model.Parameters.shear_viscosity_0 = uw.expression(r"\eta", 1.0) + return stokes, {} + + +def build_vep(): + """Visco-elasto-plastic, order 1, with a Drucker-Prager yield and floors.""" + mesh, p, stokes = _box_stokes() + cm = uw.constitutive_models.ViscoElasticPlasticFlowModel(stokes.Unknowns, order=1) + stokes.constitutive_model = cm + cm.Parameters.shear_viscosity_0 = 1.0 + cm.Parameters.shear_modulus = 1.0 + cm.Parameters.dt_elastic = sympy.Rational(1, 10) + cm.Parameters.yield_stress = uw.expression(r"C", 0.5) + uw.expression(r"\mu", 0.3) * p.sym[0] + cm.Parameters.yield_stress_min = uw.expression(r"\tau_{\min}", 0.01) + return stokes, {} + + +def build_ti(): + """Transversely isotropic VEP with a yield on the weak direction (test_1066's law).""" + mesh, p, stokes = _box_stokes() + cm = uw.constitutive_models.TransverseIsotropicVEPFlowModel(stokes.Unknowns) + stokes.constitutive_model = cm + cm.Parameters.shear_viscosity_0 = 1.0 + cm.Parameters.shear_viscosity_1 = 0.01 + cm.Parameters.shear_modulus = 100.0 + cm.Parameters.dt_elastic = 0.1 + cm.Parameters.yield_stress = 0.5 + cm.Parameters.director = sympy.Matrix([0.0, 1.0]) + cm.Parameters.strainrate_inv_II_min = 1.0e-6 + return stokes, {} + + +def build_notch(): + """The Spiegelman notch campaign law (power-mean soft minimum).""" + sys.path.insert(0, os.path.expanduser("~/+Simulations/spiegelman_hardcase/drivers")) + import notch_model as nm + S = nm.build(os.path.expanduser("~/+Simulations/spiegelman_hardcase/meshes/notch_mesh1.msh"), + 1.0e24, 2.5, 30.0, 1, xi=0.0, seed=False, p_degree=0, p_continuous=False, + floor_scalar=1.0e-3, unique_params=True) + # mat is P0 holding 0 or 1; xicap is a non-negative P0 cap + return S["stokes"], {"mat": (0.0, 1.0), "xicap": (0.0, 3.0)} + + +BUILDERS = {"box": build_box, "powerlaw": build_powerlaw, "linear": build_linear, + "vep": build_vep, "ti": build_ti, "notch": build_notch} + + +def build(name, **kwargs): + stokes, admissible = BUILDERS[name](**kwargs) + stokes.consistent_jacobian = True + return stokes, admissible diff --git a/scripts/sessions/jit_graph/graph_size_probe.py b/scripts/sessions/jit_graph/graph_size_probe.py new file mode 100644 index 000000000..123f66eda --- /dev/null +++ b/scripts/sessions/jit_graph/graph_size_probe.py @@ -0,0 +1,66 @@ +"""Size of a kernel's shared graph of named sub-expressions against the tree the +current JIT expands it into (#823). Builds the campaign notch model, does NOT solve, +and never expands: the expanded size is counted by dynamic programming over the graph. +""" +import os, sys, time +import sympy +import underworld3 as uw +from underworld3.function.expressions import UWexpression, _unwrap_atom +from underworld3.utilities._jitextension import _is_truly_constant + +params = uw.Params( + uw_mesh=uw.Param(os.path.expanduser("~/+Simulations/spiegelman_hardcase/meshes/notch_mesh1.msh")), +) +sys.path.insert(0, os.path.expanduser("~/+Simulations/spiegelman_hardcase/drivers")) +import notch_model as nm + +S = nm.build(str(params.uw_mesh), 1.0e24, 2.5, 30.0, 1, xi=0.0, seed=False, p_degree=0, + p_continuous=False, floor_scalar=1.0e-3, unique_params=True) +st, cm = S["stokes"], S["cm"] + +constant = {} +def is_const(a): + if id(a) not in constant: + constant[id(a)] = (a, _is_truly_constant(a, UWexpression)) + return constant[id(a)][1] + +named, consts = {}, {} +def body(a): + return _unwrap_atom(a, "symbolic") + +tree_size, seen_nodes = {}, set() +def size(e): + """Expanded-tree node count of e, with non-constant UW atoms expanded.""" + k = id(e) + if k in tree_size: + return tree_size[k][1] + seen_nodes.add(k) + if isinstance(e, UWexpression): + if is_const(e): + consts[id(e)] = e + n = 1 + else: + named[id(e)] = e + n = size(body(e)) + elif isinstance(e, (sympy.MatrixBase, sympy.NDimArray)): + n = sum(size(x) for x in e) + elif isinstance(e, sympy.Basic) and e.args: + n = 1 + sum(size(a) for a in e.args) + else: + n = 1 + tree_size[k] = (e, n) + return n + +t0 = time.perf_counter() +F1 = st.F1.sym +eta = cm.viscosity +n_eta, n_F1 = size(eta), size(F1) +graph_nodes = len(seen_nodes) +print(f"[probe] {time.perf_counter()-t0:.2f} s") +print(f"named non-constant sub-expressions : {len(named)}") +print(f"constants[] atoms reached : {len(consts)}") +print(f"distinct sympy nodes in the graph : {graph_nodes}") +print(f"expanded tree, viscosity : {n_eta}") +print(f"expanded tree, F1 (all entries) : {n_F1}") +for e in sorted(named.values(), key=lambda a: -tree_size[id(a)][1])[:12]: + print(f" {tree_size[id(e)][1]:>8d} {e.name}") diff --git a/scripts/sessions/jit_graph/graph_vs_library.py b/scripts/sessions/jit_graph/graph_vs_library.py new file mode 100644 index 000000000..496bb6e40 --- /dev/null +++ b/scripts/sessions/jit_graph/graph_vs_library.py @@ -0,0 +1,261 @@ +"""Kernels built two ways, compiled, and compared (tier 2 of #823). + + library: today's route on tier 1 -- the Newton source from the solver's own + ``_jacobian_source``, residual and Picard kernels unwrapped as ``getext`` + unwraps them (keep-constants), derivatives through ``diff_wrt_field`` with the + solver's explicit uu_G3 loops, each output printed whole; + graph: ``kernel_graph`` nodes, the same loops, one temporary per node. + +Three kernels per fixture: the residual flux F1 (plain nodes), the Picard uu_G3 (atoms +frozen while differentiating, plain nodes after) and the Newton uu_G3 (guarded nodes). +Each pair is compiled into C functions of the same leaves and evaluated at random +states and at a state of rest. + +Determinism checks: run under several PYTHONHASHSEED values, with ``-uw_preamble N`` +(N throwaway objects created first, shifting every creation counter) and with +``-uw_redeclare 1`` (the law built twice in one process, as a re-run notebook cell +does); the printed md5s must not change. + +Fixtures are in fixtures.py: box, powerlaw, linear, vep and ti are built there and +reproducible from the repository; notch is the Spiegelman campaign law +(~/+Simulations/spiegelman_hardcase/drivers/notch_model.py). +""" +import ctypes +import hashlib +import os +import subprocess +import sys +import tempfile +import time + +import numpy as np +import sympy +from sympy.core.function import AppliedUndef +from sympy.printing.c import c_code_printers +from sympy.vector.scalar import BaseScalar + +import underworld3 as uw +from underworld3.function.expressions import UWexpression, unwrap_expression +from underworld3.function import diff_wrt_field +from underworld3.utilities._jitextension import _extract_constants, _pack_constants + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +from kernel_graph import KernelGraph, _KernelNode, emission_order, guard # noqa: E402 +import fixtures # noqa: E402 + +params = uw.Params( + uw_fixture=uw.Param("box", description="box | powerlaw | linear | vep | ti | notch"), + uw_states=uw.Param(2000, description="random states for the value comparison"), + uw_calls=uw.Param(200000, description="kernel calls for the run-time measurement"), + uw_preamble=uw.Param(0, description="throwaway expressions and one mesh created before the law"), + uw_redeclare=uw.Param(0, description="1 = build the law twice and use the second"), + uw_numeric_n=uw.Param(0, description="powerlaw: 1 = the stress exponent as a plain number, not a constant atom"), +) +fixture = str(params.uw_fixture) + + +for i in range(int(params.uw_preamble)): + uw.expression(f"preamble_{i}", float(i)) +if int(params.uw_preamble): + uw.meshing.UnstructuredSimplexBox(cellSize=0.5) +kwargs = {"numeric_n": bool(int(params.uw_numeric_n))} if fixture == "powerlaw" else {} +stokes, admissible = fixtures.build(fixture, **kwargs) +if int(params.uw_redeclare): + stokes, admissible = fixtures.build(fixture, **kwargs) +mesh = stokes.mesh +dim = mesh.dim +L = stokes.Unknowns.L +F1 = sympy.Array(stokes.F1.sym).reshape(dim, dim) + + +def g3(F): + """The solver's explicit uu_G3 loops, differentiating through diff_wrt_field.""" + F = sympy.Array(F).reshape(dim, dim) + G = sympy.zeros(dim * dim, dim * dim) + for gc in range(dim): + for dg in range(dim): + dF = diff_wrt_field(F, L[gc, dg]) + for fc in range(dim): + for df in range(dim): + G[fc * dim + gc, df * dim + dg] = dF[fc, df] + return list(G) + + +def keep_constants(e): + return unwrap_expression(e, mode="symbolic_keep_constants") + + +timing = {} +def timed(label, fn): + t = time.perf_counter() + out = fn() + timing[label] = time.perf_counter() - t + return out + + +F1_flat = [F1[i, j] for i in range(dim) for j in range(dim)] +graph = KernelGraph() +kernels = {} +kernels["residual F1"] = ( + timed("library residual: unwrap", lambda: [keep_constants(e) for e in F1_flat]), + timed("graph residual: lower", lambda: [graph.lower(e, guarded=False) for e in F1_flat])) +picard_raw = timed("both Picard: differentiate (atoms frozen)", lambda: g3(F1)) +kernels["Picard G3"] = ( + timed("library Picard: unwrap", lambda: [keep_constants(e) for e in picard_raw]), + timed("graph Picard: lower", lambda: [graph.lower(e, guarded=False) for e in picard_raw])) +F1_lib = timed("library Newton: source (unwrap + guard)", + lambda: stokes._jacobian_source(F1, stokes._newton_flux(F1))) +G3_lib = timed("library Newton: differentiate", lambda: g3(F1_lib)) +F1_graph = timed("graph Newton: lower", + lambda: sympy.Array([guard(graph.lower(e, guarded=True)) for e in F1_flat], + (dim, dim))) +G3_graph = timed("graph Newton: differentiate", lambda: g3(F1_graph)) +kernels["Newton G3"] = (G3_lib, G3_graph) + +# ------------------------------------------------------------------ spelling of leaves +def leaves_of(exprs): + out = set() + for x in exprs: + x = sympy.sympify(x) + out |= {a for a in x.atoms(AppliedUndef) if not isinstance(a, _KernelNode)} + out |= {s for s in x.free_symbols if isinstance(s, (UWexpression, BaseScalar))} + return out + +graph_bodies = [b for _, gr in kernels.values() for _, b in emission_order(gr, str)[0]] +graph_leaves = leaves_of([x for _, gr in kernels.values() for x in gr] + graph_bodies) +library_leaves = leaves_of([x for lib, _ in kernels.values() for x in lib]) +leaves = graph_leaves | library_leaves +fields = sorted((s for s in leaves if isinstance(s, AppliedUndef)), key=KernelGraph.display_key) +slot = {f: k for k, f in enumerate(fields)} +library_consts = {e for _, e in _extract_constants( + tuple(sympy.ImmutableMatrix([list(lib)]) for lib, _ in kernels.values()), mesh)[0]} +graph_consts = {s for s in graph_leaves if isinstance(s, UWexpression)} +consts = sorted(library_consts | graph_consts, key=lambda e: (e.name, e.instance_number)) +cindex = {c: i for i, c in enumerate(consts)} +cvals = np.asarray(_pack_constants(list(enumerate(consts))), dtype=float) + + +def spell(s): + if s in slot: + return f"petsc_u[{slot[s]}]" + if s in cindex: + return f"constants[{cindex[s]}]" + if isinstance(s, BaseScalar): + return f"petsc_x[{s._id[0]}]" + return str(s) + + +spelled = {s: sympy.Symbol(spell(s)) for s in leaves} +printer = c_code_printers["c99"]({}) +def c_of(e, temps=None): + e = sympy.sympify(e) + if temps: + e = e.xreplace(temps) + return printer.doprint(e.xreplace(spelled)) + +# ------------------------------------------------------------------ compile and compare +SIG = "void k(const double *petsc_u, const double *petsc_x, const double *constants, double *out)" +HDR = ("#include \n" + "static inline double Heaviside_1(double x){return x<0?0:x>0?1:0.5;}\n") +work = tempfile.mkdtemp(dir=os.path.expanduser("~/+Simulations"), prefix=f"jit_graph_{fixture}_") +P = ctypes.POINTER(ctypes.c_double) +md5 = lambda s: hashlib.md5(s.encode()).hexdigest()[:10] + + +def compile_c(tag, body): + cfile = os.path.join(work, f"{tag}.c") + open(cfile, "w").write(HDR + SIG + " {\n" + body.replace("Heaviside(", "Heaviside_1(") + "\n}\n") + so = os.path.join(work, f"{tag}.so") + t = time.perf_counter() + subprocess.run(["cc", "-std=c99", "-O3", "-g0", "-shared", "-fPIC", cfile, "-o", so], check=True) + lib = ctypes.CDLL(so) + lib.k.restype = None + return lib, time.perf_counter() - t, os.path.getsize(so) + + +empty, _, _ = compile_c("empty", "") +print(f"[{fixture}] seed {os.environ.get('PYTHONHASHSEED', 'random')} preamble {params.uw_preamble} " + f"redeclare {params.uw_redeclare}") +for k, v in timing.items(): + print(f" {k:44s} {v:9.2f} s") +extra = sorted(c.name for c in graph_consts - library_consts) +missing = sorted(c.name for c in library_consts - graph_consts) +print(f"constants: library {len(library_consts)}, graph {len(graph_consts)}; " + f"graph has the library's: {not missing}; extra in graph: {extra}") + +rng = np.random.default_rng(7) +states = [] +Lfields = {f for f in fields if f in set(L)} +for _ in range(int(params.uw_states)): + scale = 10.0 ** rng.uniform(-6, 2) + u = rng.normal(size=len(fields)) + for f, k in slot.items(): + if f in Lfields: + u[k] *= scale + for name, (lo, hi) in admissible.items(): + if f.func.__name__.strip("{}") == name: + u[k] = rng.uniform(lo, hi) + states.append((u, rng.uniform(0, 1, size=3))) +rest = (np.zeros(len(fields)), np.full(3, 0.5)) + +for name, (lib_exprs, graph_exprs) in kernels.items(): + t = time.perf_counter() + order, key = emission_order(graph_exprs, spell) + temps_by_key = {k: sympy.Symbol(f"t{i}", real=True) for i, (k, _) in enumerate(order)} + temps = {app: temps_by_key[k] for app, k in key.items()} + src_graph = "\n".join([f"const double {temps_by_key[k]} = {c_of(body, temps)};" for k, body in order] + + [f"out[{i}] = {c_of(x, temps)};" for i, x in enumerate(graph_exprs)]) + t_emit = time.perf_counter() - t + t = time.perf_counter() + src_lib = "\n".join(f"out[{i}] = {c_of(x)};" for i, x in enumerate(lib_exprs)) + t_print = time.perf_counter() - t + tag = name.replace(" ", "_") + lib_g, cc_g, so_g = compile_c(f"{tag}_graph", src_graph) + lib_l, cc_l, so_l = compile_c(f"{tag}_library", src_lib) + nout = len(graph_exprs) + + def call(lib, u, x): + out = np.zeros(nout) + lib.k(u.ctypes.data_as(P), x.ctypes.data_as(P), cvals.ctypes.data_as(P), out.ctypes.data_as(P)) + return out + + nonzero = exact = nan_mismatch = nan_both = lib_zero_graph_not = 0 + worst_block = 0.0 + for u, x in states: + a, b = call(lib_g, u, x), call(lib_l, u, x) + nan_mismatch += int(np.any(np.isnan(a) != np.isnan(b))) + nan_both += int(np.sum(np.isnan(a) & np.isnan(b))) + lib_zero_graph_not += int(np.sum((a != 0) & (b == 0))) + m = (b != 0) & np.isfinite(b) & np.isfinite(a) + nonzero += m.sum(); exact += (a[m] == b[m]).sum() + if m.any(): + worst_block = max(worst_block, float(np.max(np.abs(a[m] - b[m])) / np.max(np.abs(b[m])))) + a0, b0 = call(lib_g, *rest), call(lib_l, *rest) + zl = sum(sympy.sympify(x) == 0 for x in lib_exprs) + zg = sum(sympy.sympify(x) == 0 for x in graph_exprs) + + per = {} + u, x = states[0] + for tg, lib in (("graph", lib_g), ("library", lib_l), ("empty", empty)): + out = np.zeros(max(nout, 1)) + args = (u.ctypes.data_as(P), x.ctypes.data_as(P), cvals.ctypes.data_as(P), out.ctypes.data_as(P)) + n = int(params.uw_calls) + t = time.perf_counter() + for _ in range(n): + lib.k(*args) + per[tg] = 1e9 * (time.perf_counter() - t) / n + + print(f"--- {name}: md5 graph {md5(src_graph)} library {md5(src_lib)}") + print(f" emit {t_emit:.2f} s, print {t_print:.2f} s; C graph {len(src_graph):,} B " + f"({len(order)} temporaries), library {len(src_lib):,} B; cc -O3 {cc_g:.2f} s / {cc_l:.2f} s; " + f".so {so_g:,} / {so_l:,} B") + print(f" {len(states)} states: {nonzero} non-zero entries, bit-identical {exact} " + f"({100 * exact / max(nonzero, 1):.0f}%), max |diff|/max|block| {worst_block:.1e}, " + f"NaN in one route {nan_mismatch}, NaN in both {nan_both}, " + f"graph non-zero where library exactly zero {lib_zero_graph_not}") + print(f" structural zeros: library {zl}, graph {zg}; state of rest: finite graph " + f"{bool(np.isfinite(a0).all())} library {bool(np.isfinite(b0).all())}, " + f"elementwise equal {bool(np.array_equal(a0, b0))}") + print(f" one call: graph {per['graph'] - per['empty']:.1f} ns, library " + f"{per['library'] - per['empty']:.1f} ns (empty-kernel ctypes call {per['empty']:.1f} ns subtracted)") +print("work dir:", work) diff --git a/scripts/sessions/jit_graph/kernel_graph.py b/scripts/sessions/jit_graph/kernel_graph.py new file mode 100644 index 000000000..92d9a7be9 --- /dev/null +++ b/scripts/sessions/jit_graph/kernel_graph.py @@ -0,0 +1,278 @@ +"""Prototype of the shared-graph lowering for JIT kernels (tier 2 of #823). + +Not library code: the design note it supports is +docs/developer/design/jit-shared-graph-codegen.md. + +A non-constant UWexpression atom becomes a NODE: an applied undefined function of the +leaves its value depends on. SymPy's own chain rule then differentiates through it, +because ``fdiff(i)`` returns the partial derivative with respect to argument slot i. + +Two kinds of identity are kept apart: + +- In Python, a node is identified by its body. Two atoms with equal bodies share a + node; bodies that use different constants are different even when the constants + share a display name, because SymPy equality sees the difference. The class name is + a hash of the body written with DISPLAY identities only (no creation counters), and + each class carries its graph's serial in its SymPy identity (``_ctx``), so classes + from two compiles never compare equal. +- In C, a temporary is identified by a hash of the C it computes: its body with every + leaf written as the C the kernel reads (``petsc_u[3]``, ``constants[2]``) and every + child as the child's hash. Temporaries are ordered and merged by that hash, so the + source is a function of the mathematics and the kernel's data layout, not of Python + names, creation counters or hash seeds. + +Partial derivatives are taken with every leaf replaced by an independent real dummy. +A derivative that is zero, a number or a single leaf is returned as itself. +""" +import hashlib +import itertools + +import sympy +from sympy.core.function import AppliedUndef, UndefinedFunction +from sympy.vector.scalar import BaseScalar + +from underworld3.function.expressions import UWexpression +from underworld3.function._function import UnderworldAppliedFunction +from underworld3.utilities._jitextension import _is_truly_constant + +EPS2 = sympy.Float(1.0e-36) +_graph_serial = itertools.count() + + +def guard(e): + """The sqrt guard of ``_jacobian_unwrap``: every half-integer power whose base has + free symbols (constants included) gets ``+1e-36`` in its base. Memoised on identity + (as tier 1 does); node applications are not entered.""" + memo = {} + + def g(n): + hit = memo.get(id(n)) + if hit is not None: + return hit[1] + out = n + args = getattr(n, "args", None) + if args and not isinstance(n, _KernelNode): + new_args = tuple(g(a) for a in args) + if any(a is not b for a, b in zip(args, new_args)) and args != new_args: + out = n.func(*new_args) + if (out.is_Pow and out.exp.is_Rational and out.exp.q == 2 + and out.args[0].free_symbols): + out = sympy.Pow(out.args[0] + EPS2, out.exp) + memo[id(n)] = (n, out) + return out + + return g(e) + + +class _KernelNode(AppliedUndef): + """A named sub-expression applied to the leaves its value depends on.""" + + def fdiff(self, argindex=1): + return self._graph.slot_derivative(self, argindex - 1) + + +def body_of(app): + """A node application's body, with the application's arguments substituted.""" + cls = app.func + if app.args == cls._deps: + return cls._body + return cls._body.xreplace(dict(zip(cls._deps, app.args))) + + +class KernelGraph: + """One lowering context: the nodes of one compile.""" + + def __init__(self): + self.serial = next(_graph_serial) + self._const = {} # id(atom) -> (atom, bool) + self._node = {} # (id(atom), guarded) -> (atom, application) + self._by_body = {} # body -> class + self._deriv = {} # (class, slot) -> derivative in terms of the class's deps + self._busy = set() + self._splitting = False + + # ------------------------------------------------------------------ leaves + def is_constant(self, atom): + hit = self._const.get(id(atom)) + if hit is None: + hit = self._const[id(atom)] = (atom, _is_truly_constant(atom, UWexpression)) + return hit[1] + + def is_leaf(self, s): + if isinstance(s, _KernelNode): + return False + if isinstance(s, (AppliedUndef, BaseScalar)): + return True + if isinstance(s, UWexpression): + return self.is_constant(s) + return isinstance(s, sympy.Symbol) + + @staticmethod + def display_key(s): + """What a leaf is called, without creation counters: stable across processes, + ranks, preambles and re-declarations. Not unique (two constants may share a + display name); used only to name and order, never to identify.""" + if isinstance(s, UnderworldAppliedFunction): + return f"F|{s.func.__name__}" + if isinstance(s, BaseScalar): + return f"X|{s._id[1]}|{s._id[0]}" + if isinstance(s, UWexpression): + return f"C|{s.name}" + if isinstance(s, AppliedUndef): + return f"A|{s.func.__name__}|{','.join(map(str, s.args))}" + return f"S|{s.name}" + + @staticmethod + def order_key(s): + """Display key, ties broken by RELATIVE creation order (as the constants + manifest orders same-named constants).""" + return (KernelGraph.display_key(s), getattr(s, "instance_number", 0)) + + def leaves(self, e): + out, seen, stack = set(), set(), [e] + while stack: + a = stack.pop() + if id(a) in seen: + continue + seen.add(id(a)) + if isinstance(a, _KernelNode): + out.update(a.args) + elif self.is_leaf(a): + out.add(a) + elif isinstance(a, sympy.Basic): + stack.extend(a.args) + return out + + # ------------------------------------------------------------------ nodes + def lower(self, e, guarded): + """``e`` with each non-constant UW atom replaced by its node.""" + if not isinstance(e, sympy.Basic): + return e + atoms = sorted((s for s in e.free_symbols + if isinstance(s, UWexpression) and not self.is_constant(s)), + key=self.order_key) + m = {s: self.node_of(s, guarded) for s in atoms} + # a coordinate atom lowers to the base scalar the kernel reads, as the tree does + m.update({s: s.sym for s in e.free_symbols + if isinstance(s, BaseScalar) and type(s).__name__ == "UWCoordinate"}) + return e.xreplace(m) if m else e + + def node_of(self, atom, guarded): + key = (id(atom), guarded) + hit = self._node.get(key) + if hit is not None: + return hit[1] + if key in self._busy: + raise ValueError(f"cyclic expression: {atom} contains itself") + self._busy.add(key) + try: + body = self.lower(atom.sym, guarded) + if guarded: + body = guard(body) + app = self.make_node(body) + finally: + self._busy.discard(key) + self._node[key] = (atom, app) + return app + + def _display_name(self, body): + canon = {a: sympy.Symbol(a.func.__name__) for a in body.atoms(_KernelNode)} + for leaf in self.leaves(body): + canon.setdefault(leaf, sympy.Symbol(self.display_key(leaf))) + text = sympy.srepr(body.xreplace(canon)) + return "N" + hashlib.sha1(text.encode()).hexdigest()[:12] + + def make_node(self, body): + """The node for ``body``: one per distinct body. A body that is a number, a leaf + or a single node needs no temporary and is returned as itself.""" + body = sympy.sympify(body) + if not self._splitting: + body = self._split_shared(body) + if body.is_Atom or isinstance(body, (AppliedUndef, BaseScalar)): + return body + cls = self._by_body.get(body) + if cls is None: + deps = tuple(sorted(self.leaves(body), key=self.order_key)) + cls = UndefinedFunction(self._display_name(body), bases=(_KernelNode,), + real=True, _ctx=self.serial, _n=len(self._by_body), + __dict__={"_graph": self}) + cls._body, cls._deps = body, deps + self._by_body[body] = cls + return cls(*cls._deps) + + def _split_shared(self, body): + """Repeated unnamed sub-expressions of one body become anonymous nodes.""" + if body.is_Atom: + return body + repl, (reduced,) = sympy.cse([body], symbols=sympy.numbered_symbols("_cse", real=True), + order="none") + if not repl: + return body + self._splitting = True + try: + m = {} + for sym, e in repl: + m[sym] = self.make_node(e.xreplace(m)) + return reduced.xreplace(m) + finally: + self._splitting = False + + # ------------------------------------------------------------------ derivatives + def slot_derivative(self, app, i): + cls = app.func + key = (cls, i) + if key not in self._deriv: + deps = cls._deps + dummies = [sympy.Dummy(real=True) for _ in deps] + b = cls._body.xreplace(dict(zip(deps, dummies))) + d = sympy.diff(b, dummies[i]).xreplace(dict(zip(dummies, deps))) + self._deriv[key] = self.make_node(d) + d = self._deriv[key] + if app.args == cls._deps: + return d + return d.xreplace(dict(zip(cls._deps, app.args))) + + +# ---------------------------------------------------------------------- emission +def emission_order(outputs, spell): + """The temporaries a kernel needs, in the order to write them. + + ``spell(leaf)`` is the C the kernel reads for a leaf. Each node application gets a + canonical key: a hash of its body with leaves spelled as C and children replaced by + their own keys. Applications with equal keys compute the same C and share one + temporary; temporaries follow the ones they use, ties broken by key. + + Returns ``(order, key)``: ``order`` is a list of ``(key, body)``; ``key`` maps each + node application reached to its key. + """ + key = {} + + def key_of(app): + hit = key.get(app) + if hit is not None: + return hit + body = body_of(app) + sub = {c: sympy.Symbol("@" + key_of(c)) for c in body.atoms(_KernelNode)} + for leaf in body.free_symbols | body.atoms(AppliedUndef): + if not isinstance(leaf, _KernelNode) and leaf not in sub: + sub[leaf] = sympy.Symbol(spell(leaf)) + k = hashlib.sha1(sympy.srepr(body.xreplace(sub)).encode()).hexdigest()[:16] + key[app] = k + return k + + order, placed = [], set() + + def visit(app): + k = key_of(app) + if k in placed: + return + placed.add(k) + body = body_of(app) + for child in sorted(body.atoms(_KernelNode), key=key_of): + visit(child) + order.append((k, body)) + + for out in outputs: + for child in sorted(sympy.sympify(out).atoms(_KernelNode), key=key_of): + visit(child) + return order, key diff --git a/scripts/sessions/jit_graph/library_setup_profile.py b/scripts/sessions/jit_graph/library_setup_profile.py new file mode 100644 index 000000000..d350b993c --- /dev/null +++ b/scripts/sessions/jit_graph/library_setup_profile.py @@ -0,0 +1,86 @@ +"""The library's whole Newton setup on one fixture: wall time by phase, an identity hash +of the generated C, and optionally a cProfile with the hotspots grouped. + +The identity hash is the md5 of the generated header with the module name and symbol +prefix canonicalised, as ``getext`` canonicalises them before it hashes; it is +independent of the Underworld version string, so it compares two builds directly. +""" +import cProfile +import hashlib +import io +import os +import pstats +import sys +import time + +import underworld3 as uw +import underworld3.utilities._jitextension as jx + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +import fixtures # noqa: E402 + +params = uw.Params( + uw_fixture=uw.Param("notch", description="box | powerlaw | linear | vep | ti | notch"), + uw_profile=uw.Param(0, description="1 = cProfile the setup and print the hotspots"), +) +fixture = str(params.uw_fixture) +stokes, _ = fixtures.build(fixture) + +seen = {} +generate, compile_ = jx.generate_c_source, jx.compile_and_load + + +def timed_generate(*a, **k): + t = time.perf_counter() + modname, codeguys, diag = generate(*a, **k) + seen["generate"] = time.perf_counter() - t + header = dict(codeguys)["cy_ext.h"] + header = header.replace(modname, "__MOD__").replace(diag["randstr"], "__RS__") + seen["header md5"] = hashlib.md5(header.encode()).hexdigest()[:10] + seen["header bytes"] = len(header) + return modname, codeguys, diag + + +def timed_compile(*a, **k): + t = time.perf_counter() + out = compile_(*a, **k) + seen["compile"] = time.perf_counter() - t + return out + + +jx.generate_c_source, jx.compile_and_load = timed_generate, timed_compile + +profiler = cProfile.Profile() if int(params.uw_profile) else None +t = time.perf_counter() +if profiler: + profiler.enable() +stokes._setup_pointwise_functions() +if profiler: + profiler.disable() +wall = time.perf_counter() - t + +print(f"[{fixture}] Newton setup {wall:.1f} s{' (under cProfile)' if profiler else ''}; " + f"C generation {seen.get('generate', float('nan')):.1f} s; compile {seen.get('compile', float('nan')):.1f} s; " + f"header {seen.get('header bytes', 0):,} B, md5 {seen.get('header md5')}") + +if profiler: + out = os.path.expanduser(f"~/+Simulations/jit_graph/{fixture}_setup.prof") + profiler.dump_stats(out) + st = pstats.Stats(profiler) + groups = { + "C printing (CodePrinter.doprint)": ("sympy/printing/codeprinter.py", "doprint"), + " term ordering (printer._as_ordered_terms)": ("sympy/printing/printer.py", "_as_ordered_terms"), + " _handle_UnevaluatedExpr": ("sympy/printing/codeprinter.py", "_handle_UnevaluatedExpr"), + "differentiation (_dispatch_eval_derivative_n_times)": ("sympy/core/function.py", "_dispatch_eval_derivative_n_times"), + "replace": ("sympy/core/basic.py", "replace"), + "xreplace": ("sympy/core/basic.py", "xreplace"), + "subs": ("sympy/core/basic.py", "subs"), + "as_real_imag (Pow)": ("sympy/core/power.py", "as_real_imag"), + "_jacobian_unwrap": ("underworld3/cython/petsc_generic_snes_solvers", "_jacobian_unwrap"), + "_extract_constants": ("underworld3/utilities/_jitextension.py", "_extract_constants"), + "compile_and_load": ("underworld3/utilities/_jitextension.py", "compile_and_load"), + } + for label, (path, func) in groups.items(): + total = sum(v[3] for k, v in st.stats.items() if path in k[0] and k[2] == func) + print(f" {label:52s} {total:7.2f} s") + print(f" profile: {out}") From d64de4af3fc8146725439230d47df3f84f0d511d Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 07:46:11 +1100 Subject: [PATCH 02/14] JIT graph route behind UW_JIT_GRAPH (#823, tier 2, staging steps 2 and 3) The JIT lowers each non-constant UWexpression to a node, an applied function of the leaves it reads, instead of expanding it into the tree. getext() emits one C temporary per distinct computation, ordered and merged by a hash of the C it computes, and takes the constants manifest from the leaves. _jacobian_unwrap builds guarded nodes, so the Newton tangent is formed by SymPy's chain rule through fdiff. The unwrappers expand nodes on request, for code that evaluates a lowered block. Selected by the private development switch UW_JIT_GRAPH=1; unset, every path is the tier 1 route unchanged. test_0024 closes the design note's risks (slot-keyed partial derivatives, two constants with one name, a stale node class, the guard placement, canonical emission, the manifest), each shown to fail under a mutation of the lowering. test_0022's guard test is pinned to the tree route. Session scripts: route_ab.py (end-to-end setup, solve, assembly timing), route_assemble.py (residual and Jacobian at a fixed state, both routes), rank_agreement_752.py, plain_diff_probe.py. Underworld development team with AI support from Claude Code --- scripts/sessions/jit_graph/fixtures.py | 25 + .../jit_graph/library_setup_profile.py | 4 + .../sessions/jit_graph/plain_diff_probe.py | 42 ++ .../sessions/jit_graph/rank_agreement_752.py | 48 ++ scripts/sessions/jit_graph/route_ab.py | 122 +++++ .../sessions/jit_graph/route_ab_compare.py | 28 ++ scripts/sessions/jit_graph/route_assemble.py | 55 +++ .../jit_graph/route_assemble_compare.py | 35 ++ .../cython/petsc_generic_snes_solvers.pyx | 12 +- src/underworld3/function/expressions.py | 12 + src/underworld3/utilities/_jit_graph.py | 442 ++++++++++++++++++ src/underworld3/utilities/_jitextension.py | 59 ++- ...022_unwrap_memoised_matches_fixed_point.py | 6 +- tests/test_0024_jit_graph_lowering.py | 267 +++++++++++ 14 files changed, 1149 insertions(+), 8 deletions(-) create mode 100644 scripts/sessions/jit_graph/plain_diff_probe.py create mode 100644 scripts/sessions/jit_graph/rank_agreement_752.py create mode 100644 scripts/sessions/jit_graph/route_ab.py create mode 100644 scripts/sessions/jit_graph/route_ab_compare.py create mode 100644 scripts/sessions/jit_graph/route_assemble.py create mode 100644 scripts/sessions/jit_graph/route_assemble_compare.py create mode 100644 src/underworld3/utilities/_jit_graph.py create mode 100644 tests/test_0024_jit_graph_lowering.py diff --git a/scripts/sessions/jit_graph/fixtures.py b/scripts/sessions/jit_graph/fixtures.py index e77c51fc7..28c1104fa 100644 --- a/scripts/sessions/jit_graph/fixtures.py +++ b/scripts/sessions/jit_graph/fixtures.py @@ -103,3 +103,28 @@ def build(name, **kwargs): stokes, admissible = BUILDERS[name](**kwargs) stokes.consistent_jacobian = True return stokes, admissible + + +def prepare_solve(stokes, name): + """Boundary conditions and the initial guess for a solve of fixture ``name``; + returns the keyword arguments for ``stokes.solve``. + + The box fixtures become a lid-driven box (shear everywhere, yielding at the lid + corners) started from simple shear, the lid's own profile: an uncapped power law + is singular at rest. The notch carries its own conditions and starts from rest. + """ + import numpy as np + + if name == "notch": + return {"zero_init_guess": True} + stokes.add_dirichlet_bc((1.0, 0.0), "Top") + stokes.add_dirichlet_bc((0.0, 0.0), "Bottom") + stokes.add_dirichlet_bc((0.0, 0.0), "Left") + stokes.add_dirichlet_bc((0.0, 0.0), "Right") + v0 = np.zeros_like(stokes.u.array) + v0[:, 0, 0] = stokes.u.coords[:, 1] + stokes.u.array[...] = v0 + kwargs = {"zero_init_guess": False} + if name in ("vep", "ti"): + kwargs["timestep"] = 0.1 + return kwargs diff --git a/scripts/sessions/jit_graph/library_setup_profile.py b/scripts/sessions/jit_graph/library_setup_profile.py index d350b993c..96557af34 100644 --- a/scripts/sessions/jit_graph/library_setup_profile.py +++ b/scripts/sessions/jit_graph/library_setup_profile.py @@ -22,6 +22,7 @@ params = uw.Params( uw_fixture=uw.Param("notch", description="box | powerlaw | linear | vep | ti | notch"), uw_profile=uw.Param(0, description="1 = cProfile the setup and print the hotspots"), + uw_dump=uw.Param("", description="write the canonicalised header to this path"), ) fixture = str(params.uw_fixture) stokes, _ = fixtures.build(fixture) @@ -38,6 +39,9 @@ def timed_generate(*a, **k): header = header.replace(modname, "__MOD__").replace(diag["randstr"], "__RS__") seen["header md5"] = hashlib.md5(header.encode()).hexdigest()[:10] seen["header bytes"] = len(header) + if str(params.uw_dump): + with open(os.path.expanduser(str(params.uw_dump)), "w") as fh: + fh.write(header) return modname, codeguys, diag diff --git a/scripts/sessions/jit_graph/plain_diff_probe.py b/scripts/sessions/jit_graph/plain_diff_probe.py new file mode 100644 index 000000000..8aafb710d --- /dev/null +++ b/scripts/sessions/jit_graph/plain_diff_probe.py @@ -0,0 +1,42 @@ +"""Whether a route still needs tier 1's real stand-in at the solvers' derivative sites +(#823, tier 2). + +Swaps ``diff_wrt_field`` / ``derive_by_array_wrt_field`` in the solvers for plain +``sympy.diff`` / ``sympy.derive_by_array`` and runs the Newton solve that motivated the +stand-in: the Drucker-Prager yield with a yield-stress floor at softness 0, whose floor +``(a + b + sqrt((a - b)**2)) / 2`` becomes an ``Abs`` of the pressure once fields are +real (``test_0023``). The route is chosen by ``UW_JIT_GRAPH``. +""" +import sympy +import underworld3 as uw +import underworld3.cython.generic_solvers as gs +from underworld3.utilities import _jit_graph + +gs.diff_wrt_field = lambda e, w: sympy.diff(e, w) +gs.derive_by_array_wrt_field = lambda e, dx: sympy.derive_by_array(e, dx) +route = "graph" if _jit_graph.enabled() else "tree" + +mesh = uw.meshing.UnstructuredSimplexBox(minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), + cellSize=0.25) +v = uw.discretisation.MeshVariable("V", mesh, 2, degree=2) +p = uw.discretisation.MeshVariable("P", mesh, 1, degree=1) +stokes = uw.systems.Stokes(mesh, velocityField=v, pressureField=p) +stokes.constitutive_model = uw.constitutive_models.ViscoPlasticFlowModel +stokes.add_dirichlet_bc((0.0, 0.0), "Bottom") +stokes.add_dirichlet_bc((1.0, 0.0), "Top") +stokes.bodyforce = sympy.Matrix([0, -mesh.X[0]]) +stokes.tolerance = 1.0e-10 +P = stokes.constitutive_model.Parameters +P.shear_viscosity_0 = 1.0 +P.yield_stress = (uw.expression(r"C", 100.0, "cohesion") + + uw.expression(r"\mu", 0.6, "friction") * p.sym[0]) +P.yield_stress_min = uw.expression(r"\tau", 0.01, "yield floor") +stokes.constitutive_model.yield_softness = 0 +stokes.consistent_jacobian = True +try: + stokes.solve() + print(f"[plain sympy.diff, {route}] solved: {stokes.solve_report.reason_str}, " + f"nl={stokes.solve_report.nl_its}") +except Exception as failure: + print(f"[plain sympy.diff, {route}] FAILED: {type(failure).__name__}: " + f"{str(failure).splitlines()[0][:200]}") diff --git a/scripts/sessions/jit_graph/rank_agreement_752.py b/scripts/sessions/jit_graph/rank_agreement_752.py new file mode 100644 index 000000000..b779352a9 --- /dev/null +++ b/scripts/sessions/jit_graph/rank_agreement_752.py @@ -0,0 +1,48 @@ +"""The #752 fixture: whether the ranks generate the same C (#823, tier 2). + +A power-law transversely isotropic Stokes on an annulus with rotated free-slip, the +law of ``tests/test_0022_rotated_adjoint.py`` (adjoint branches), whose JIT source +differed across ranks in about one np=2 run in five on development. Builds the +kernels and reports, on rank 0, whether ``_agree_source_across_ranks`` saw different +hashes. The route is chosen by ``UW_JIT_GRAPH``. Run under ``mpirun -n 2``. +""" +import sympy +import underworld3 as uw +import underworld3.utilities._jitextension as jx + +R_I, R_O = 0.5, 1.0 +seen = [] +agree = jx._agree_source_across_ranks + + +def counted(codeguys, source, source_hash): + seen.append(len(set(uw.mpi.comm.allgather(source_hash)))) + return agree(codeguys, source, source_hash) + + +jx._agree_source_across_ranks = counted + +mesh = uw.meshing.Annulus(radiusInner=R_I, radiusOuter=R_O, cellSize=0.15, qdegree=3) +x, y = mesh.X +r = sympy.sqrt(x ** 2 + y ** 2) +unit_r = sympy.Matrix([[x / r, y / r]]) +th = sympy.atan2(y, x) +v = uw.discretisation.MeshVariable("v_rot_adj", mesh, 2, degree=2) +p = uw.discretisation.MeshVariable("p_rot_adj", mesh, 1, degree=1) +eta_1 = uw.expression(r"\eta_1", 0.2, "weak-plane viscosity ratio") +stokes = uw.systems.Stokes(mesh, velocityField=v, pressureField=p) +edot = mesh.vector.strain_tensor(v.sym) +eII = sympy.sqrt(sympy.Rational(1, 2) * (edot[0, 0] ** 2 + edot[1, 1] ** 2) + edot[0, 1] ** 2) +eta_0 = (sympy.Float(0.01) + eII) ** sympy.Rational(-1, 3) +stokes.constitutive_model = uw.constitutive_models.TransverseIsotropicFlowModel +stokes.constitutive_model.Parameters.shear_viscosity_0 = eta_0 +stokes.constitutive_model.Parameters.shear_viscosity_1 = eta_1 * eta_0 +stokes.constitutive_model.Parameters.director = unit_r +stokes.bodyforce = 1.0e2 * sympy.cos(3 * th) * (r - R_I) / (R_O - R_I) * unit_r +stokes.add_dirichlet_bc((0.0, 0.0), "Lower") +stokes.add_rotated_freeslip_bc(0.0, "Upper") +stokes.consistent_jacobian = True +stokes.petsc_options["snes_max_it"] = 1 +stokes.solve(zero_init_guess=True) +uw.pprint(f"[752] getext calls {len(seen)}; calls whose ranks disagreed: " + f"{sum(1 for k in seen if k > 1)}") diff --git a/scripts/sessions/jit_graph/route_ab.py b/scripts/sessions/jit_graph/route_ab.py new file mode 100644 index 000000000..4fca2c91b --- /dev/null +++ b/scripts/sessions/jit_graph/route_ab.py @@ -0,0 +1,122 @@ +"""End-to-end A/B of the two JIT routes on one fixture (#823, tier 2). + +The route is chosen by the environment, so one build serves both: ``UW_JIT_GRAPH=1`` is +the graph route, unset is the expanded tree (tier 1). Run once per route, then +``route_ab_compare.py`` on the two ``.npz`` files. + +Records the whole pointwise setup (generation and compile separately), the size of the +generated header, the Newton solve (SNES reason, nonlinear and linear iterations, the +residual history), and the time of repeated residual and Jacobian assemblies at the +converged state. Run with ``UW_JIT_CACHE=0 UW_NO_USAGE_METRICS=1``. +""" +import hashlib +import os +import sys +import time + +import numpy as np +import underworld3 as uw +import underworld3.utilities._jitextension as jx +from underworld3.utilities import _jit_graph + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +import fixtures # noqa: E402 + +params = uw.Params( + uw_fixture=uw.Param("box", description="box | powerlaw | linear | vep | ti | notch"), + uw_tangent=uw.Param("newton", description="newton | picard | continuation"), + uw_maxit=uw.Param(30, description="SNES iteration limit"), + uw_tol=uw.Param(1.0e-8, description="SNES relative tolerance"), + uw_assemble=uw.Param(20, description="residual and Jacobian assemblies to time"), + uw_out=uw.Param("~/+Simulations/jit_graph/tier2", description="output directory"), +) +route = "graph" if _jit_graph.enabled() else "tree" +fixture = str(params.uw_fixture) +tangent = str(params.uw_tangent) + +stokes, _ = fixtures.build(fixture) +stokes.consistent_jacobian = {"newton": True, "picard": False, + "continuation": "continuation"}[tangent] +solve_kwargs = fixtures.prepare_solve(stokes, fixture) +stokes.tolerance = float(params.uw_tol) +stokes.petsc_options["snes_max_it"] = int(params.uw_maxit) +if "ksp_monitor" in stokes.petsc_options: + stokes.petsc_options.delValue("ksp_monitor") + +seen = {"generate": 0.0, "compile": 0.0, "pointwise": 0.0, "n_pointwise": 0, + "header bytes": 0} +generate, compile_ = jx.generate_c_source, jx.compile_and_load + + +def timed_generate(*a, **k): + t = time.perf_counter() + modname, codeguys, diag = generate(*a, **k) + seen["generate"] += time.perf_counter() - t + header = dict(codeguys)["cy_ext.h"] + header = header.replace(modname, "__MOD__").replace(diag["randstr"], "__RS__") + seen["header md5"] = hashlib.md5(header.encode()).hexdigest()[:10] + seen["header bytes"] += len(header) + return modname, codeguys, diag + + +def timed_compile(*a, **k): + t = time.perf_counter() + out = compile_(*a, **k) + seen["compile"] += time.perf_counter() - t + return out + + +jx.generate_c_source, jx.compile_and_load = timed_generate, timed_compile +pointwise = type(stokes)._setup_pointwise_functions + + +def timed_pointwise(self, *a, **k): + t = time.perf_counter() + out = pointwise(self, *a, **k) + seen["pointwise"] += time.perf_counter() - t + seen["n_pointwise"] += 1 + return out + + +type(stokes)._setup_pointwise_functions = timed_pointwise + +t = time.perf_counter() +stokes.solve(**solve_kwargs) +wall = time.perf_counter() - t +r = stokes.solve_report + +# repeated assemblies at the converged state +snes = stokes.snes +X = snes.getSolution() +F = X.duplicate() +J, P = snes.getJacobian()[0], snes.getJacobian()[1] +n = int(params.uw_assemble) +snes.computeFunction(X, F) +t = time.perf_counter() +for _ in range(n): + snes.computeFunction(X, F) +t_res = (time.perf_counter() - t) / n +snes.computeJacobian(X, J, P) +t = time.perf_counter() +for _ in range(n): + snes.computeJacobian(X, J, P) +t_jac = (time.perf_counter() - t) / n + +print(f"[{fixture} {tangent} {route}] pointwise setup {seen['pointwise']:.2f} s " + f"({seen['n_pointwise']} call) = generate {seen['generate']:.2f} s + compile " + f"{seen['compile']:.2f} s + other; header {seen['header bytes']:,} B " + f"md5 {seen.get('header md5')}") +print(f"[{fixture} {tangent} {route}] solve {wall:.2f} s: {r.reason_str} nl={r.nl_its} " + f"ksp={r.ksp_its} fnorm={r.fnorm:.3e}") +print(f"[{fixture} {tangent} {route}] assembly at the solution: residual {1e3 * t_res:.2f} ms, " + f"Jacobian {1e3 * t_jac:.2f} ms (mean of {n})") + +out = os.path.expanduser(str(params.uw_out)) +os.makedirs(out, exist_ok=True) +np.savez(os.path.join(out, f"{fixture}_{tangent}_{route}.npz"), + v=np.asarray(stokes.u.array), p=np.asarray(stokes.p.array), + X=np.asarray(X.array).copy(), + history=np.asarray(r.history, dtype=float), + reason=r.reason_str, nl=r.nl_its, ksp=r.ksp_its, + pointwise=seen["pointwise"], generate=seen["generate"], compile=seen["compile"], + header_bytes=seen["header bytes"], solve=wall, t_res=t_res, t_jac=t_jac) diff --git a/scripts/sessions/jit_graph/route_ab_compare.py b/scripts/sessions/jit_graph/route_ab_compare.py new file mode 100644 index 000000000..95b79130d --- /dev/null +++ b/scripts/sessions/jit_graph/route_ab_compare.py @@ -0,0 +1,28 @@ +"""Compare the two routes' records written by ``route_ab.py`` for one fixture.""" +import os +import sys + +import numpy as np + +out = os.path.expanduser(sys.argv[1] if len(sys.argv) > 1 else "~/+Simulations/jit_graph/tier2") +cases = sorted({f.rsplit("_", 1)[0] for f in os.listdir(out) if f.endswith(".npz")}) +for case in cases: + paths = [os.path.join(out, f"{case}_{r}.npz") for r in ("tree", "graph")] + if not all(os.path.exists(p) for p in paths): + continue + a, b = (np.load(p) for p in paths) + dv = np.abs(a["v"] - b["v"]).max() / max(np.abs(a["v"]).max(), 1e-300) + dp = np.abs(a["p"] - b["p"]).max() / max(np.abs(a["p"]).max(), 1e-300) + ha, hb = a["history"], b["history"] + k = min(len(ha), len(hb)) + dh = (np.abs(ha[:k] - hb[:k]) / np.maximum(np.abs(ha[:k]), 1e-300)).max() if k else np.nan + print(f"{case}") + print(f" {'':12s} {'tree':>12s} {'graph':>12s}") + for key, fmt in (("pointwise", ".2f"), ("generate", ".2f"), ("compile", ".2f"), + ("header_bytes", ",d"), ("solve", ".2f"), ("t_res", ".2e"), + ("t_jac", ".2e"), ("nl", "d"), ("ksp", "d")): + va, vb = a[key].item(), b[key].item() + print(f" {key:12s} {va:>12{fmt}} {vb:>12{fmt}}") + print(f" reason {str(a['reason']):>12s} {str(b['reason']):>12s}") + print(f" solution: max|dv|/max|v| {dv:.2e}, max|dp|/max|p| {dp:.2e}; " + f"residual history (first {k}) max rel diff {dh:.2e}") diff --git a/scripts/sessions/jit_graph/route_assemble.py b/scripts/sessions/jit_graph/route_assemble.py new file mode 100644 index 000000000..fa44e50a1 --- /dev/null +++ b/scripts/sessions/jit_graph/route_assemble.py @@ -0,0 +1,55 @@ +"""Assemble the residual and the Jacobian of one fixture at a given state, on the route +the environment selects (``UW_JIT_GRAPH=1`` or unset), and save them (#823, tier 2). + +Two runs at the same state, one per route, are compared by +``route_assemble_compare.py``. The state is the SNES solution vector: ``zero`` (rest, +with the boundary values), or the ``X`` saved by ``route_ab.py``. +""" +import os +import sys + +import numpy as np +import underworld3 as uw +from underworld3.utilities import _jit_graph + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +import fixtures # noqa: E402 + +params = uw.Params( + uw_fixture=uw.Param("box", description="box | powerlaw | linear | vep | ti | notch"), + uw_tangent=uw.Param("newton", description="newton | picard | continuation"), + uw_state=uw.Param("zero", description="zero, or a route_ab .npz whose X is the state"), + uw_tag=uw.Param("zero", description="label for the output file"), + uw_out=uw.Param("~/+Simulations/jit_graph/tier2/assemble", description="output directory"), +) +route = "graph" if _jit_graph.enabled() else "tree" +fixture, tangent = str(params.uw_fixture), str(params.uw_tangent) + +stokes, _ = fixtures.build(fixture) +stokes.consistent_jacobian = {"newton": True, "picard": False, + "continuation": "continuation"}[tangent] +kwargs = fixtures.prepare_solve(stokes, fixture) +stokes.petsc_options["snes_max_it"] = 0 +if "ksp_monitor" in stokes.petsc_options: + stokes.petsc_options.delValue("ksp_monitor") +stokes.solve(**kwargs) + +snes = stokes.snes +X = snes.getSolution() +state = str(params.uw_state) +if state == "zero": + X.set(0.0) +else: + X.array[:] = np.load(os.path.expanduser(state))["X"] +F = X.duplicate() +snes.computeFunction(X, F) +J, P = snes.getJacobian()[0], snes.getJacobian()[1] +snes.computeJacobian(X, J, P) +ai, aj, av = J.getValuesCSR() + +out = os.path.expanduser(str(params.uw_out)) +os.makedirs(out, exist_ok=True) +path = os.path.join(out, f"{fixture}_{tangent}_{params.uw_tag}_{route}.npz") +np.savez(path, F=np.asarray(F.array).copy(), ai=ai, aj=aj, av=av) +print(f"[{fixture} {tangent} {params.uw_tag} {route}] |F| {F.norm():.6e} " + f"nnz {len(av)} saved {path}") diff --git a/scripts/sessions/jit_graph/route_assemble_compare.py b/scripts/sessions/jit_graph/route_assemble_compare.py new file mode 100644 index 000000000..606318791 --- /dev/null +++ b/scripts/sessions/jit_graph/route_assemble_compare.py @@ -0,0 +1,35 @@ +"""Compare the residuals and Jacobians ``route_assemble.py`` saved for the two routes.""" +import os +import sys + +import numpy as np + +out = os.path.expanduser(sys.argv[1] if len(sys.argv) > 1 else + "~/+Simulations/jit_graph/tier2/assemble") +cases = sorted({f.rsplit("_", 1)[0] for f in os.listdir(out) if f.endswith(".npz")}) +for case in cases: + paths = [os.path.join(out, f"{case}_{r}.npz") for r in ("tree", "graph")] + if not all(os.path.exists(p) for p in paths): + continue + a, b = (np.load(p) for p in paths) + fa, fb = a["F"], b["F"] + nan = (int((~np.isfinite(fa)).sum()), int((~np.isfinite(fb)).sum())) + df = np.abs(fa - fb).max() / max(np.abs(fa).max(), 1e-300) + same_pattern = np.array_equal(a["ai"], b["ai"]) and np.array_equal(a["aj"], b["aj"]) + ja, jb = a["av"], b["av"] + jnan = (int((~np.isfinite(ja)).sum()), int((~np.isfinite(jb)).sum())) + if same_pattern: + dj = np.abs(ja - jb) + rel = dj.max() / max(np.abs(ja).max(), 1e-300) + # entry by entry, against the largest entry of its row + rows = np.repeat(np.arange(len(a["ai"]) - 1), np.diff(a["ai"])) + rowmax = np.zeros(len(a["ai"]) - 1) + np.maximum.at(rowmax, rows, np.abs(ja)) + rowrel = (dj / np.maximum(rowmax[rows], 1e-300)).max() + bit = (ja == jb).mean() + jtxt = (f"J: max|dJ|/max|J| {rel:.2e}, worst entry against its row {rowrel:.2e}, " + f"bit-identical {100 * bit:.1f}%") + else: + jtxt = "J: SPARSITY PATTERNS DIFFER" + print(f"{case}: F max|dF|/max|F| {df:.2e}, bit-identical {100 * (fa == fb).mean():.1f}%; " + f"{jtxt}; non-finite F {nan}, J {jnan}") diff --git a/src/underworld3/cython/petsc_generic_snes_solvers.pyx b/src/underworld3/cython/petsc_generic_snes_solvers.pyx index 0e73b97b9..46c5a81c8 100644 --- a/src/underworld3/cython/petsc_generic_snes_solvers.pyx +++ b/src/underworld3/cython/petsc_generic_snes_solvers.pyx @@ -134,8 +134,16 @@ def _jacobian_unwrap(expr): return guard(e) - f = lambda e: _guard_sqrts( - _unwrap_expression(e, mode="symbolic_keep_constants")) + from underworld3.utilities import _jit_graph + if _jit_graph.enabled(): + # The graph route (#823, tier 2): each non-constant atom becomes its guarded + # node instead of being expanded, and the derivative passes through it by + # the chain rule. The guard is applied in each node body and at the top level. + graph = _jit_graph.KernelGraph() + f = lambda e: _jit_graph.guard_half_integer_powers(graph.lower(e, guarded=True)) + else: + f = lambda e: _guard_sqrts( + _unwrap_expression(e, mode="symbolic_keep_constants")) if isinstance(expr, sympy.MatrixBase): return expr.applyfunc(f) if isinstance(expr, sympy.NDimArray): diff --git a/src/underworld3/function/expressions.py b/src/underworld3/function/expressions.py index 05bde6184..b119f98c2 100644 --- a/src/underworld3/function/expressions.py +++ b/src/underworld3/function/expressions.py @@ -169,11 +169,14 @@ def _unwrap_expression_once(expr, mode='nondimensional'): Expression with UW atoms substituted """ from underworld3.coordinates import UWCoordinate + from underworld3.utilities import _jit_graph # Handle non-expression types directly if isinstance(expr, UWQuantity) and not isinstance(expr, UWexpression): return _unwrap_atom(expr, mode) + expr = _jit_graph.expand_nodes(expr) + if isinstance(expr, UWCoordinate): return _unwrap_atom(expr, mode) @@ -222,6 +225,11 @@ def _unwrap_expression_complete(expr, mode): same Python object, which identity-memoised walks downstream exploit. """ from underworld3.coordinates import UWCoordinate + from underworld3.utilities import _jit_graph + + # a block lowered onto the JIT's graph holds node applications: each is one more + # atom, whose expansion is its body (#823) + expr = _jit_graph.expand_nodes(expr) uw_types = (UWexpression, UWQuantity, UWCoordinate) expanded = {} # id(atom) -> (atom, expansion); the atom is held so its id stays valid @@ -619,6 +627,10 @@ def unwrap_for_evaluate(expr, scaling_active=None): else: sym_expr = expr + # a block lowered onto the JIT's graph: nodes expand to their bodies (#823) + from underworld3.utilities import _jit_graph + sym_expr = _jit_graph.expand_nodes(sym_expr) + # Step 4: Process composite expressions - TYPE-BASED DISPATCH if isinstance(sym_expr, sympy.Expr): substitutions = {} diff --git a/src/underworld3/utilities/_jit_graph.py b/src/underworld3/utilities/_jit_graph.py new file mode 100644 index 000000000..d9cb624d8 --- /dev/null +++ b/src/underworld3/utilities/_jit_graph.py @@ -0,0 +1,442 @@ +r"""JIT kernels lowered onto the shared expression graph (#823, tier 2). + +Design: ``docs/developer/design/jit-shared-graph-codegen.md``. + +A constitutive law is a graph of named sub-expressions (``UWexpression`` atoms). The +expanded-tree route copies every atom into every place that refers to it before it +differentiates and prints. This module keeps the graph instead: + +- each non-constant atom becomes a NODE, an applied undefined function of the leaves + its value depends on (field values and gradients, coordinates, constant atoms), whose + body is the atom's content with each child atom replaced by the child's node; +- SymPy's own chain rule differentiates through a node, because ``fdiff(i)`` returns the + partial derivative of the body with respect to argument slot ``i``, itself a node; +- the emitter writes one C temporary per distinct computation, ordered and merged by a + hash of the C it computes, so the generated source is a function of the mathematics + and of the kernel's data layout alone. + +Selected by the private development switch ``UW_JIT_GRAPH=1`` while both routes exist. +""" +import hashlib +import itertools +import os + +import sympy +from sympy.core.function import AppliedUndef, UndefinedFunction +from sympy.tensor.array import NDimArray +from sympy.vector.scalar import BaseScalar + +_EPS2 = sympy.Float(1.0e-36) +_serial = itertools.count(1) +_nodes_made = False + + +def enabled(): + """Whether the graph route is selected (``UW_JIT_GRAPH=1``).""" + return os.environ.get("UW_JIT_GRAPH", "").lower() in ("1", "true", "yes") + + +def nodes_exist(): + """Whether any node has been made in this process; when not, no expression can + hold one and the unwrappers skip the search.""" + return _nodes_made + + +def guard_half_integer_powers(e): + r"""``e`` with :math:`10^{-36}` added to the base of every half-integer power + whose base has free symbols: the sqrt guard of ``_jacobian_unwrap``, the same rule + and the same rebuild. Node applications are not entered (a guarded node's body was + guarded when it was made).""" + memo = {} + + def guard(n): + hit = memo.get(id(n)) + if hit is not None: + return hit[1] + out = n + args = getattr(n, "args", None) + if args and not isinstance(n, _KernelNode): + new_args = tuple(guard(a) for a in args) + if any(a is not b for a, b in zip(args, new_args)) and args != new_args: + out = n.func(*new_args) + # replace(simultaneous=True): a rebuild that collapses to one of the + # changed arguments is not matched again + if any(out == a and a != b for a, b in zip(args, new_args)): + memo[id(n)] = (n, out) + return out + if (out.is_Pow and out.exp.is_Rational and out.exp.q == 2 + and out.args[0].free_symbols): + out = sympy.Pow(out.args[0] + _EPS2, out.exp) + memo[id(n)] = (n, out) + return out + + return guard(e) + + +class _KernelNode(AppliedUndef): + """A named quantity of a kernel, applied to the leaves its value depends on.""" + + def fdiff(self, argindex=1): + return self._graph.slot_derivative(self, argindex - 1) + + def _ccode(self, printer): + raise RuntimeError( + f"JIT: node {self.func.__name__} reached the C printer; every node must be " + f"emitted as a temporary (a defect in the graph lowering, #823)") + + +class _Temporary(sympy.Symbol): + """A C temporary of one kernel, written ``uwt_``.""" + + __slots__ = ("_ccodestr",) + + def __new__(cls, index): + obj = sympy.Symbol.__xnew__(cls, f"uwt_{index}", real=True) + obj._ccodestr = f"uwt_{index}" + return obj + + def _ccode(self, printer): + return self._ccodestr + + +def body_of(app): + """A node application's body, with the application's arguments substituted.""" + cls = app.func + if app.args == cls._deps: + return cls._body + return cls._body.xreplace(dict(zip(cls._deps, app.args))) + + +def _holds_node(expr): + from underworld3.utilities._jitextension import _holds_instance + return _nodes_made and _holds_instance(expr, _KernelNode) + + +def expand_nodes(expr): + """``expr`` (an expression, Matrix or Array) with every node application replaced + by its body, recursively: the expanded tree of the expression. For code outside the + JIT that evaluates a lowered block (a test's oracle, ``uw.function.evaluate``).""" + if not _holds_node(expr): + return expr + memo = {} + + def expand(e): + apps = e.atoms(_KernelNode) + if not apps: + return e + rule = {} + for app in apps: + hit = memo.get(app) + if hit is None: + hit = memo[app] = expand(body_of(app)) + rule[app] = hit + return e.xreplace(rule) + + if isinstance(expr, sympy.MatrixBase): + return expr.applyfunc(expand) + if isinstance(expr, NDimArray): + return type(expr)([expand(e) for e in expr], expr.shape) + return expand(expr) + + +class KernelGraph: + """One lowering context: the nodes made for one compile, or for one Jacobian + source. Each node class holds its body and this context, so a node is read + back without the context being passed.""" + + def __init__(self): + self.serial = next(_serial) + self._const = {} # id(atom) -> (atom, bool) + self._node = {} # (id(atom), guarded) -> (atom, replacement) + self._by_body = {} # body -> node class + self._deriv = {} # (node class, slot) -> derivative in the class's deps + self._busy = set() + self._splitting = False + + # ------------------------------------------------------------------ leaves + def is_constant(self, atom): + hit = self._const.get(id(atom)) + if hit is None: + from underworld3.function.expressions import UWexpression + from underworld3.utilities._jitextension import _is_truly_constant + hit = self._const[id(atom)] = (atom, _is_truly_constant(atom, UWexpression)) + return hit[1] + + def is_leaf(self, s): + from underworld3.function.expressions import UWexpression + + if isinstance(s, _KernelNode): + return False + if isinstance(s, (AppliedUndef, BaseScalar)): + return True + if isinstance(s, UWexpression): + return self.is_constant(s) + return isinstance(s, sympy.Symbol) + + @staticmethod + def display_key(s): + """What a leaf is called, without creation counters. Not unique (two + constants may share a name); used to name and order, never to identify.""" + from underworld3.function.expressions import UWexpression + + if isinstance(s, AppliedUndef): + return f"F|{s.func.__name__}|{','.join(map(str, s.args))}" + if isinstance(s, BaseScalar): + return f"X|{s._id[1]}|{s._id[0]}" + if isinstance(s, UWexpression): + return f"C|{s.name}" + return f"S|{getattr(s, 'name', s)}" + + @staticmethod + def order_key(s): + """Display key, ties broken by relative creation order (as the constants + manifest breaks them).""" + return (KernelGraph.display_key(s), getattr(s, "instance_number", 0)) + + def leaves(self, e): + out, seen, stack = set(), set(), [e] + while stack: + a = stack.pop() + if id(a) in seen: + continue + seen.add(id(a)) + if isinstance(a, _KernelNode): + out.update(a.args) + elif self.is_leaf(a): + out.add(a) + elif isinstance(a, sympy.Basic): + stack.extend(a.args) + return out + + # ------------------------------------------------------------------ lowering + def lower(self, e, guarded=False): + """``e`` (an expression, Matrix or Array) with each UW atom lowered: a + non-constant ``UWexpression`` to its node, a coordinate to its base scalar, + a bare quantity to its non-dimensional value. Constant atoms stay as they + are; they are the ``constants[]`` leaves. ``guarded`` selects the node + variant with the sqrt guard in every body (a Newton source).""" + from underworld3.coordinates import UWCoordinate + from underworld3.function.expressions import UWexpression + from underworld3.function.quantities import UWQuantity + + if isinstance(e, sympy.MatrixBase): + return e.applyfunc(lambda x: self.lower(x, guarded)) + if isinstance(e, NDimArray): + return type(e)([self.lower(x, guarded) for x in e], e.shape) + if isinstance(e, (UWexpression, UWCoordinate)) or ( + isinstance(e, UWQuantity) and not isinstance(e, UWexpression)): + return self._replacement(e, guarded) + if not isinstance(e, sympy.Basic): + return sympy.sympify(e) + uw_types = (UWexpression, UWQuantity, UWCoordinate) + atoms = sorted((s for s in e.free_symbols if isinstance(s, uw_types)), + key=self.order_key) + rule = {} + for s in atoms: + r = self._replacement(s, guarded) + if r is not s: + rule[s] = r + return e.xreplace(rule) if rule else e + + def _replacement(self, atom, guarded): + from underworld3.coordinates import UWCoordinate + from underworld3.function.expressions import UWexpression, _unwrap_atom + + if isinstance(atom, UWCoordinate): + return atom.sym + if not isinstance(atom, UWexpression): # a bare quantity + return sympy.sympify(_unwrap_atom(atom, "nondimensional")) + if self.is_constant(atom): + return atom + return self.node_of(atom, guarded) + + def node_of(self, atom, guarded): + key = (id(atom), guarded) + hit = self._node.get(key) + if hit is not None: + return hit[1] + if key in self._busy: + raise ValueError(f"cyclic expression: {atom} contains itself") + self._busy.add(key) + try: + body = self.lower(atom.sym, guarded) + if guarded: + body = guard_half_integer_powers(body) + # a matrix-valued atom cannot be one scalar temporary: it is expanded + # in place, as the tree route expands it + out = self.make_node(body) if isinstance(body, sympy.Expr) else body + finally: + self._busy.discard(key) + self._node[key] = (atom, out) + return out + + def _display_name(self, body): + canon = {a: sympy.Symbol(a.func.__name__) for a in body.atoms(_KernelNode)} + for leaf in self.leaves(body): + canon.setdefault(leaf, sympy.Symbol(self.display_key(leaf))) + text = sympy.srepr(body.xreplace(canon)) + return "N" + hashlib.sha1(text.encode()).hexdigest()[:12] + + def make_node(self, body): + """The node for ``body``, one per distinct body. A body that is a number, a + single leaf, a single node, or reads no leaf is returned as itself.""" + global _nodes_made + + body = sympy.sympify(body) + if not self._splitting: + body = self._split_shared(body) + if body.is_Atom or isinstance(body, (AppliedUndef, BaseScalar)): + return body + cls = self._by_body.get(body) + if cls is None: + deps = tuple(sorted(self.leaves(body), key=self.order_key)) + if not deps: + return body + cls = UndefinedFunction(self._display_name(body), bases=(_KernelNode,), + real=True, _ctx=self.serial, _n=len(self._by_body), + __dict__={"_graph": self}) + cls._body, cls._deps = body, deps + self._by_body[body] = cls + _nodes_made = True + return cls(*cls._deps) + + def _split_shared(self, body): + """Repeated unnamed sub-expressions of one body become anonymous nodes. + Canonical order, so the split does not depend on the hash seed.""" + if body.is_Atom: + return body + repl, (reduced,) = sympy.cse([body], symbols=sympy.numbered_symbols("_cse", real=True)) + if not repl: + return body + self._splitting = True + try: + rule = {} + for sym, e in repl: + rule[sym] = self.make_node(e.xreplace(rule)) + return reduced.xreplace(rule) + finally: + self._splitting = False + + # ------------------------------------------------------------------ derivatives + def slot_derivative(self, app, i): + """The partial derivative of ``app`` with respect to its argument slot + ``i``, as a node of the same leaves. Taken with every leaf replaced by an + independent real dummy, so a field that is itself a function of the + coordinates is not differentiated a second time through them.""" + cls = app.func + key = (cls, i) + if key not in self._deriv: + deps = cls._deps + dummies = [sympy.Dummy(real=True) for _ in deps] + b = cls._body.xreplace(dict(zip(deps, dummies))) + d = sympy.diff(b, dummies[i]).xreplace(dict(zip(dummies, deps))) + self._deriv[key] = self.make_node(d) + d = self._deriv[key] + if app.args == cls._deps: + return d + return d.xreplace(dict(zip(cls._deps, app.args))) + + +def lower_callbacks(fns, mesh): + """Each callback lowered in one context and shaped as ``generate_c_source`` + prints it: a Matrix whose entries hold nodes, constant atoms and leaves.""" + from underworld3.function.expressions import UWDerivativeExpression + + graph = KernelGraph() + out = [] + for fn in fns: + if isinstance(fn, UWDerivativeExpression): + fn = fn.doit() + fn = graph.lower(fn) + if isinstance(fn, sympy.vector.Vector): + fn = fn.to_matrix(mesh.N)[0:mesh.dim, 0] + elif isinstance(fn, sympy.vector.Dyadic): + fn = fn.to_matrix(mesh.N)[0:mesh.dim, 0:mesh.dim] + else: + fn = sympy.Matrix([fn]) + out.append(fn) + return out + + +def constant_leaves(lowered): + """The constant atoms the lowered callbacks read. A node application's arguments + are every leaf its body reads, child nodes included, so the free symbols of the + outputs reach every constant of every temporary.""" + from underworld3.function.expressions import UWexpression + + found = set() + for fn in lowered: + found.update(s for s in fn.free_symbols if isinstance(s, UWexpression)) + return found + + +def emission_order(outputs, spell): + """The temporaries a kernel needs, in the order to write them. + + ``spell(leaf)`` is the C the kernel reads for a leaf. Each node application gets a + canonical key: a hash of its body with leaves spelled as C and children replaced by + their own keys. Applications with equal keys compute the same C and share one + temporary; temporaries follow the ones they use, ties broken by key. + + Returns ``(order, key)``: ``order`` is a list of ``(key, body)``; ``key`` maps each + node application reached to its key. + """ + key = {} + spelt = {} + + def spelling(leaf): + hit = spelt.get(leaf) + if hit is None: + hit = spelt[leaf] = sympy.Symbol(spell(leaf)) + return hit + + def key_of(app): + hit = key.get(app) + if hit is not None: + return hit + body = body_of(app) + sub = {c: sympy.Symbol("@" + key_of(c)) for c in body.atoms(_KernelNode)} + for leaf in body.free_symbols | body.atoms(AppliedUndef): + if not isinstance(leaf, _KernelNode) and leaf not in sub: + sub[leaf] = spelling(leaf) + k = hashlib.sha1(sympy.srepr(body.xreplace(sub)).encode()).hexdigest()[:16] + key[app] = k + return k + + order, placed = [], set() + + def visit(app): + k = key_of(app) + if k in placed: + return + placed.add(k) + body = body_of(app) + for child in sorted(body.atoms(_KernelNode), key=key_of): + visit(child) + order.append((k, body)) + + for out in outputs: + for child in sorted(sympy.sympify(out).atoms(_KernelNode), key=key_of): + visit(child) + return order, key + + +def emit(fn, spell, constants_rule): + """``fn`` (a lowered Matrix) as temporaries and outputs. + + Returns ``(temporaries, outputs)``: ``temporaries`` is a list of + ``(_Temporary, body)`` in the order to write them, and ``outputs`` is ``fn`` with + each node application replaced by its temporary. Constant atoms are replaced by + ``constants_rule`` (their ``constants[]`` placeholders) in both. + """ + order, key = emission_order(list(fn), spell) + temp_of, temporaries = {}, [] + for i, (k, body) in enumerate(order): + t = _Temporary(i) + rule = {c: temp_of[key[c]] for c in body.atoms(_KernelNode)} + rule.update(constants_rule) + temporaries.append((t, body.xreplace(rule))) + temp_of[k] = t + rule = {app: temp_of[key[app]] for app in fn.atoms(_KernelNode)} + rule.update(constants_rule) + return temporaries, fn.xreplace(rule) diff --git a/src/underworld3/utilities/_jitextension.py b/src/underworld3/utilities/_jitextension.py index 3daef9703..05a3389d2 100644 --- a/src/underworld3/utilities/_jitextension.py +++ b/src/underworld3/utilities/_jitextension.py @@ -6,6 +6,7 @@ import sympy import underworld3 import underworld3.timing as timing +from underworld3.utilities import _jit_graph from collections import namedtuple from dataclasses import dataclass from pathlib import Path @@ -447,6 +448,12 @@ def _extract_constants(all_fns, mesh): else: _collect_constant_atoms(fn, constant_exprs, is_constant_expr, UWexpression) + return _manifest_from(constant_exprs) + + +def _manifest_from(constant_exprs): + """The ``constants[]`` manifest of a set of constant atoms: ``(manifest, + subs_map)``, slots ordered by name, then creation order.""" if not constant_exprs: return [], {} @@ -821,7 +828,15 @@ def getext( # Extract constant UWexpressions that are routed through PETSc's # constants[] array. Value changes don't affect the C source — they # only alter what we pass to PetscDSSetConstants at solve time. - constants_manifest, constants_subs_map = _extract_constants(callbacks.flat(), mesh) + if _jit_graph.enabled(): + # The graph route (#823, tier 2): lower every callback once, and take the + # manifest from the constant leaves of what was lowered. + lowered_fns = _jit_graph.lower_callbacks(callbacks.flat(), mesh) + constants_manifest, constants_subs_map = _manifest_from( + _jit_graph.constant_leaves(lowered_fns)) + else: + lowered_fns = None + constants_manifest, constants_subs_map = _extract_constants(callbacks.flat(), mesh) if debug and underworld3.mpi.rank == 0: if constants_manifest: @@ -842,6 +857,7 @@ def getext( verbose=verbose, debug=debug, debug_name=debug_name, + lowered_fns=lowered_fns, ) gen_randstr = diag["randstr"] @@ -1098,6 +1114,7 @@ def generate_c_source( verbose: Optional[bool] = False, debug: Optional[bool] = False, debug_name=None, + lowered_fns=None, ): """Generate the setup.py / C header / Cython wrapper for a JIT bundle. @@ -1118,6 +1135,10 @@ def generate_c_source( Variables that map to PETSc primary variable arrays (``petsc_u[]``). constants_subs_map : dict, optional Mapping from UWexpression → ``_JITConstant`` placeholder. + lowered_fns : list of sympy.Matrix, optional + The callbacks lowered onto the shared graph (``_jit_graph.lower_callbacks``), + one per entry of ``callbacks.flat()``. When given, each kernel is emitted as + temporaries and outputs instead of being unwrapped and printed whole. Returns ------- @@ -1353,11 +1374,20 @@ def _handle_UnevaluatedExpr(expr): underworld3._libdirs.clear() underworld3._libfiles.clear() + def _spell(leaf): + # the C a kernel reads for a leaf of the graph + placeholder = constants_subs_map.get(leaf) if constants_subs_map else None + return placeholder._ccodestr if placeholder is not None else printer.doprint(leaf) + eqns = [] for index, fn in enumerate(fns): # Save original for debugging fn_original = fn + temporaries = () + if lowered_fns is not None: + temporaries, fn = _jit_graph.emit( + lowered_fns[index], _spell, constants_subs_map or {}) # --- Gate the UW lowering (issue #302 pipeline) on the presence of # UW-expression atoms. Plain-sympy components — the derivative @@ -1368,7 +1398,7 @@ def _handle_UnevaluatedExpr(expr): from underworld3.function.expressions import UWexpression as _UWexpr from underworld3.function.expressions import UWDerivativeExpression as _UWderiv - _needs_lowering = ( + _needs_lowering = lowered_fns is None and ( isinstance(fn, (_UWexpr, _UWderiv)) # `has` is a bare traversal (no atom-set build) — the atoms() # form built a set of every node, which cost seconds per 100k-node @@ -1425,7 +1455,9 @@ def _handle_UnevaluatedExpr(expr): f"(issue #302)." ) - if isinstance(fn, sympy.vector.Vector): + if lowered_fns is not None: + pass # shaped by _jit_graph.lower_callbacks + elif isinstance(fn, sympy.vector.Vector): fn = fn.to_matrix(mesh.N)[0 : mesh.dim, 0] elif isinstance(fn, sympy.vector.Dyadic): fn = fn.to_matrix(mesh.N)[0 : mesh.dim, 0 : mesh.dim] @@ -1441,7 +1473,10 @@ def _handle_UnevaluatedExpr(expr): # We recover _ccodestr from the coordinate's _id attribute. from sympy.vector.scalar import BaseScalar - free_syms = tuple(_stable_sorted(fn.free_symbols)) + free_syms = fn.free_symbols + for _t, _body in temporaries: + free_syms = free_syms | _body.free_symbols + free_syms = tuple(_stable_sorted(free_syms)) for sym in free_syms: if isinstance(sym, BaseScalar) and not hasattr(sym, '_ccodestr'): idx = sym._id[0] # 0, 1, or 2 for x, y, z @@ -1508,7 +1543,21 @@ def _handle_UnevaluatedExpr(expr): # Semantics-preserving: temps are exact aliases of repeated # subexpressions, so the generated kernel evaluates identical values. # Opt in with UW_JIT_CSE=1 (default is off to preserve original behavior). - if os.environ.get("UW_JIT_CSE") in ("1", "true", "True", "yes", "YES"): + if temporaries: + # the graph route: one temporary per distinct computation, then the + # outputs in terms of them (#823) + _temp_code = [] + for _t, _body in temporaries: + _code = printer.doprint(_body) + if _code.startswith("// Not supported in C:"): + _temp_code = None + eqn = ("eqn_" + str(index), _code) + break + _temp_code.append(f"const double {_t._ccodestr} = {_code};") + if _temp_code is not None: + eqn = ("eqn_" + str(index), + "\n".join(_temp_code) + "\n" + printer.doprint(fn, out)) + elif lowered_fns is None and os.environ.get("UW_JIT_CSE") in ("1", "true", "True", "yes", "YES"): from sympy.simplify.cse_main import cse from sympy.vector.scalar import BaseScalar diff --git a/tests/test_0022_unwrap_memoised_matches_fixed_point.py b/tests/test_0022_unwrap_memoised_matches_fixed_point.py index 573f84302..be7477277 100644 --- a/tests/test_0022_unwrap_memoised_matches_fixed_point.py +++ b/tests/test_0022_unwrap_memoised_matches_fixed_point.py @@ -101,9 +101,13 @@ def test_memoised_unwrap_is_the_fixed_point(laws, mode): assert sympy.srepr(new) == sympy.srepr(old), (name, mode) -def test_the_jacobian_sqrt_guard_matches_replace(laws): +def test_the_jacobian_sqrt_guard_matches_replace(laws, monkeypatch): + """The expanded-tree route's guard. The graph route guards each node body + instead; test_0024 holds it to this one by value.""" from underworld3.cython.generic_solvers import _jacobian_unwrap + monkeypatch.setenv("UW_JIT_GRAPH", "0") + eps2 = sympy.Float(1.0e-36) def old_guard(e): diff --git a/tests/test_0024_jit_graph_lowering.py b/tests/test_0024_jit_graph_lowering.py new file mode 100644 index 000000000..e6bf785c8 --- /dev/null +++ b/tests/test_0024_jit_graph_lowering.py @@ -0,0 +1,267 @@ +"""The JIT's graph lowering (#823, tier 2): nodes, their derivatives, and the C emitted +from them. + +Each non-constant ``UWexpression`` becomes a node, an applied function of the leaves its +value depends on, and SymPy's chain rule differentiates through it by ``fdiff``. Each +test closes one risk of the design note +(``docs/developer/design/jit-shared-graph-codegen.md``), against the expanded tree as +the reference. +""" +import numpy as np +import pytest +import sympy +from sympy.core.function import AppliedUndef +from sympy.vector.scalar import BaseScalar + +import underworld3 as uw +import underworld3.function.expressions as ex +from underworld3.function import diff_wrt_field +from underworld3.utilities import _jit_graph as jg + +pytestmark = [pytest.mark.level_1, pytest.mark.tier_b] + + +@pytest.fixture(scope="module") +def box(): + uw.reset_default_model() + mesh = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) + T = uw.discretisation.MeshVariable("T0024", mesh, 1, degree=1) + v = uw.discretisation.MeshVariable("V0024", mesh, 2, degree=2) + return mesh, T, v + + +def _numbers(exprs, seed=0, rest=()): + """Every leaf of the expanded ``exprs`` (field values and gradients, coordinates) + as a number, the same number wherever it occurs; leaves in ``rest`` are zero.""" + rng = np.random.default_rng(seed) + leaves = set() + for e in exprs: + leaves |= e.atoms(AppliedUndef) + leaves |= {s for s in e.free_symbols if isinstance(s, BaseScalar)} + leaves = sorted(leaves, key=sympy.default_sort_key) + return {s: sympy.Float(0.0 if s in rest else rng.uniform(0.2, 1.5)) for s in leaves} + + +def _value(e, numbers): + e = ex.unwrap_expression(e, mode="nondimensional") + return complex(e.xreplace(numbers).evalf()) + + +def test_a_derivative_through_nodes_is_the_derivative_of_the_tree(box): + """By field value, by field gradient and by coordinate, through + ``diff_wrt_field`` and through ``sympy.diff`` (which swaps the variable for a + ``Dummy`` before ``fdiff`` sees it: a derivative keyed by the leaf OBJECT would be + zero). The coordinate case holds a field and a coordinate in one node: a slot + derivative taken as a total derivative would count the field's gradient twice.""" + mesh, T, v = box + x, y = mesh.N.x, mesh.N.y + u, ux = T.sym[0], T.sym[0].diff(x) + k = uw.expression(r"k_{0024a}", 0.7, "constant") + # by field value and gradient + a = uw.expression(r"a_{0024a}", u ** 2 + u * ux, "inner") + b = uw.expression(r"b_{0024a}", sympy.exp(k * a) + a * u + sympy.sqrt(a + ux ** 2), + "outer") + # by coordinate: a field and a coordinate in one node (a gradient leaf would need + # a second derivative of the field, which Underworld refuses on either route) + c = uw.expression(r"c_{0024a}", u ** 2 + x * u + y, "inner, with coordinates") + d = uw.expression(r"d_{0024a}", sympy.exp(k * c) + c * u * x, "outer, with coordinates") + for law, wrt in ((b * u + a, (u, ux)), (d * u + c, (u, x, y))): + graph = jg.KernelGraph() + lowered = graph.lower(law) + tree = ex.unwrap_expression(law, mode="symbolic_keep_constants") + assert lowered.atoms(jg._KernelNode), "nothing was lowered to a node" + for w in wrt: + d_tree = diff_wrt_field(tree, w) + n = _numbers([d_tree, tree]) + reference = _value(d_tree, n) + assert abs(reference) > 1.0e-3, w + for dg in (diff_wrt_field(lowered, w), sympy.diff(lowered, w)): + assert abs(_value(dg, n) - reference) <= 1.0e-12 * abs(reference), (w, dg) + + +def test_second_derivatives_and_a_constant_sensitivity_pass_through_nodes(box): + mesh, T, v = box + u = T.sym[0] + c = uw.expression(r"c_{0024b}", 1.3, "constant") + a = uw.expression(r"a_{0024b}", c * u ** 3 + sympy.log(u + c), "inner") + law = a ** 2 + lowered = jg.KernelGraph().lower(law) + tree = ex.unwrap_expression(law, mode="symbolic_keep_constants") + for d_graph, d_tree in ( + (diff_wrt_field(diff_wrt_field(lowered, u), u), + diff_wrt_field(diff_wrt_field(tree, u), u)), + (sympy.diff(lowered, c), sympy.diff(tree, c))): + n = _numbers([d_tree]) + reference = _value(d_tree, n) + assert abs(reference) > 1.0e-3 + assert abs(_value(d_graph, n) - reference) <= 1.0e-12 * abs(reference) + + +def test_two_constants_with_one_name_stay_two(box): + """Two viscosities are both called eta (one per material). Nodes are identified by + body, and SymPy's equality tells the two apart, so each keeps its own slot.""" + mesh, T, v = box + from underworld3.utilities._jitextension import _manifest_from + + vv = uw.discretisation.MeshVariable("V0024c", mesh, 2, degree=2) + p = uw.discretisation.MeshVariable("P0024c", mesh, 1, degree=1) + stokes = uw.systems.Stokes(mesh, velocityField=vv, pressureField=p) + lower = uw.constitutive_models.ViscousFlowModel(stokes.Unknowns, material_name="lo") + lower.Parameters.shear_viscosity_0 = 1.0 + upper = uw.constitutive_models.ViscousFlowModel(stokes.Unknowns, material_name="up") + upper.Parameters.shear_viscosity_0 = 1000.0 + a = lower.Parameters.shear_viscosity_0 + b = upper.Parameters.shear_viscosity_0 + assert a != b and a.name == b.name + u = T.sym[0] + ea = uw.expression(r"e_{0024c}", a * u ** 2, "first material") + eb = uw.expression(r"e_{0024c}", b * u ** 2, "second material", + _unique_name_generation=True) + lowered = jg.lower_callbacks([ea + 2 * eb], mesh) + manifest, subs = _manifest_from(jg.constant_leaves(lowered)) + assert len(manifest) == 2 + nodes = lowered[0].atoms(jg._KernelNode) + assert len({n.func for n in nodes}) == 2, nodes + + +def test_a_changed_body_is_lowered_afresh_without_clearing_the_cache(box): + """Each lowering's node classes carry its serial in their SymPy identity, so + SymPy's cache cannot hand back a node whose body is an earlier version. The hard + case: two constants share a display name and the re-declared atom swaps their + roles, so the new node has the old one's name AND arguments.""" + mesh, T, v = box + u = T.sym[0] + c1 = uw.expression(r"c_{0024d}", 2.0, "first") + c2 = uw.expression(r"c_{0024d}", 5.0, "second", _unique_name_generation=True) + assert c1 != c2 and c1.name == c2.name + a = uw.expression(r"a_{0024d}", c1 * u ** 3 + c2 * u, "re-declared") + first = jg.KernelGraph().lower(a) + a.sym = c2 * u ** 3 + c1 * u + second = jg.KernelGraph().lower(a) + assert first.func.__name__ == second.func.__name__ # the hard case holds + assert first.args == second.args + assert jg.expand_nodes(second) == c2 * u ** 3 + c1 * u + assert jg.expand_nodes(diff_wrt_field(second, u)) == 3 * c2 * u ** 2 + c1 + + +def test_the_guarded_lowering_is_the_guarded_tree(box, monkeypatch): + """The Newton source on the graph (the sqrt guard in each node body) evaluates to + the guarded tree, and so does its tangent, including at a state of rest where the + unguarded tangent is 0/0.""" + from underworld3.cython.generic_solvers import _jacobian_unwrap + + mesh, T, v = box + vv = uw.discretisation.MeshVariable("V0024e", mesh, 2, degree=2) + p = uw.discretisation.MeshVariable("P0024e", mesh, 1, degree=1) + stokes = uw.systems.Stokes(mesh, velocityField=vv, pressureField=p) + stokes.constitutive_model = uw.constitutive_models.ViscoPlasticFlowModel + cm = stokes.constitutive_model + cm.Parameters.shear_viscosity_0 = 1.0 + cm.Parameters.yield_stress = uw.expression(r"C_{0024e}", 0.5) + 0.3 * p.sym[0] + cm.Parameters.yield_stress_min = uw.expression(r"\tau_{0024e}", 0.01) + flux = sympy.Matrix(stokes.F1.sym) + + monkeypatch.setenv("UW_JIT_GRAPH", "0") + tree = _jacobian_unwrap(flux) + monkeypatch.setenv("UW_JIT_GRAPH", "1") + graph = _jacobian_unwrap(flux) + assert graph.atoms(jg._KernelNode), "the Newton source was not lowered" + + L = stokes.Unknowns.L + rest = {L[i, j] for i in range(2) for j in range(2)} + for state in (dict(seed=1), dict(seed=2, rest=rest)): + for i in range(2): + for j in range(2): + d_tree = diff_wrt_field(tree[i, j], L[0, 1]) + d_graph = diff_wrt_field(graph[i, j], L[0, 1]) + n = _numbers([tree[i, j], d_tree], **state) + for t, g in ((tree[i, j], graph[i, j]), (d_tree, d_graph)): + vt, vg = _value(t, n), _value(g, n) + assert np.isfinite(vt) and np.isfinite(vg), (state, t) + assert abs(vg - vt) <= 1.0e-12 * max(abs(vt), 1.0), (state, i, j) + + +def _header(solver, monkeypatch, route): + """The generated header of ``solver`` on ``route``, with the module name and the + symbol prefix canonicalised as ``getext`` canonicalises them.""" + import underworld3.utilities._jitextension as jx + + seen = {} + generate = jx.generate_c_source + + def keep(*a, **k): + modname, codeguys, diag = generate(*a, **k) + seen["h"] = dict(codeguys)["cy_ext.h"].replace(modname, "M").replace( + diag["randstr"], "R") + return modname, codeguys, diag + + monkeypatch.setattr(jx, "generate_c_source", keep) + monkeypatch.setenv("UW_JIT_GRAPH", route) + solver.is_setup = False + solver._setup_pointwise_functions() + return seen["h"] + + +def _poisson(name, mesh): + u = uw.discretisation.MeshVariable("U" + name, mesh, 1, degree=1) + pois = uw.systems.Poisson(mesh, u_Field=u) + pois.constitutive_model = uw.constitutive_models.DiffusionModel + k0 = uw.expression(r"k_{0024f}", 2.0, "conductivity") + g = uw.expression(r"g_{0024f}", sympy.sqrt(1 + u.sym[0].diff(mesh.N.x) ** 2), "slope") + pois.constitutive_model.Parameters.diffusivity = k0 * g + u.sym[0] ** 2 + pois.f = 1.0 + pois.consistent_jacobian = True + return pois + + +def test_the_emitted_source_does_not_depend_on_what_came_before(monkeypatch): + """The C is a function of the mathematics and the data layout: the same law, + declared after a preamble of unrelated objects (every creation counter shifted) + and declared twice, emits the same header.""" + uw.reset_default_model() + mesh = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) + first = _header(_poisson("0024f", mesh), monkeypatch, "1") + for k in range(7): + uw.expression(rf"junk_{{0024f,{k}}}", float(k), "preamble") + other = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) + uw.discretisation.MeshVariable("J0024f", other, 1, degree=1) + mesh2 = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) + second = _header(_poisson("0024f", mesh2), monkeypatch, "1") + assert first == second + assert "uwt_0" in first, "the kernel has no temporaries: nothing was lowered" + + +def test_a_law_with_no_named_quantity_emits_the_tree_route_source(monkeypatch): + uw.reset_default_model() + mesh = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) + u = uw.discretisation.MeshVariable("U0024g", mesh, 1, degree=1) + pois = uw.systems.Poisson(mesh, u_Field=u) + pois.constitutive_model = uw.constitutive_models.DiffusionModel + pois.constitutive_model.Parameters.diffusivity = 1.0 + pois.f = 2.0 + assert _header(pois, monkeypatch, "0") == _header(pois, monkeypatch, "1") + + +def test_the_manifest_from_the_leaves_is_the_scanned_manifest(box): + """The constants the lowered kernels read are the constants the expressions hold, + in the same slots.""" + from underworld3.utilities._jitextension import _extract_constants, _manifest_from + + mesh, T, v = box + u = T.sym[0] + c1 = uw.expression(r"c_{0024h,1}", 0.5, "constant") + c2 = uw.expression(r"c_{0024h,2}", 2.0, "constant") + c3 = uw.expression(r"c_{0024h,3}", 3.0, "constant in a constant") + c4 = uw.expression(r"c_{0024h,4}", c3 * 2, "constant of a constant") + a = uw.expression(r"a_{0024h}", c1 * u + c4 * u ** 2, "inner") + b = uw.expression(r"b_{0024h}", a / (c2 + a ** 2), "outer") + fns = (b, sympy.Matrix([[b * u, a]])) + scanned, _ = _extract_constants(fns, mesh) + leaves, _ = _manifest_from(jg.constant_leaves(jg.lower_callbacks(fns, mesh))) + assert [e for _, e in leaves] == [e for _, e in scanned] + assert len(scanned) == 3 # c4 is a slot of its own; c3 is folded into it From 3e370d6f538d228c866585eb53bf1dee8b1e0faa Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 08:42:41 +1100 Subject: [PATCH 03/14] Graph route: conditions stay inline, coordinates spelled from their index (#823) Two defects the serial suite found on the graph route, both in the fault-network end-to-end tests (test_0850, test_0851): - the per-body common sub-expression split made a repeated Piecewise condition into a node, a value, which Piecewise refuses as a condition. A body that is not an Expr now stays inline (test_0024::test_a_repeated_condition_stays_a_condition, shown failing first); - a coordinate leaf can be a UWCoordinate that SymPy's cache returned for its equal base scalar, without the C name the mesh set. The tree route recovers the name before printing; the graph route spells leaves earlier, during emission, so the speller now recovers it from the coordinate's index and system. Design note: the measurements in the library against tier 1 (setup, C size, assembly, Newton iterations, operator agreement at three states, hash seeds, the plain sympy.diff probe), and the code-size estimate corrected: about even, not 530 lines out for 360 in. route_ab.py gains a body-force perturbation for the notch's iteration spread. Underworld development team with AI support from Claude Code --- .../design/jit-shared-graph-codegen.md | 145 +++++++++++++----- scripts/sessions/jit_graph/route_ab.py | 7 +- src/underworld3/utilities/_jit_graph.py | 5 +- src/underworld3/utilities/_jitextension.py | 16 +- tests/test_0024_jit_graph_lowering.py | 21 +++ 5 files changed, 156 insertions(+), 38 deletions(-) diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index df109a34b..488585021 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -1,8 +1,11 @@ # Generating JIT kernels from the shared expression graph -**Status**: Proposed, 2026-10-07. Tier 2 of [#823](https://github.com/underworldcode/underworld3/issues/823). -Tier 1 ([#830](https://github.com/underworldcode/underworld3/pull/830), -`bugfix/jit-setup-cost-823`) is the fix for today and lands first: the memoised +**Status**: Proposed, 2026-10-07. Staging steps 2 and 3 implemented behind the private +switch `UW_JIT_GRAPH=1` on `feature/jit-graph-codegen`, 2026-10-08, and measured against +tier 1 end to end ({ref}`jit-graph-in-the-library`). +Tier 2 of [#823](https://github.com/underworldcode/underworld3/issues/823). +Tier 1 ([#830](https://github.com/underworldcode/underworld3/pull/830), merged +2026-10-07) is the fix for today: the memoised unwrap, realness for field values, coordinates and constant slots, derivatives with respect to fields through `uw.function.diff_wrt_field`, and the power-mean sharpness as an atom of its own. Tier 2 is the design we would have chosen from the start, and tier @@ -220,14 +223,17 @@ which SymPy does not merge, and on a power law over a named invariant, with $n$ constant atom or as the number 3, both routes give a finite Newton tangent at a state of rest, equal entry for entry. -### One lowering per setup, read back by `getext()` +### Each lowering is read back by `getext()` from its nodes -`_setup_pointwise_functions` creates one lowering context, and `_jacobian_unwrap` builds -its nodes in it. Each node class holds its body and its context, so `getext()` reads -every body from the nodes it is handed and needs no new argument. `getext()` lowers the -remaining atoms — those of the residual and the Picard blocks — in the same way; a -caller that is not a solver (`Integral`, `CellWiseIntegral`, `BdIntegral`) gets a -context of its own. +Each call of `_jacobian_unwrap` lowers its source in a context of its own, and +`getext()` lowers the remaining atoms — those of the residual and the Picard blocks — in +another. Each node class holds its body and its context, so `getext()` reads every body +from the nodes it is handed and needs no new argument, and a caller that is not a solver +(`Integral`, `CellWiseIntegral`, `BdIntegral`) needs no change. An atom lowered in two +contexts is two nodes in Python and one temporary in C, because emission merges +temporaries by the C they compute. The solvers are therefore unchanged apart from +`_jacobian_unwrap`; one context per setup would only save the constancy test of an atom +that both lowerings meet. `getext()` no longer runs `unwrap_expression` over a whole kernel. Its present phases — reveal the constants, substitute them, unwrap the rest — become: lower atoms to nodes, @@ -372,20 +378,82 @@ cancellation leaves rounding noise and the tangent is finite but meaningless. Th singular there. The notch cannot reach it — its material fraction is a P0 field holding exactly 0 or 1 — but a model with a projected or higher-degree material field could. +(jit-graph-in-the-library)= +## In the library, against tier 1 + +Steps 2 and 3 are implemented on `feature/jit-graph-codegen` +(`src/underworld3/utilities/_jit_graph.py`, with branches in `getext()`, +`generate_c_source()`, `_jacobian_unwrap` and the two unwrappers), selected by +`UW_JIT_GRAPH=1`. Unset, every path is tier 1's. Both routes therefore run on one build, +and each fixture is run once per route in a fresh process with the JIT cache off +(`scripts/sessions/jit_graph/route_ab.py`): the solver's whole pointwise setup, a Newton +solve, then the residual and the Jacobian assembled repeatedly at the solution. The +fixtures are those of the prototype, with the box ones solved as a lid-driven box +started from simple shear; the notch is solved from rest to a relative tolerance of +$10^{-6}$. Apple clang `-O3`, macOS, single runs on a machine shared with another +session's seven solver processes, so the times are indicative. Tree first, graph second: + +| | notch | VEP | box | TI VEP | power law | linear | +|---|---|---|---|---|---|---| +| pointwise setup, s | 15.9 / 2.5 | 3.4 / 2.3 | 2.5 / 1.8 | 1.5 / 1.5 | 1.5 / 1.3 | 1.2 / 1.0 | +| of which C generation, s | 8.6 / 0.23 | 0.61 / 0.21 | 0.33 / 0.10 | 0.17 / 0.14 | 0.05 / 0.04 | 0.02 / 0.01 | +| generated C, all modules of the solve | 4.0 MB / 22 KB | 380 KB / 40 KB | 104 KB / 16 KB | 42 KB / 31 KB | 28 KB / 14 KB | 10.7 KB, byte-identical | +| Jacobian assembly, ms | 180 / 132 | 2.40 / 2.19 | 1.58 / 1.45 | 2.23 / 2.22 | 1.48 / 1.40 | 1.39 / 1.44 | +| residual assembly, ms | 21.4 / 18.7 | 0.62 / 0.59 | 0.35 / 0.33 | 0.61 / 0.60 | 0.32 / 0.32 | 0.33 / 0.33 | +| Newton iterations, nonlinear / linear | 75 / 496 and 57 / 386 | 6 / 6, both | 30 / 30, both (limit) | 2 / 2, both | 16 / 16, both | 1 / 1, both | + +The compile time is not shown separately: it counts every module the solve builds (the +VEP's history projections among them), and on these fixtures it is 1–4 s on either +route, never slower on the graph. + +**The operators agree to round-off.** `route_assemble.py` assembles the residual and the +Jacobian on each route at the same state — rest (zero velocity, boundary values +imposed), the tree's final state and the graph's final state — and compares them entry +by entry. On every fixture, at every state, and for the Newton, Picard and continuation +tangents, no Jacobian entry differs from the other route's by more than $5.1\times10^{-16}$ +of the largest entry in its row, and the sparsity patterns are identical. Residuals +agree to $10^{-14}$ of their largest entry, except at converged states, where the +residual is itself round-off ($10^{-13}$). The uncapped power law is singular at rest: +both routes give a non-finite residual there, in the same 126 entries. + +**The solves agree, and the notch's iteration count is not a property of the route.** On +every fixture but the notch, both routes take the same path: equal nonlinear and linear +iteration counts, residual histories equal to $10^{-5}$ and solutions to $10^{-11}$ or +better. On the notch both converge, the tree in 75 Newton iterations and the graph in +57, to solutions that agree to $8\times10^{-7}$ of the largest velocity and +$2\times10^{-4}$ of the largest pressure, consistent with the tolerance. The two +residual histories agree to $10^{-9}$ for four iterations and part from the fifth: the +notch's Newton iteration amplifies differences at the level of round-off, and the +operators that drive it agree at that level at three states. + +**The source is canonical.** On the notch, box and VEP fixtures the graph route's C is +byte-identical under `PYTHONHASHSEED` 0, 1 and 2; `test_0024` checks that a law +declared after a preamble of unrelated objects and declared twice emits the same +header. + +**The stand-in derivative is no longer load-bearing.** `plain_diff_probe.py` replaces +`diff_wrt_field` and `derive_by_array_wrt_field` in the solvers by plain `sympy.diff` +and `sympy.derive_by_array` and runs the Newton solve that motivated them (the +Drucker–Prager yield with a yield-stress floor at softness 0). The tree route fails in +the C printer, on the `Derivative` SymPy leaves beside the `Abs`; the graph route +converges, because every `Abs` sits in a node body, whose partials are taken against +real dummies. + ## Against tier 1: no slower, more robust, less code Tier 1 is the competitor. Each criterion is measured against it, not against the code before it. -**No slower.** On every fixture so far, every stage is as fast or faster: lowering, -differentiating and emitting the notch's Newton block take 0.5 s against 9 s, the C is -6.6 KB against 3.2 MB, and the kernel runs in about an eighth of the time. On a -constant-viscosity Stokes the two routes emit byte-identical C, because a constant law -has no node to lower. On a small power law the graph spends 0.02 s more differentiating -(0.09 s against 0.07 s), and the run time of its small kernels is equal within the -noise of a loaded machine. The claim needs the benchmark plan's measurement: an idle -machine, the minimum of repeated runs, the library's whole setup end to end, on Linux -as well as macOS. +**No slower.** In the library, end to end, every stage on every fixture is as fast or +faster: the notch's pointwise setup takes 2.5 s against 15.9 s, its C is 22 KB against +4.0 MB, and its Jacobian assembles in 132 ms against 180 ms. A constant-viscosity Stokes +emits byte-identical C, because a constant law has no node to lower. The assembly gain +is smaller than the prototype's kernel timings suggested (an eighth of the time per +call) because the pointwise kernel is one part of the assembly, beside quadrature and +the element loop. Still to measure: an idle machine with repeated runs, Linux with gcc +(whose default `-fmath-errno` keeps it from merging repeated `pow`, `exp` and `sqrt` +calls, which the graph's temporaries do for every named quantity), a three-dimensional fixture, +and the small kernels of `Integral` and `BdIntegral`. **More robust.** Each failure class below needed a patch in tier 1, placed where it surfaced; in the graph it cannot arise, or arises only in one body: @@ -393,26 +461,31 @@ surfaced; in the graph it cannot arise, or arises only in one body: | failure class | tier 1 | graph | shown | |---|---|---|---| | SymPy's automatic algebra on an expanded base is slow (`im()` inside `Pow.__new__`, 40 s on the notch) | a constant atom added to the power-mean law; realness for unit-carrying parameters | bases are bodies, not trees | graph lowering took 0.04 s on every tier 1 commit, including those where the library took 40 s | -| realness turns `sqrt(x**2)` into `Abs`, whose derivative with respect to a field SymPy leaves unevaluated | a stand-in derivative at 49 call sites | body partials are taken against real dummies | the graph differentiated the Drucker–Prager floor law cleanly with plain `sympy.diff` where the library failed | +| realness turns `sqrt(x**2)` into `Abs`, whose derivative with respect to a field SymPy leaves unevaluated | a stand-in derivative at 49 call sites | body partials are taken against real dummies | with plain `sympy.diff` at every solver site, the Drucker–Prager floor law fails to compile on the tree route and solves on the graph route | | manifest and C built by two walks (#302) | two consistency guards and `_reveal_constants` | one walk | by construction; manifests identical on all fixtures | -| generated C that differs between ranks (#752, open) | rank 0's source adopted | canonical emission | byte-identical under hash seeds, preambles and re-declaration; ranks not yet tested | -| field symbols given their C names by mutating their classes, in an order that matters | `ccode_patch_fns`, the coordinate recovery block | an explicit map from leaf to C | by construction | -| generated C too large to read or to compile (#547) | opt-in CSE, lower optimisation flags | one line per named quantity | 6.6 KB against 3.2 MB on the notch | +| generated C that differs between ranks (#752, open) | rank 0's source adopted | canonical emission | byte-identical under hash seeds 0–2 (notch, box, VEP), after a preamble and when re-declared (`test_0024`); ranks not yet tested | +| field symbols given their C names by mutating their classes, in an order that matters | `ccode_patch_fns`, the coordinate recovery block | an explicit map from leaf to C | not yet: steps 2–3 still spell leaves through the patched printer; the map belongs to step 4 | +| generated C too large to read or to compile (#547) | opt-in CSE, lower optimisation flags | one line per named quantity | 22 KB against 4.0 MB for the notch's whole solve | The graph brings one failure class tier 1 does not have: it cannot cancel a quantity against its own reciprocal across a name. -**Less code, less global state.** By line ranges in `_jitextension.py`, the graph -replaces about 530 lines — `_reveal_constants`, `_extract_constants` and -`_collect_constant_atoms`, the two consistency guards, the global `_ccode` patching and -the coordinate recovery, the expanded-tree lowering, the opt-in CSE path, the -identity-walk scans that exist for large trees, and the dead `prepare_for_cache_key` and -`_createext` — with about 360: one lowering module, a leaf-to-C map, a manifest built -from the leaves, and node expansion in two unwrappers. Several tier 1 measures become -belt-and-braces rather than load-bearing: the stand-in derivative at every call site, -realness for unit-carrying parameters, and the patch to SymPy's private -`BaseScalar._prop_handler` table, the last of which we would want to remove. These -counts are estimates until the change exists; the PR that makes it shows them. +**Less global state, not less code.** The earlier estimate here (530 lines replaced by +360) does not survive the implementation. The lowering module is 442 lines, about a +third of them docstrings, and the branches it needs elsewhere about 60. Step 4 deletes +about 390 lines of the tree route: `_reveal_constants`, the scanning half of +`_extract_constants`, `_collect_constant_atoms`, `_xreplace_shared`, `_unique_symbols`, +the tree lowering and its two consistency guards in `generate_c_source`, the coordinate +recovery, the opt-in CSE path, the tree guard in `_jacobian_unwrap`, and the unused +`prepare_for_cache_key` and `_createext`. Replacing the class patching of +`ccode_patch_fns` (148 lines with its comments) by a map from leaf to C would remove +about 100 more. The line count comes out about even. What changes is where the +correctness lives: in one module whose every risk has a test that fails when the +mechanism is broken (`test_0024`, each test checked against a mutation of the lowering), +instead of in patches at the places each failure surfaced. The stand-in derivative at +49 call sites becomes belt-and-braces, as the probe above shows; so would realness for +unit-carrying parameters and the patch to SymPy's private `BaseScalar._prop_handler` +table, which we would want to remove — neither is yet shown. ## Results agree to round-off, not bit for bit @@ -531,7 +604,9 @@ repository; the notch is the one exception, and its driver is named. expanded and print as before. Nodes exist only inside `getext()`. 3. `_jacobian_unwrap` builds nodes instead of expanding, so nodes reach the solver's blocks; in the same change the two unwrappers learn to expand them and `test_0022`'s - two tree-shaped tests are rewritten. + tree-shaped guard test is pinned to the tree route. + + Steps 2 and 3 are implemented together behind `UW_JIT_GRAPH=1` (2026-10-08). 4. The expanded route and the switch are removed, together with the two functions that mirror it and have no callers (`prepare_for_cache_key`, `_createext`), and `jit-cache.md`, `expressions-functions.md` and `jacobian-consistent-tangent.md` are diff --git a/scripts/sessions/jit_graph/route_ab.py b/scripts/sessions/jit_graph/route_ab.py index 4fca2c91b..5b97a1875 100644 --- a/scripts/sessions/jit_graph/route_ab.py +++ b/scripts/sessions/jit_graph/route_ab.py @@ -29,6 +29,9 @@ uw_tol=uw.Param(1.0e-8, description="SNES relative tolerance"), uw_assemble=uw.Param(20, description="residual and Jacobian assemblies to time"), uw_out=uw.Param("~/+Simulations/jit_graph/tier2", description="output directory"), + uw_force_scale=uw.Param(0.0, description="relative change of the body force, to " + "measure how the Newton path responds to round-off-sized changes"), + uw_label=uw.Param("", description="suffix for the output file"), ) route = "graph" if _jit_graph.enabled() else "tree" fixture = str(params.uw_fixture) @@ -38,6 +41,8 @@ stokes.consistent_jacobian = {"newton": True, "picard": False, "continuation": "continuation"}[tangent] solve_kwargs = fixtures.prepare_solve(stokes, fixture) +if float(params.uw_force_scale): + stokes.bodyforce = stokes.bodyforce * (1 + float(params.uw_force_scale)) stokes.tolerance = float(params.uw_tol) stokes.petsc_options["snes_max_it"] = int(params.uw_maxit) if "ksp_monitor" in stokes.petsc_options: @@ -113,7 +118,7 @@ def timed_pointwise(self, *a, **k): out = os.path.expanduser(str(params.uw_out)) os.makedirs(out, exist_ok=True) -np.savez(os.path.join(out, f"{fixture}_{tangent}_{route}.npz"), +np.savez(os.path.join(out, f"{fixture}_{tangent}{params.uw_label}_{route}.npz"), v=np.asarray(stokes.u.array), p=np.asarray(stokes.p.array), X=np.asarray(X.array).copy(), history=np.asarray(r.history, dtype=float), diff --git a/src/underworld3/utilities/_jit_graph.py b/src/underworld3/utilities/_jit_graph.py index d9cb624d8..b97064e98 100644 --- a/src/underworld3/utilities/_jit_graph.py +++ b/src/underworld3/utilities/_jit_graph.py @@ -279,10 +279,13 @@ def _display_name(self, body): def make_node(self, body): """The node for ``body``, one per distinct body. A body that is a number, a - single leaf, a single node, or reads no leaf is returned as itself.""" + single leaf, a single node, or reads no leaf is returned as itself, and so is a + condition (a node is a value; a Piecewise refuses one as its condition).""" global _nodes_made body = sympy.sympify(body) + if not isinstance(body, sympy.Expr): + return body if not self._splitting: body = self._split_shared(body) if body.is_Atom or isinstance(body, (AppliedUndef, BaseScalar)): diff --git a/src/underworld3/utilities/_jitextension.py b/src/underworld3/utilities/_jitextension.py index 05a3389d2..0c7baee77 100644 --- a/src/underworld3/utilities/_jitextension.py +++ b/src/underworld3/utilities/_jitextension.py @@ -1377,7 +1377,21 @@ def _handle_UnevaluatedExpr(expr): def _spell(leaf): # the C a kernel reads for a leaf of the graph placeholder = constants_subs_map.get(leaf) if constants_subs_map else None - return placeholder._ccodestr if placeholder is not None else printer.doprint(leaf) + if placeholder is not None: + return placeholder._ccodestr + if isinstance(leaf, sympy.vector.scalar.BaseScalar): + # a coordinate may be a fresh instance, or a UWCoordinate SymPy's cache + # handed back for its equal base scalar, without the name the mesh set: + # recover it from the coordinate's index and system, as the coordinate + # recovery below does for the printer + try: + return leaf._ccodestr + except AttributeError: + idx, system = leaf._id[0], str(leaf._id[1]) + leaf._ccodestr = (f"petsc_n[{idx}]" if "Gamma" in system + else f"petsc_x[{idx}]") + return leaf._ccodestr + return printer.doprint(leaf) eqns = [] for index, fn in enumerate(fns): diff --git a/tests/test_0024_jit_graph_lowering.py b/tests/test_0024_jit_graph_lowering.py index e6bf785c8..ff59cb312 100644 --- a/tests/test_0024_jit_graph_lowering.py +++ b/tests/test_0024_jit_graph_lowering.py @@ -265,3 +265,24 @@ def test_the_manifest_from_the_leaves_is_the_scanned_manifest(box): leaves, _ = _manifest_from(jg.constant_leaves(jg.lower_callbacks(fns, mesh))) assert [e for _, e in leaves] == [e for _, e in scanned] assert len(scanned) == 3 # c4 is a slot of its own; c3 is folded into it + + +def test_a_repeated_condition_stays_a_condition(box): + """A condition that repeats inside one body (two Piecewise on the same test) is + shared by the body's common sub-expression split; it must stay a Boolean, not + become a node, which Piecewise refuses as a condition (the fault-network laws, + test_0850 and test_0851).""" + mesh, T, v = box + x = mesh.N.x + u = T.sym[0] + a = uw.expression(r"a_{0024i}", + sympy.Piecewise((u, x > 0.5), (u ** 2, True)) + + sympy.Piecewise((2 * u, x > 0.5), (u ** 3, True)), + "two branches on one test") + lowered = jg.KernelGraph().lower(a * u) + tree = ex.unwrap_expression(a * u, mode="symbolic_keep_constants") + for xv in (0.2, 0.8): + n = _numbers([tree, diff_wrt_field(tree, u)]) + n[x] = sympy.Float(xv) + for t, g in ((tree, lowered), (diff_wrt_field(tree, u), diff_wrt_field(lowered, u))): + assert abs(_value(g, n) - _value(t, n)) <= 1.0e-12 * abs(_value(t, n)) From 890606e36d3cbcf9377d3f296f8189dc079a47a5 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 09:53:11 +1100 Subject: [PATCH 04/14] Design note: suite, iteration spread and #752 results for the graph route (#823) The serial suite passes on the graph route (3,023 passed). The notch's Newton iteration count over 17 round-off-sized perturbations per route: tree 60-111 and one non-convergence, graph 43-86; the same distribution within the sample. The #752 fixture no longer disagrees across ranks on either route (0 of 10 at np = 2), so it cannot test canonical emission's effect on #752. Underworld development team with AI support from Claude Code --- .../design/jit-shared-graph-codegen.md | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index 488585021..882b889b9 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -400,7 +400,7 @@ session's seven solver processes, so the times are indicative. Tree first, graph | generated C, all modules of the solve | 4.0 MB / 22 KB | 380 KB / 40 KB | 104 KB / 16 KB | 42 KB / 31 KB | 28 KB / 14 KB | 10.7 KB, byte-identical | | Jacobian assembly, ms | 180 / 132 | 2.40 / 2.19 | 1.58 / 1.45 | 2.23 / 2.22 | 1.48 / 1.40 | 1.39 / 1.44 | | residual assembly, ms | 21.4 / 18.7 | 0.62 / 0.59 | 0.35 / 0.33 | 0.61 / 0.60 | 0.32 / 0.32 | 0.33 / 0.33 | -| Newton iterations, nonlinear / linear | 75 / 496 and 57 / 386 | 6 / 6, both | 30 / 30, both (limit) | 2 / 2, both | 16 / 16, both | 1 / 1, both | +| Newton iterations, nonlinear / linear | 75 / 496 and 57 / 386 (see below) | 6 / 6, both | 30 / 30, both (limit) | 2 / 2, both | 16 / 16, both | 1 / 1, both | The compile time is not shown separately: it counts every module the solve builds (the VEP's history projections among them), and on these fixtures it is 1–4 s on either @@ -424,7 +424,18 @@ better. On the notch both converge, the tree in 75 Newton iterations and the gra $2\times10^{-4}$ of the largest pressure, consistent with the tolerance. The two residual histories agree to $10^{-9}$ for four iterations and part from the fifth: the notch's Newton iteration amplifies differences at the level of round-off, and the -operators that drive it agree at that level at three states. +operators that drive it agree at that level at three states. Changing the body force by +$\pm1$ to $\pm7\times10^{-12}$ of itself shows how far: over 17 such solves per route, +the tree converged in 60 to 111 iterations (median 75) and once failed to converge in +300, and the graph converged every time, in 43 to 86 (median 69). The two distributions +are the same within the sample; the notch's iteration count is a property of the +problem, not of the route. + +**The test suite passes on the graph route.** Every serial batch of `scripts/test.sh` +(tier A and B, levels 1 to 3) run with `UW_JIT_GRAPH=1`: 3,023 passed, none failed. The +first run found two defects, both in the fault-network laws and both fixed with a test: +a repeated Piecewise condition shared as a node (a value, which Piecewise refuses as a +condition), and a coordinate leaf without the C name the mesh sets. **The source is canonical.** On the notch, box and VEP fixtures the graph route's C is byte-identical under `PYTHONHASHSEED` 0, 1 and 2; `test_0024` checks that a law @@ -463,7 +474,7 @@ surfaced; in the graph it cannot arise, or arises only in one body: | SymPy's automatic algebra on an expanded base is slow (`im()` inside `Pow.__new__`, 40 s on the notch) | a constant atom added to the power-mean law; realness for unit-carrying parameters | bases are bodies, not trees | graph lowering took 0.04 s on every tier 1 commit, including those where the library took 40 s | | realness turns `sqrt(x**2)` into `Abs`, whose derivative with respect to a field SymPy leaves unevaluated | a stand-in derivative at 49 call sites | body partials are taken against real dummies | with plain `sympy.diff` at every solver site, the Drucker–Prager floor law fails to compile on the tree route and solves on the graph route | | manifest and C built by two walks (#302) | two consistency guards and `_reveal_constants` | one walk | by construction; manifests identical on all fixtures | -| generated C that differs between ranks (#752, open) | rank 0's source adopted | canonical emission | byte-identical under hash seeds 0–2 (notch, box, VEP), after a preamble and when re-declared (`test_0024`); ranks not yet tested | +| generated C that differs between ranks (#752, open) | rank 0's source adopted | canonical emission | byte-identical under hash seeds 0–2 (notch, box, VEP), after a preamble and when re-declared (`test_0024`). Across ranks untested: the #752 fixture (`rank_agreement_752.py`, np = 2) disagreed in 0 of 10 runs on either route, against about 2 in 10 before tier 1, so it no longer reproduces the defect | | field symbols given their C names by mutating their classes, in an order that matters | `ccode_patch_fns`, the coordinate recovery block | an explicit map from leaf to C | not yet: steps 2–3 still spell leaves through the patched printer; the map belongs to step 4 | | generated C too large to read or to compile (#547) | opt-in CSE, lower optimisation flags | one line per named quantity | 22 KB against 4.0 MB for the notch's whole solve | From 94d60da9e82e2452d247a38e505550033d87e9b9 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 10:10:07 +1100 Subject: [PATCH 05/14] Design note: the callbacks' share of the notch assembly, per route (#823) assembly_split.py separates the law's callbacks from the finite-element machinery by swapping in a constant viscosity on the same mesh and state. Underworld development team with AI support from Claude Code --- .../design/jit-shared-graph-codegen.md | 11 +++ scripts/sessions/jit_graph/assembly_split.py | 79 +++++++++++++++++++ 2 files changed, 90 insertions(+) create mode 100644 scripts/sessions/jit_graph/assembly_split.py diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index 882b889b9..bd03e7303 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -406,6 +406,17 @@ The compile time is not shown separately: it counts every module the solve build VEP's history projections among them), and on these fixtures it is 1–4 s on either route, never slower on the graph. +**The callbacks run about ten times faster; the assembly is mostly not callbacks.** +`assembly_split.py` times the notch's assembly at a solved state (minimum of 30), then +swaps the law for a constant viscosity on the same mesh, fields and boundary +conditions and times it again; the difference is the cost of the law's callbacks. Per +quadrature point and per Jacobian assembly, which builds the Jacobian and its +preconditioner as two matrices, the callbacks cost 3.1 µs on the tree and 0.29 µs on the +graph, 30% and 4% of the assembly; per residual, 300 ns and 63 ns. The rest of the +assembly, about 150–165 ms here, is the finite-element machinery, the same on both +routes, so on the notch the assembly can gain at most a third. A law with more named +layers, or a three-dimensional `G3` with 81 entries, gives the callbacks a larger share. + **The operators agree to round-off.** `route_assemble.py` assembles the residual and the Jacobian on each route at the same state — rest (zero velocity, boundary values imposed), the tree's final state and the graph's final state — and compares them entry diff --git a/scripts/sessions/jit_graph/assembly_split.py b/scripts/sessions/jit_graph/assembly_split.py new file mode 100644 index 000000000..2aefbbad4 --- /dev/null +++ b/scripts/sessions/jit_graph/assembly_split.py @@ -0,0 +1,79 @@ +"""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 ``UW_JIT_GRAPH``. +""" +import os +import sys +import time + +import numpy as np +import underworld3 as uw +from underworld3.utilities import _jit_graph + +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 _jit_graph.enabled() else "tree" +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): + stokes.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") + +# 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 +floor = timed_assembly("constant viscosity") + +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") From 3f982e6a6a77acc9d27361b31ca7c23a16241ed6 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 11:11:36 +1100 Subject: [PATCH 06/14] assembly_split.py: the constant-viscosity floor drops a viscoelastic fixture's history (#823) The floor pass kept the VEP fixture's stress-history store and timestep, so the solve set dt_elastic on a law without one (found on Hyperion). Underworld development team with AI support from Claude Code --- scripts/sessions/jit_graph/assembly_split.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/scripts/sessions/jit_graph/assembly_split.py b/scripts/sessions/jit_graph/assembly_split.py index 2aefbbad4..386ce44e6 100644 --- a/scripts/sessions/jit_graph/assembly_split.py +++ b/scripts/sessions/jit_graph/assembly_split.py @@ -33,8 +33,8 @@ X_state = np.load(os.path.expanduser(str(params.uw_state)))["X"] -def timed_assembly(label): - stokes.solve(**kwargs) # (re)builds the kernels; no iteration +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 @@ -59,7 +59,7 @@ def timed_assembly(label): return t_res, t_jac -law = timed_assembly("law") +law = timed_assembly("law", kwargs) # quadrature points of the velocity block dm = stokes.mesh.dm @@ -71,7 +71,11 @@ def timed_assembly(label): stokes.constitutive_model = cm stokes.constitutive_model.Parameters.shear_viscosity_0 = 1.0 stokes.is_setup = False -floor = timed_assembly("constant viscosity") +# 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])): From 771bc4db0f2495a3e77c56e2114246a9769d5d52 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 11:17:57 +1100 Subject: [PATCH 07/14] Design note: the Linux/gcc measurements and the -fmath-errno test (#823) Run on Hyperion by another session: same solver paths and ratios as the Mac; gcc's -fmath-errno accounts for about two-thirds of the tree's box callback cost and none of the VEP's; #752 fixture agrees across ranks at np 3 and 4; ptest_jit_cache passes at np 4 on both routes. Underworld development team with AI support from Claude Code --- .../design/jit-shared-graph-codegen.md | 26 ++++++++++++++++--- 1 file changed, 22 insertions(+), 4 deletions(-) diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index bd03e7303..368422810 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -417,6 +417,25 @@ assembly, about 150–165 ms here, is the finite-element machinery, the same on routes, so on the notch the assembly can gain at most a third. A law with more named layers, or a three-dimensional `G3` with 81 entries, gives the callbacks a larger share. +**On Linux with gcc the verdict holds, and gcc's `-fmath-errno` explains part of the +tree's cost on some laws.** On a Linux server (2× Xeon Gold 6240R, the environment's +gcc 14.3 at the default `-O3 -g0`, so with gcc's `-fmath-errno`; the machine fully +loaded by other jobs, a noise floor of about ±10%), the five self-contained fixtures take +identical solver paths on both routes. Setup falls from 10.3 to 8.8 s (box) and from +12.0 to 4.8 s (VEP); the Jacobian assembly ratios, graph to tree, are 0.87 (box), 0.94 +(VEP, power law), 1.00 (TI) and 1.04 (linear, byte-identical C, so noise). Per +quadrature point, the box's Jacobian callbacks cost 8.2 µs on the tree and about 1 µs +on the graph, the noise floor of the subtraction. With `-fno-math-errno`, which lets gcc +treat `sqrt`, `pow` and `exp` as pure and merge repeated calls, the box's tree callbacks +fall to 2.8 µs (reproduced in a reverse-order repeat) and the graph's do not move. The +VEP's do not move on either route (tree 2.7–3.1 µs, graph 0.7–1.0 µs). So on gcc, part +of the tree's cost is repeated math calls that the flag alone would recover for tier 1; +the rest, and all of it in the VEP, is repeated arithmetic that only computing each +named quantity once removes. Apple clang does not set `errno` by default, which is why +the Mac does not show the effect. On the same server the #752 fixture agreed across +ranks in all 10 runs at np = 3 and at np = 4 on both routes, and +`tests/parallel/ptest_jit_cache.py` passed at np = 4 on both. + **The operators agree to round-off.** `route_assemble.py` assembles the residual and the Jacobian on each route at the same state — rest (zero velocity, boundary values imposed), the tree's final state and the graph's final state — and compares them entry @@ -472,10 +491,9 @@ faster: the notch's pointwise setup takes 2.5 s against 15.9 s, its C is 22 KB a emits byte-identical C, because a constant law has no node to lower. The assembly gain is smaller than the prototype's kernel timings suggested (an eighth of the time per call) because the pointwise kernel is one part of the assembly, beside quadrature and -the element loop. Still to measure: an idle machine with repeated runs, Linux with gcc -(whose default `-fmath-errno` keeps it from merging repeated `pow`, `exp` and `sqrt` -calls, which the graph's temporaries do for every named quantity), a three-dimensional fixture, -and the small kernels of `Integral` and `BdIntegral`. +the element loop. On Linux with gcc the ratios are the same (below). Still to measure: +an idle machine with repeated runs, a three-dimensional fixture, and the small kernels +of `Integral` and `BdIntegral`. **More robust.** Each failure class below needed a patch in tier 1, placed where it surfaced; in the graph it cannot arise, or arises only in one body: From 93ef499ec953ad8428c2cafcc85196d06c0fe01a Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 17:48:05 +1100 Subject: [PATCH 08/14] JIT: the graph route is the only route (#823, tier 2, staging step 4) getext() always lowers its callbacks onto the shared graph and emits one C temporary per distinct computation; _jacobian_unwrap always builds guarded nodes. The UW_JIT_GRAPH switch is gone, and with it the expanded-tree route: _reveal_constants, the two manifest consistency guards, the whole-kernel unwrap, the coordinate recovery, the opt-in UW_JIT_CSE path, _collect_constant_atoms, _xreplace_shared, _unique_symbols, the tree's sqrt guard, and the unused prepare_for_cache_key and _createext. Leaves are spelled by an explicit map built for each compile (_leaf_spellings, _spell_leaf) instead of C names patched onto the field classes: a field class no longer carries the last compile's array slot into the next, and a field the compile does not own is an unconvertible-symbol error instead of a silent read of another field's data. The integration-point gradient refusal and the unconvertible-symbol message are kept. _extract_constants is the manifest of the lowered kernels. The verbose "Processing JIT" line prints the lowered kernel (the mathematics, with named quantities as nodes); test_0004 reads it. Tests: test_0022 loses the three tests of deleted helpers; test_0023 reads free_symbols; test_0024 holds the guarded lowering to a guarded tree built in the test, and pins the manifest of a nested law. Session scripts take the route from the build (this branch or development); the prototype lowering and its A/B harness are removed, superseded by the library module. Underworld development team with AI support from Claude Code --- scripts/sessions/jit_graph/assembly_split.py | 8 +- .../sessions/jit_graph/graph_vs_library.py | 261 ------ scripts/sessions/jit_graph/kernel_graph.py | 278 ------ .../sessions/jit_graph/plain_diff_probe.py | 9 +- .../sessions/jit_graph/rank_agreement_752.py | 3 +- scripts/sessions/jit_graph/route_ab.py | 9 +- scripts/sessions/jit_graph/route_assemble.py | 7 +- .../cython/petsc_generic_snes_solvers.pyx | 90 +- src/underworld3/utilities/_jit_graph.py | 48 +- src/underworld3/utilities/_jitextension.py | 827 ++++-------------- ...022_unwrap_memoised_matches_fixed_point.py | 60 +- ...est_0023_field_realness_and_derivatives.py | 6 +- tests/test_0024_jit_graph_lowering.py | 51 +- 13 files changed, 305 insertions(+), 1352 deletions(-) delete mode 100644 scripts/sessions/jit_graph/graph_vs_library.py delete mode 100644 scripts/sessions/jit_graph/kernel_graph.py diff --git a/scripts/sessions/jit_graph/assembly_split.py b/scripts/sessions/jit_graph/assembly_split.py index 386ce44e6..b141f6045 100644 --- a/scripts/sessions/jit_graph/assembly_split.py +++ b/scripts/sessions/jit_graph/assembly_split.py @@ -6,15 +6,16 @@ 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 ``UW_JIT_GRAPH``. +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 -from underworld3.utilities import _jit_graph sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) import fixtures # noqa: E402 @@ -25,7 +26,8 @@ description="route_ab .npz whose X is the state"), uw_repeat=uw.Param(20, description="assemblies timed"), ) -route = "graph" if _jit_graph.enabled() else "tree" +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) diff --git a/scripts/sessions/jit_graph/graph_vs_library.py b/scripts/sessions/jit_graph/graph_vs_library.py deleted file mode 100644 index 496bb6e40..000000000 --- a/scripts/sessions/jit_graph/graph_vs_library.py +++ /dev/null @@ -1,261 +0,0 @@ -"""Kernels built two ways, compiled, and compared (tier 2 of #823). - - library: today's route on tier 1 -- the Newton source from the solver's own - ``_jacobian_source``, residual and Picard kernels unwrapped as ``getext`` - unwraps them (keep-constants), derivatives through ``diff_wrt_field`` with the - solver's explicit uu_G3 loops, each output printed whole; - graph: ``kernel_graph`` nodes, the same loops, one temporary per node. - -Three kernels per fixture: the residual flux F1 (plain nodes), the Picard uu_G3 (atoms -frozen while differentiating, plain nodes after) and the Newton uu_G3 (guarded nodes). -Each pair is compiled into C functions of the same leaves and evaluated at random -states and at a state of rest. - -Determinism checks: run under several PYTHONHASHSEED values, with ``-uw_preamble N`` -(N throwaway objects created first, shifting every creation counter) and with -``-uw_redeclare 1`` (the law built twice in one process, as a re-run notebook cell -does); the printed md5s must not change. - -Fixtures are in fixtures.py: box, powerlaw, linear, vep and ti are built there and -reproducible from the repository; notch is the Spiegelman campaign law -(~/+Simulations/spiegelman_hardcase/drivers/notch_model.py). -""" -import ctypes -import hashlib -import os -import subprocess -import sys -import tempfile -import time - -import numpy as np -import sympy -from sympy.core.function import AppliedUndef -from sympy.printing.c import c_code_printers -from sympy.vector.scalar import BaseScalar - -import underworld3 as uw -from underworld3.function.expressions import UWexpression, unwrap_expression -from underworld3.function import diff_wrt_field -from underworld3.utilities._jitextension import _extract_constants, _pack_constants - -sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) -from kernel_graph import KernelGraph, _KernelNode, emission_order, guard # noqa: E402 -import fixtures # noqa: E402 - -params = uw.Params( - uw_fixture=uw.Param("box", description="box | powerlaw | linear | vep | ti | notch"), - uw_states=uw.Param(2000, description="random states for the value comparison"), - uw_calls=uw.Param(200000, description="kernel calls for the run-time measurement"), - uw_preamble=uw.Param(0, description="throwaway expressions and one mesh created before the law"), - uw_redeclare=uw.Param(0, description="1 = build the law twice and use the second"), - uw_numeric_n=uw.Param(0, description="powerlaw: 1 = the stress exponent as a plain number, not a constant atom"), -) -fixture = str(params.uw_fixture) - - -for i in range(int(params.uw_preamble)): - uw.expression(f"preamble_{i}", float(i)) -if int(params.uw_preamble): - uw.meshing.UnstructuredSimplexBox(cellSize=0.5) -kwargs = {"numeric_n": bool(int(params.uw_numeric_n))} if fixture == "powerlaw" else {} -stokes, admissible = fixtures.build(fixture, **kwargs) -if int(params.uw_redeclare): - stokes, admissible = fixtures.build(fixture, **kwargs) -mesh = stokes.mesh -dim = mesh.dim -L = stokes.Unknowns.L -F1 = sympy.Array(stokes.F1.sym).reshape(dim, dim) - - -def g3(F): - """The solver's explicit uu_G3 loops, differentiating through diff_wrt_field.""" - F = sympy.Array(F).reshape(dim, dim) - G = sympy.zeros(dim * dim, dim * dim) - for gc in range(dim): - for dg in range(dim): - dF = diff_wrt_field(F, L[gc, dg]) - for fc in range(dim): - for df in range(dim): - G[fc * dim + gc, df * dim + dg] = dF[fc, df] - return list(G) - - -def keep_constants(e): - return unwrap_expression(e, mode="symbolic_keep_constants") - - -timing = {} -def timed(label, fn): - t = time.perf_counter() - out = fn() - timing[label] = time.perf_counter() - t - return out - - -F1_flat = [F1[i, j] for i in range(dim) for j in range(dim)] -graph = KernelGraph() -kernels = {} -kernels["residual F1"] = ( - timed("library residual: unwrap", lambda: [keep_constants(e) for e in F1_flat]), - timed("graph residual: lower", lambda: [graph.lower(e, guarded=False) for e in F1_flat])) -picard_raw = timed("both Picard: differentiate (atoms frozen)", lambda: g3(F1)) -kernels["Picard G3"] = ( - timed("library Picard: unwrap", lambda: [keep_constants(e) for e in picard_raw]), - timed("graph Picard: lower", lambda: [graph.lower(e, guarded=False) for e in picard_raw])) -F1_lib = timed("library Newton: source (unwrap + guard)", - lambda: stokes._jacobian_source(F1, stokes._newton_flux(F1))) -G3_lib = timed("library Newton: differentiate", lambda: g3(F1_lib)) -F1_graph = timed("graph Newton: lower", - lambda: sympy.Array([guard(graph.lower(e, guarded=True)) for e in F1_flat], - (dim, dim))) -G3_graph = timed("graph Newton: differentiate", lambda: g3(F1_graph)) -kernels["Newton G3"] = (G3_lib, G3_graph) - -# ------------------------------------------------------------------ spelling of leaves -def leaves_of(exprs): - out = set() - for x in exprs: - x = sympy.sympify(x) - out |= {a for a in x.atoms(AppliedUndef) if not isinstance(a, _KernelNode)} - out |= {s for s in x.free_symbols if isinstance(s, (UWexpression, BaseScalar))} - return out - -graph_bodies = [b for _, gr in kernels.values() for _, b in emission_order(gr, str)[0]] -graph_leaves = leaves_of([x for _, gr in kernels.values() for x in gr] + graph_bodies) -library_leaves = leaves_of([x for lib, _ in kernels.values() for x in lib]) -leaves = graph_leaves | library_leaves -fields = sorted((s for s in leaves if isinstance(s, AppliedUndef)), key=KernelGraph.display_key) -slot = {f: k for k, f in enumerate(fields)} -library_consts = {e for _, e in _extract_constants( - tuple(sympy.ImmutableMatrix([list(lib)]) for lib, _ in kernels.values()), mesh)[0]} -graph_consts = {s for s in graph_leaves if isinstance(s, UWexpression)} -consts = sorted(library_consts | graph_consts, key=lambda e: (e.name, e.instance_number)) -cindex = {c: i for i, c in enumerate(consts)} -cvals = np.asarray(_pack_constants(list(enumerate(consts))), dtype=float) - - -def spell(s): - if s in slot: - return f"petsc_u[{slot[s]}]" - if s in cindex: - return f"constants[{cindex[s]}]" - if isinstance(s, BaseScalar): - return f"petsc_x[{s._id[0]}]" - return str(s) - - -spelled = {s: sympy.Symbol(spell(s)) for s in leaves} -printer = c_code_printers["c99"]({}) -def c_of(e, temps=None): - e = sympy.sympify(e) - if temps: - e = e.xreplace(temps) - return printer.doprint(e.xreplace(spelled)) - -# ------------------------------------------------------------------ compile and compare -SIG = "void k(const double *petsc_u, const double *petsc_x, const double *constants, double *out)" -HDR = ("#include \n" - "static inline double Heaviside_1(double x){return x<0?0:x>0?1:0.5;}\n") -work = tempfile.mkdtemp(dir=os.path.expanduser("~/+Simulations"), prefix=f"jit_graph_{fixture}_") -P = ctypes.POINTER(ctypes.c_double) -md5 = lambda s: hashlib.md5(s.encode()).hexdigest()[:10] - - -def compile_c(tag, body): - cfile = os.path.join(work, f"{tag}.c") - open(cfile, "w").write(HDR + SIG + " {\n" + body.replace("Heaviside(", "Heaviside_1(") + "\n}\n") - so = os.path.join(work, f"{tag}.so") - t = time.perf_counter() - subprocess.run(["cc", "-std=c99", "-O3", "-g0", "-shared", "-fPIC", cfile, "-o", so], check=True) - lib = ctypes.CDLL(so) - lib.k.restype = None - return lib, time.perf_counter() - t, os.path.getsize(so) - - -empty, _, _ = compile_c("empty", "") -print(f"[{fixture}] seed {os.environ.get('PYTHONHASHSEED', 'random')} preamble {params.uw_preamble} " - f"redeclare {params.uw_redeclare}") -for k, v in timing.items(): - print(f" {k:44s} {v:9.2f} s") -extra = sorted(c.name for c in graph_consts - library_consts) -missing = sorted(c.name for c in library_consts - graph_consts) -print(f"constants: library {len(library_consts)}, graph {len(graph_consts)}; " - f"graph has the library's: {not missing}; extra in graph: {extra}") - -rng = np.random.default_rng(7) -states = [] -Lfields = {f for f in fields if f in set(L)} -for _ in range(int(params.uw_states)): - scale = 10.0 ** rng.uniform(-6, 2) - u = rng.normal(size=len(fields)) - for f, k in slot.items(): - if f in Lfields: - u[k] *= scale - for name, (lo, hi) in admissible.items(): - if f.func.__name__.strip("{}") == name: - u[k] = rng.uniform(lo, hi) - states.append((u, rng.uniform(0, 1, size=3))) -rest = (np.zeros(len(fields)), np.full(3, 0.5)) - -for name, (lib_exprs, graph_exprs) in kernels.items(): - t = time.perf_counter() - order, key = emission_order(graph_exprs, spell) - temps_by_key = {k: sympy.Symbol(f"t{i}", real=True) for i, (k, _) in enumerate(order)} - temps = {app: temps_by_key[k] for app, k in key.items()} - src_graph = "\n".join([f"const double {temps_by_key[k]} = {c_of(body, temps)};" for k, body in order] - + [f"out[{i}] = {c_of(x, temps)};" for i, x in enumerate(graph_exprs)]) - t_emit = time.perf_counter() - t - t = time.perf_counter() - src_lib = "\n".join(f"out[{i}] = {c_of(x)};" for i, x in enumerate(lib_exprs)) - t_print = time.perf_counter() - t - tag = name.replace(" ", "_") - lib_g, cc_g, so_g = compile_c(f"{tag}_graph", src_graph) - lib_l, cc_l, so_l = compile_c(f"{tag}_library", src_lib) - nout = len(graph_exprs) - - def call(lib, u, x): - out = np.zeros(nout) - lib.k(u.ctypes.data_as(P), x.ctypes.data_as(P), cvals.ctypes.data_as(P), out.ctypes.data_as(P)) - return out - - nonzero = exact = nan_mismatch = nan_both = lib_zero_graph_not = 0 - worst_block = 0.0 - for u, x in states: - a, b = call(lib_g, u, x), call(lib_l, u, x) - nan_mismatch += int(np.any(np.isnan(a) != np.isnan(b))) - nan_both += int(np.sum(np.isnan(a) & np.isnan(b))) - lib_zero_graph_not += int(np.sum((a != 0) & (b == 0))) - m = (b != 0) & np.isfinite(b) & np.isfinite(a) - nonzero += m.sum(); exact += (a[m] == b[m]).sum() - if m.any(): - worst_block = max(worst_block, float(np.max(np.abs(a[m] - b[m])) / np.max(np.abs(b[m])))) - a0, b0 = call(lib_g, *rest), call(lib_l, *rest) - zl = sum(sympy.sympify(x) == 0 for x in lib_exprs) - zg = sum(sympy.sympify(x) == 0 for x in graph_exprs) - - per = {} - u, x = states[0] - for tg, lib in (("graph", lib_g), ("library", lib_l), ("empty", empty)): - out = np.zeros(max(nout, 1)) - args = (u.ctypes.data_as(P), x.ctypes.data_as(P), cvals.ctypes.data_as(P), out.ctypes.data_as(P)) - n = int(params.uw_calls) - t = time.perf_counter() - for _ in range(n): - lib.k(*args) - per[tg] = 1e9 * (time.perf_counter() - t) / n - - print(f"--- {name}: md5 graph {md5(src_graph)} library {md5(src_lib)}") - print(f" emit {t_emit:.2f} s, print {t_print:.2f} s; C graph {len(src_graph):,} B " - f"({len(order)} temporaries), library {len(src_lib):,} B; cc -O3 {cc_g:.2f} s / {cc_l:.2f} s; " - f".so {so_g:,} / {so_l:,} B") - print(f" {len(states)} states: {nonzero} non-zero entries, bit-identical {exact} " - f"({100 * exact / max(nonzero, 1):.0f}%), max |diff|/max|block| {worst_block:.1e}, " - f"NaN in one route {nan_mismatch}, NaN in both {nan_both}, " - f"graph non-zero where library exactly zero {lib_zero_graph_not}") - print(f" structural zeros: library {zl}, graph {zg}; state of rest: finite graph " - f"{bool(np.isfinite(a0).all())} library {bool(np.isfinite(b0).all())}, " - f"elementwise equal {bool(np.array_equal(a0, b0))}") - print(f" one call: graph {per['graph'] - per['empty']:.1f} ns, library " - f"{per['library'] - per['empty']:.1f} ns (empty-kernel ctypes call {per['empty']:.1f} ns subtracted)") -print("work dir:", work) diff --git a/scripts/sessions/jit_graph/kernel_graph.py b/scripts/sessions/jit_graph/kernel_graph.py deleted file mode 100644 index 92d9a7be9..000000000 --- a/scripts/sessions/jit_graph/kernel_graph.py +++ /dev/null @@ -1,278 +0,0 @@ -"""Prototype of the shared-graph lowering for JIT kernels (tier 2 of #823). - -Not library code: the design note it supports is -docs/developer/design/jit-shared-graph-codegen.md. - -A non-constant UWexpression atom becomes a NODE: an applied undefined function of the -leaves its value depends on. SymPy's own chain rule then differentiates through it, -because ``fdiff(i)`` returns the partial derivative with respect to argument slot i. - -Two kinds of identity are kept apart: - -- In Python, a node is identified by its body. Two atoms with equal bodies share a - node; bodies that use different constants are different even when the constants - share a display name, because SymPy equality sees the difference. The class name is - a hash of the body written with DISPLAY identities only (no creation counters), and - each class carries its graph's serial in its SymPy identity (``_ctx``), so classes - from two compiles never compare equal. -- In C, a temporary is identified by a hash of the C it computes: its body with every - leaf written as the C the kernel reads (``petsc_u[3]``, ``constants[2]``) and every - child as the child's hash. Temporaries are ordered and merged by that hash, so the - source is a function of the mathematics and the kernel's data layout, not of Python - names, creation counters or hash seeds. - -Partial derivatives are taken with every leaf replaced by an independent real dummy. -A derivative that is zero, a number or a single leaf is returned as itself. -""" -import hashlib -import itertools - -import sympy -from sympy.core.function import AppliedUndef, UndefinedFunction -from sympy.vector.scalar import BaseScalar - -from underworld3.function.expressions import UWexpression -from underworld3.function._function import UnderworldAppliedFunction -from underworld3.utilities._jitextension import _is_truly_constant - -EPS2 = sympy.Float(1.0e-36) -_graph_serial = itertools.count() - - -def guard(e): - """The sqrt guard of ``_jacobian_unwrap``: every half-integer power whose base has - free symbols (constants included) gets ``+1e-36`` in its base. Memoised on identity - (as tier 1 does); node applications are not entered.""" - memo = {} - - def g(n): - hit = memo.get(id(n)) - if hit is not None: - return hit[1] - out = n - args = getattr(n, "args", None) - if args and not isinstance(n, _KernelNode): - new_args = tuple(g(a) for a in args) - if any(a is not b for a, b in zip(args, new_args)) and args != new_args: - out = n.func(*new_args) - if (out.is_Pow and out.exp.is_Rational and out.exp.q == 2 - and out.args[0].free_symbols): - out = sympy.Pow(out.args[0] + EPS2, out.exp) - memo[id(n)] = (n, out) - return out - - return g(e) - - -class _KernelNode(AppliedUndef): - """A named sub-expression applied to the leaves its value depends on.""" - - def fdiff(self, argindex=1): - return self._graph.slot_derivative(self, argindex - 1) - - -def body_of(app): - """A node application's body, with the application's arguments substituted.""" - cls = app.func - if app.args == cls._deps: - return cls._body - return cls._body.xreplace(dict(zip(cls._deps, app.args))) - - -class KernelGraph: - """One lowering context: the nodes of one compile.""" - - def __init__(self): - self.serial = next(_graph_serial) - self._const = {} # id(atom) -> (atom, bool) - self._node = {} # (id(atom), guarded) -> (atom, application) - self._by_body = {} # body -> class - self._deriv = {} # (class, slot) -> derivative in terms of the class's deps - self._busy = set() - self._splitting = False - - # ------------------------------------------------------------------ leaves - def is_constant(self, atom): - hit = self._const.get(id(atom)) - if hit is None: - hit = self._const[id(atom)] = (atom, _is_truly_constant(atom, UWexpression)) - return hit[1] - - def is_leaf(self, s): - if isinstance(s, _KernelNode): - return False - if isinstance(s, (AppliedUndef, BaseScalar)): - return True - if isinstance(s, UWexpression): - return self.is_constant(s) - return isinstance(s, sympy.Symbol) - - @staticmethod - def display_key(s): - """What a leaf is called, without creation counters: stable across processes, - ranks, preambles and re-declarations. Not unique (two constants may share a - display name); used only to name and order, never to identify.""" - if isinstance(s, UnderworldAppliedFunction): - return f"F|{s.func.__name__}" - if isinstance(s, BaseScalar): - return f"X|{s._id[1]}|{s._id[0]}" - if isinstance(s, UWexpression): - return f"C|{s.name}" - if isinstance(s, AppliedUndef): - return f"A|{s.func.__name__}|{','.join(map(str, s.args))}" - return f"S|{s.name}" - - @staticmethod - def order_key(s): - """Display key, ties broken by RELATIVE creation order (as the constants - manifest orders same-named constants).""" - return (KernelGraph.display_key(s), getattr(s, "instance_number", 0)) - - def leaves(self, e): - out, seen, stack = set(), set(), [e] - while stack: - a = stack.pop() - if id(a) in seen: - continue - seen.add(id(a)) - if isinstance(a, _KernelNode): - out.update(a.args) - elif self.is_leaf(a): - out.add(a) - elif isinstance(a, sympy.Basic): - stack.extend(a.args) - return out - - # ------------------------------------------------------------------ nodes - def lower(self, e, guarded): - """``e`` with each non-constant UW atom replaced by its node.""" - if not isinstance(e, sympy.Basic): - return e - atoms = sorted((s for s in e.free_symbols - if isinstance(s, UWexpression) and not self.is_constant(s)), - key=self.order_key) - m = {s: self.node_of(s, guarded) for s in atoms} - # a coordinate atom lowers to the base scalar the kernel reads, as the tree does - m.update({s: s.sym for s in e.free_symbols - if isinstance(s, BaseScalar) and type(s).__name__ == "UWCoordinate"}) - return e.xreplace(m) if m else e - - def node_of(self, atom, guarded): - key = (id(atom), guarded) - hit = self._node.get(key) - if hit is not None: - return hit[1] - if key in self._busy: - raise ValueError(f"cyclic expression: {atom} contains itself") - self._busy.add(key) - try: - body = self.lower(atom.sym, guarded) - if guarded: - body = guard(body) - app = self.make_node(body) - finally: - self._busy.discard(key) - self._node[key] = (atom, app) - return app - - def _display_name(self, body): - canon = {a: sympy.Symbol(a.func.__name__) for a in body.atoms(_KernelNode)} - for leaf in self.leaves(body): - canon.setdefault(leaf, sympy.Symbol(self.display_key(leaf))) - text = sympy.srepr(body.xreplace(canon)) - return "N" + hashlib.sha1(text.encode()).hexdigest()[:12] - - def make_node(self, body): - """The node for ``body``: one per distinct body. A body that is a number, a leaf - or a single node needs no temporary and is returned as itself.""" - body = sympy.sympify(body) - if not self._splitting: - body = self._split_shared(body) - if body.is_Atom or isinstance(body, (AppliedUndef, BaseScalar)): - return body - cls = self._by_body.get(body) - if cls is None: - deps = tuple(sorted(self.leaves(body), key=self.order_key)) - cls = UndefinedFunction(self._display_name(body), bases=(_KernelNode,), - real=True, _ctx=self.serial, _n=len(self._by_body), - __dict__={"_graph": self}) - cls._body, cls._deps = body, deps - self._by_body[body] = cls - return cls(*cls._deps) - - def _split_shared(self, body): - """Repeated unnamed sub-expressions of one body become anonymous nodes.""" - if body.is_Atom: - return body - repl, (reduced,) = sympy.cse([body], symbols=sympy.numbered_symbols("_cse", real=True), - order="none") - if not repl: - return body - self._splitting = True - try: - m = {} - for sym, e in repl: - m[sym] = self.make_node(e.xreplace(m)) - return reduced.xreplace(m) - finally: - self._splitting = False - - # ------------------------------------------------------------------ derivatives - def slot_derivative(self, app, i): - cls = app.func - key = (cls, i) - if key not in self._deriv: - deps = cls._deps - dummies = [sympy.Dummy(real=True) for _ in deps] - b = cls._body.xreplace(dict(zip(deps, dummies))) - d = sympy.diff(b, dummies[i]).xreplace(dict(zip(dummies, deps))) - self._deriv[key] = self.make_node(d) - d = self._deriv[key] - if app.args == cls._deps: - return d - return d.xreplace(dict(zip(cls._deps, app.args))) - - -# ---------------------------------------------------------------------- emission -def emission_order(outputs, spell): - """The temporaries a kernel needs, in the order to write them. - - ``spell(leaf)`` is the C the kernel reads for a leaf. Each node application gets a - canonical key: a hash of its body with leaves spelled as C and children replaced by - their own keys. Applications with equal keys compute the same C and share one - temporary; temporaries follow the ones they use, ties broken by key. - - Returns ``(order, key)``: ``order`` is a list of ``(key, body)``; ``key`` maps each - node application reached to its key. - """ - key = {} - - def key_of(app): - hit = key.get(app) - if hit is not None: - return hit - body = body_of(app) - sub = {c: sympy.Symbol("@" + key_of(c)) for c in body.atoms(_KernelNode)} - for leaf in body.free_symbols | body.atoms(AppliedUndef): - if not isinstance(leaf, _KernelNode) and leaf not in sub: - sub[leaf] = sympy.Symbol(spell(leaf)) - k = hashlib.sha1(sympy.srepr(body.xreplace(sub)).encode()).hexdigest()[:16] - key[app] = k - return k - - order, placed = [], set() - - def visit(app): - k = key_of(app) - if k in placed: - return - placed.add(k) - body = body_of(app) - for child in sorted(body.atoms(_KernelNode), key=key_of): - visit(child) - order.append((k, body)) - - for out in outputs: - for child in sorted(sympy.sympify(out).atoms(_KernelNode), key=key_of): - visit(child) - return order, key diff --git a/scripts/sessions/jit_graph/plain_diff_probe.py b/scripts/sessions/jit_graph/plain_diff_probe.py index 8aafb710d..d8535b18b 100644 --- a/scripts/sessions/jit_graph/plain_diff_probe.py +++ b/scripts/sessions/jit_graph/plain_diff_probe.py @@ -5,16 +5,19 @@ ``sympy.diff`` / ``sympy.derive_by_array`` and runs the Newton solve that motivated the stand-in: the Drucker-Prager yield with a yield-stress floor at softness 0, whose floor ``(a + b + sqrt((a - b)**2)) / 2`` becomes an ``Abs`` of the pressure once fields are -real (``test_0023``). The route is chosen by ``UW_JIT_GRAPH``. +real (``test_0023``). 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 sympy import underworld3 as uw import underworld3.cython.generic_solvers as gs -from underworld3.utilities import _jit_graph gs.diff_wrt_field = lambda e, w: sympy.diff(e, w) gs.derive_by_array_wrt_field = lambda e, dx: sympy.derive_by_array(e, dx) -route = "graph" if _jit_graph.enabled() else "tree" +route = ("graph" if importlib.util.find_spec("underworld3.utilities._jit_graph") + else "tree") # the route is the build's: this branch, or development mesh = uw.meshing.UnstructuredSimplexBox(minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.25) diff --git a/scripts/sessions/jit_graph/rank_agreement_752.py b/scripts/sessions/jit_graph/rank_agreement_752.py index b779352a9..cb849747b 100644 --- a/scripts/sessions/jit_graph/rank_agreement_752.py +++ b/scripts/sessions/jit_graph/rank_agreement_752.py @@ -4,7 +4,8 @@ law of ``tests/test_0022_rotated_adjoint.py`` (adjoint branches), whose JIT source differed across ranks in about one np=2 run in five on development. Builds the kernels and reports, on rank 0, whether ``_agree_source_across_ranks`` saw different -hashes. The route is chosen by ``UW_JIT_GRAPH``. Run under ``mpirun -n 2``. +hashes. The route is chosen by the build: run it in a build of this branch +(graph) and in a build of development (tree). Run under ``mpirun -n 2``. """ import sympy import underworld3 as uw diff --git a/scripts/sessions/jit_graph/route_ab.py b/scripts/sessions/jit_graph/route_ab.py index 5b97a1875..e3baf26bb 100644 --- a/scripts/sessions/jit_graph/route_ab.py +++ b/scripts/sessions/jit_graph/route_ab.py @@ -1,7 +1,7 @@ """End-to-end A/B of the two JIT routes on one fixture (#823, tier 2). -The route is chosen by the environment, so one build serves both: ``UW_JIT_GRAPH=1`` is -the graph route, unset is the expanded tree (tier 1). Run once per route, then +The route is the build's: run once in a build of this branch (the graph route) and once +in a build of development (the expanded tree, tier 1), then ``route_ab_compare.py`` on the two ``.npz`` files. Records the whole pointwise setup (generation and compile separately), the size of the @@ -10,6 +10,7 @@ converged state. Run with ``UW_JIT_CACHE=0 UW_NO_USAGE_METRICS=1``. """ import hashlib +import importlib.util import os import sys import time @@ -17,7 +18,6 @@ import numpy as np import underworld3 as uw import underworld3.utilities._jitextension as jx -from underworld3.utilities import _jit_graph sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) import fixtures # noqa: E402 @@ -33,7 +33,8 @@ "measure how the Newton path responds to round-off-sized changes"), uw_label=uw.Param("", description="suffix for the output file"), ) -route = "graph" if _jit_graph.enabled() else "tree" +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) tangent = str(params.uw_tangent) diff --git a/scripts/sessions/jit_graph/route_assemble.py b/scripts/sessions/jit_graph/route_assemble.py index fa44e50a1..0c14bab3c 100644 --- a/scripts/sessions/jit_graph/route_assemble.py +++ b/scripts/sessions/jit_graph/route_assemble.py @@ -1,16 +1,16 @@ """Assemble the residual and the Jacobian of one fixture at a given state, on the route -the environment selects (``UW_JIT_GRAPH=1`` or unset), and save them (#823, tier 2). +the build provides (this branch: graph; development: tree), and save them (#823, tier 2). Two runs at the same state, one per route, are compared by ``route_assemble_compare.py``. The state is the SNES solution vector: ``zero`` (rest, with the boundary values), or the ``X`` saved by ``route_ab.py``. """ +import importlib.util import os import sys import numpy as np import underworld3 as uw -from underworld3.utilities import _jit_graph sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) import fixtures # noqa: E402 @@ -22,7 +22,8 @@ uw_tag=uw.Param("zero", description="label for the output file"), uw_out=uw.Param("~/+Simulations/jit_graph/tier2/assemble", description="output directory"), ) -route = "graph" if _jit_graph.enabled() else "tree" +route = ("graph" if importlib.util.find_spec("underworld3.utilities._jit_graph") + else "tree") # the route is the build's: this branch, or development fixture, tangent = str(params.uw_fixture), str(params.uw_tangent) stokes, _ = fixtures.build(fixture) diff --git a/src/underworld3/cython/petsc_generic_snes_solvers.pyx b/src/underworld3/cython/petsc_generic_snes_solvers.pyx index 46c5a81c8..47c6790b2 100644 --- a/src/underworld3/cython/petsc_generic_snes_solvers.pyx +++ b/src/underworld3/cython/petsc_generic_snes_solvers.pyx @@ -50,7 +50,6 @@ class _StrategyName(str): from underworld3.function import expression as public_expression expression = lambda *x, **X: public_expression(*x, _unique_name_generation=True, **X) -from underworld3.function.expressions import unwrap_expression as _unwrap_expression from underworld3.function._function import diff_wrt_field, derive_by_array_wrt_field @@ -65,24 +64,24 @@ def _public_names(cls): def _jacobian_unwrap(expr): - """Expand UWexpressions down to (but NOT including) constant atoms, for use - as the input to a Jacobian derivative (``derive_by_array`` / ``diff``). - - Applied element-wise over a sympy ``Matrix``/``Array`` so atoms embedded in - the residual flux are reached. Non-constant UWexpressions (e.g. the - effective viscosity ``Min(eta0, tau_y/2/eps_II)``) are expanded so the - derivative sees their field / grad-v dependence and forms the full Newton - tangent. Truly-constant atoms (``eta0``, ``tau_y``, ...) are kept as the - *same* symbol object so the JIT ``constants[]`` runtime-update mechanism is - preserved — the keep-constants predicate is shared with - ``getext()._extract_constants`` so the two cannot drift apart. - - This is a no-op for constant-viscosity problems (eta has no grad-v - dependence), so those Jacobians stay bit-identical. - - The unwrapped result is additionally made DIFFERENTIATION-SAFE: any - ``sqrt(g)`` whose argument carries non-constant symbols becomes - ``sqrt(g + 1e-36)``. Differentiating a bare invariant + r"""The Newton source of a residual flux: each non-constant UWexpression + replaced by its GUARDED node (``underworld3.utilities._jit_graph``), so that the + Jacobian derivative passes through it by the chain rule. + + Applied element-wise over a sympy ``Matrix``/``Array``. A node is an applied + function of the leaves its value depends on (field values and gradients, + coordinates, constant atoms), and its ``fdiff`` is the partial derivative of its + body with respect to that argument: the derivative of the effective viscosity + ``Min(eta0, tau_y/2/eps_II)`` with respect to grad v reaches the yield switch + (full Newton), where an opaque atom would freeze it (a Picard / defect-correction + tangent). Truly-constant atoms (``eta0``, ``tau_y``, ...) stay the same symbol + objects, so they keep their ``constants[]`` slots. + + No-op for constant-viscosity problems (eta has no grad-v dependence). + + The source is DIFFERENTIATION-SAFE: in every node body and at the top level, + each half-integer power whose base has free symbols gets ``+1e-36`` in its base. + Differentiating a bare invariant :math:`\dot\varepsilon_{II} = \sqrt{g}` produces :math:`\partial\sqrt{g}/\partial L = \dot\varepsilon/(2\dot\varepsilon_{II})` — the DIRECTION of the strain rate, which is 0/0 at a state of rest — @@ -94,56 +93,15 @@ def _jacobian_unwrap(expr): guard makes the derivative exactly zero at the singular point and perturbs it by under one part in 1e24 at any resolvable strain rate. The RESIDUAL is never routed through here, and the default (Picard) - tangent never calls this function, so both remain bit-identical. + tangent never calls this function. - See ``docs/developer/design/jacobian-unwrap-constants-bug.md``. + See ``docs/developer/design/jit-shared-graph-codegen.md`` and + ``docs/developer/design/jacobian-unwrap-constants-bug.md``. """ - eps2 = sympy.Float(1.0e-36) - - def _guard_sqrts(e): - # every HALF-INTEGER power: +1/2 (the invariant itself), -1/2 - # (its reciprocal in eta_pl = tau_y/(2 edot_II)), -3/2 (their - # derivatives), ... — all singular in value or derivative at a - # zero-argument state. - # The same bottom-up rebuild as `e.replace(query, value)`, memoised on node - # identity: the unwrapped flux repeats its shared sub-expressions as the same - # object, and `replace` walked every occurrence (measured 16 s of a 108 s - # notch compile, #823). - memo = {} - - def guard(n): - hit = memo.get(id(n)) - if hit is not None: - return hit[1] - out = n - args = getattr(n, "args", None) - if args: - new_args = tuple(guard(a) for a in args) - if any(a is not b for a, b in zip(args, new_args)) and args != new_args: - out = n.func(*new_args) - # replace(simultaneous=True): a rebuild that collapses to one of - # the changed arguments is not matched again - if any(out == a and a != b for a, b in zip(args, new_args)): - memo[id(n)] = (n, out) - return out - if (out.is_Pow and out.exp.is_Rational and out.exp.q == 2 - and out.args[0].free_symbols): - out = sympy.Pow(out.args[0] + eps2, out.exp) - memo[id(n)] = (n, out) - return out - - return guard(e) - from underworld3.utilities import _jit_graph - if _jit_graph.enabled(): - # The graph route (#823, tier 2): each non-constant atom becomes its guarded - # node instead of being expanded, and the derivative passes through it by - # the chain rule. The guard is applied in each node body and at the top level. - graph = _jit_graph.KernelGraph() - f = lambda e: _jit_graph.guard_half_integer_powers(graph.lower(e, guarded=True)) - else: - f = lambda e: _guard_sqrts( - _unwrap_expression(e, mode="symbolic_keep_constants")) + + graph = _jit_graph.KernelGraph() + f = lambda e: _jit_graph.guard_half_integer_powers(graph.lower(e, guarded=True)) if isinstance(expr, sympy.MatrixBase): return expr.applyfunc(f) if isinstance(expr, sympy.NDimArray): diff --git a/src/underworld3/utilities/_jit_graph.py b/src/underworld3/utilities/_jit_graph.py index b97064e98..63db6af39 100644 --- a/src/underworld3/utilities/_jit_graph.py +++ b/src/underworld3/utilities/_jit_graph.py @@ -14,12 +14,9 @@ - the emitter writes one C temporary per distinct computation, ordered and merged by a hash of the C it computes, so the generated source is a function of the mathematics and of the kernel's data layout alone. - -Selected by the private development switch ``UW_JIT_GRAPH=1`` while both routes exist. """ import hashlib import itertools -import os import sympy from sympy.core.function import AppliedUndef, UndefinedFunction @@ -31,11 +28,6 @@ _nodes_made = False -def enabled(): - """Whether the graph route is selected (``UW_JIT_GRAPH=1``).""" - return os.environ.get("UW_JIT_GRAPH", "").lower() in ("1", "true", "yes") - - def nodes_exist(): """Whether any node has been made in this process; when not, no expression can hold one and the unwrappers skip the search.""" @@ -85,6 +77,17 @@ def _ccode(self, printer): f"emitted as a temporary (a defect in the graph lowering, #823)") +class _CName(sympy.Symbol): + """A leaf written as the C the kernel reads: ``petsc_u[3]``, ``petsc_x[0]``, + ``constants[2]``.""" + + def __new__(cls, text): + return sympy.Symbol.__new__(cls, text, real=True) + + def _ccode(self, printer): + return self.name + + class _Temporary(sympy.Symbol): """A C temporary of one kernel, written ``uwt_``.""" @@ -424,22 +427,31 @@ def visit(app): return order, key -def emit(fn, spell, constants_rule): - """``fn`` (a lowered Matrix) as temporaries and outputs. +def emit(fn, spell): + """``fn`` (a lowered Matrix) as temporaries and outputs, ready to print. Returns ``(temporaries, outputs)``: ``temporaries`` is a list of ``(_Temporary, body)`` in the order to write them, and ``outputs`` is ``fn`` with - each node application replaced by its temporary. Constant atoms are replaced by - ``constants_rule`` (their ``constants[]`` placeholders) in both. + each node application replaced by its temporary. In both, every leaf is replaced + by a ``_CName`` holding ``spell(leaf)``. """ order, key = emission_order(list(fn), spell) + spelt = {} + + def rule_for(e, temp_of): + rule = {c: temp_of[key[c]] for c in e.atoms(_KernelNode)} + for leaf in e.free_symbols | e.atoms(AppliedUndef): + if leaf in rule or isinstance(leaf, (_KernelNode, _Temporary)): + continue + hit = spelt.get(leaf) + if hit is None: + hit = spelt[leaf] = _CName(spell(leaf)) + rule[leaf] = hit + return rule + temp_of, temporaries = {}, [] for i, (k, body) in enumerate(order): t = _Temporary(i) - rule = {c: temp_of[key[c]] for c in body.atoms(_KernelNode)} - rule.update(constants_rule) - temporaries.append((t, body.xreplace(rule))) + temporaries.append((t, body.xreplace(rule_for(body, temp_of)))) temp_of[k] = t - rule = {app: temp_of[key[app]] for app in fn.atoms(_KernelNode)} - rule.update(constants_rule) - return temporaries, fn.xreplace(rule) + return temporaries, fn.xreplace(rule_for(fn, temp_of)) diff --git a/src/underworld3/utilities/_jitextension.py b/src/underworld3/utilities/_jitextension.py index 0c7baee77..0a0ece09c 100644 --- a/src/underworld3/utilities/_jitextension.py +++ b/src/underworld3/utilities/_jitextension.py @@ -251,7 +251,7 @@ def flat(self) -> tuple: """Concatenate all slots into a single ordered tuple. The ordering (residual, bcs, jacobian, bd_residual, bd_jacobian) - matches what ``_createext()`` expects. + matches what ``generate_c_source()`` expects. """ return self.residual + self.bcs + self.jacobian + self.bd_residual + self.bd_jacobian @@ -275,70 +275,11 @@ def map(self, fn) -> 'JITCallbackSet': @property def counts(self): - """Lengths of each slot, for ``_createext()`` offset calculation.""" + """Lengths of each slot, for the offsets in ``generate_c_source()``.""" return (len(self.residual), len(self.bcs), len(self.jacobian), len(self.bd_residual), len(self.bd_jacobian)) -def _reveal_constants(fn): - """Expand non-constant UWexpressions to fixpoint, KEEPING truly-constant - atoms symbolic — so constants nested at ANY depth surface as atoms. - - This must run BEFORE the ``constants[]`` substitution: a top-level - ``xreplace`` cannot see a constant hidden inside a nested UWexpression - (every ``Parameters.*`` value is template-wrapped in one), so nested - constants were silently folded to C literals while the manifest still - listed them as live — issue #302. The keep-constants predicate here is - the same ``_is_truly_constant`` used to build the manifest, so the set - of atoms revealed is exactly the set the manifest routes to - ``constants[]``. - """ - from underworld3.function.expressions import ( - unwrap_expression, - UWDerivativeExpression, - ) - - if fn is None: - return fn - if isinstance(fn, UWDerivativeExpression): - fn = fn.doit() - if isinstance(fn, sympy.MatrixBase): - return fn.applyfunc( - lambda e: unwrap_expression(e, mode='symbolic_keep_constants')) - return unwrap_expression(fn, mode='symbolic_keep_constants') - - -def prepare_for_cache_key(fn, constants_subs_map): - """Prepare a single expression for JIT cache hashing. - - Three-phase process (mirrors the codegen lowering in ``_createext`` so - the cache key and the generated C agree — issue #302): - 1. Reveal nested constants (``_reveal_constants``). - 2. Substitute manifested constants with ``_JITConstant`` placeholders - so that changing a constant's *value* does not invalidate the cache. - 3. Unwrap the remaining UW atoms to pure SymPy so the hash is - deterministic. - """ - # Phase 1: reveal constants nested inside other UWexpressions. Loud on - # failure, exactly like the codegen path — a silent fallback here would - # hash constant VALUES into the cache key and force a recompile on - # every ramp (adversarial-review finding). - fn_structural = _reveal_constants(fn) - - # Phase 2: Substitute constants with _JITConstant placeholders - if constants_subs_map and fn_structural is not None: - try: - if hasattr(fn_structural, "xreplace"): - fn_structural = _xreplace_shared(fn_structural, constants_subs_map) - except Exception: - pass - - # Phase 3: Unwrap remaining (non-constant) expressions - return underworld3.function.expressions.unwrap( - fn_structural, keep_constants=False, return_self=False - ) - - # ============================================================================ # JIT Constants Support # ============================================================================ @@ -409,46 +350,19 @@ def _ccode(self, printer): def _extract_constants(all_fns, mesh): - """Extract constant UWexpressions from a list of pre-unwrap functions. - - Scans all expressions for UWexpression atoms where is_constant_expr() - is True (no spatial/field dependencies). Assigns deterministic indices - sorted by expression name for MPI consistency. - - Parameters - ---------- - all_fns : tuple of sympy expressions - The raw (pre-unwrap) function list. - mesh : underworld3.discretisation.Mesh - The mesh (currently unused, reserved for future mesh.t support). + """The ``constants[]`` manifest of a list of callback expressions: the constant + atoms their lowered kernels read (``_jit_graph``), ordered as ``_manifest_from`` + orders them. Returns ------- list of (int, UWexpression) Ordered mapping from constants[] index to UWexpression reference. dict - Mapping from UWexpression to _JITConstant symbol for substitution. + Mapping from UWexpression to _JITConstant symbol. """ - from underworld3.function.expressions import ( - is_constant_expr, - extract_expressions, - UWexpression, - ) - - constant_exprs = set() - - for fn in all_fns: - if fn is None: - continue - - # Handle Matrix expressions - if isinstance(fn, sympy.MatrixBase): - for elem in fn: - _collect_constant_atoms(elem, constant_exprs, is_constant_expr, UWexpression) - else: - _collect_constant_atoms(fn, constant_exprs, is_constant_expr, UWexpression) - - return _manifest_from(constant_exprs) + lowered = _jit_graph.lower_callbacks([fn for fn in all_fns if fn is not None], mesh) + return _manifest_from(_jit_graph.constant_leaves(lowered)) def _manifest_from(constant_exprs): @@ -486,37 +400,6 @@ def _manifest_from(constant_exprs): return manifest, subs_map -def _xreplace_shared(expr, rule): - """``expr.xreplace(rule)`` (a sympy expression or Matrix), visiting each node - OBJECT once. - - A compiled kernel repeats its shared sub-expressions as the same object (the - memoised unwrap inserts one object wherever an atom occurs, #823), and - ``xreplace`` rebuilds every occurrence: 4.3 s of the notch C generation for the - constants[] substitution alone. The result is the same expression. - """ - memo = {} - - def walk(e): - hit = memo.get(id(e)) - if hit is not None: - return hit[1] - if e in rule: - out = rule[e] - elif isinstance(e, sympy.Basic) and e.args: - new_args = tuple(walk(a) if isinstance(a, sympy.Basic) else a for a in e.args) - changed = any(n is not a for n, a in zip(new_args, e.args)) - out = e.func(*new_args) if changed else e - else: - out = e - memo[id(e)] = (e, out) - return out - - if isinstance(expr, sympy.MatrixBase): - return expr.applyfunc(walk) - return walk(expr) - - def _holds_instance(expr, types): """Whether ``expr`` holds a node of ``types``, visiting each node object once.""" from sympy.tensor.array import NDimArray @@ -580,36 +463,6 @@ def _without_dirac_deltas(expr, where): return expr.xreplace({d: sympy.S.Zero for d in deltas}) -def _unique_symbols(expr): - """The Symbol atoms of ``expr`` (a sympy expression, Matrix or Array): the same set - as ``expr.atoms(sympy.Symbol)``, found by visiting each node OBJECT once. - - ``atoms`` walks every occurrence of every node. An unwrapped constitutive law - repeats its shared sub-expressions as the SAME Python object (the memoised unwrap - inserts one object wherever an atom occurs, #823), so an identity walk is - proportional to the shared graph rather than the expanded tree: measured on the - Spiegelman notch kernels, ``atoms`` was 35 s of a 108 s compile. - """ - # keyed by id, holding the object so that no id is reused while the walk runs - seen = {} - found = set() - stack = [expr] - while stack: - e = stack.pop() - if id(e) in seen: - continue - seen[id(e)] = e - if isinstance(e, (sympy.MatrixBase, sympy.NDimArray)): - stack.extend(e) - continue - if isinstance(e, sympy.Symbol): - found.add(e) - continue - if isinstance(e, sympy.Basic): - stack.extend(e.args) - return found - - def _is_truly_constant(expr, UWexpression): """Check if a UWexpression resolves to a pure constant (no spatial deps). @@ -656,31 +509,6 @@ def _is_truly_constant(expr, UWexpression): return True -def _collect_constant_atoms(expr, result_set, is_constant_expr, UWexpression): - """Recursively collect constant UWexpression atoms from an expression.""" - - if isinstance(expr, UWexpression): - if _is_truly_constant(expr, UWexpression): - result_set.add(expr) - return # Don't recurse into constant expressions - # Non-constant UWexpression: check its inner sym for nested constants - if hasattr(expr, '_sym') and expr._sym is not None: - _collect_constant_atoms(expr._sym, result_set, is_constant_expr, UWexpression) - return - - if not hasattr(expr, 'atoms'): - return - - # Check all UWexpression atoms - for atom in _stable_sorted(_unique_symbols(expr)): - if isinstance(atom, UWexpression) and _is_truly_constant(atom, UWexpression): - result_set.add(atom) - elif isinstance(atom, UWexpression): - # Non-constant UWexpression: recurse into its sym - if hasattr(atom, '_sym') and atom._sym is not None: - _collect_constant_atoms(atom._sym, result_set, is_constant_expr, UWexpression) - - def _pack_constants(manifest): """Pack current values from a constants manifest into a flat array. @@ -828,15 +656,11 @@ def getext( # Extract constant UWexpressions that are routed through PETSc's # constants[] array. Value changes don't affect the C source — they # only alter what we pass to PetscDSSetConstants at solve time. - if _jit_graph.enabled(): - # The graph route (#823, tier 2): lower every callback once, and take the - # manifest from the constant leaves of what was lowered. - lowered_fns = _jit_graph.lower_callbacks(callbacks.flat(), mesh) - constants_manifest, constants_subs_map = _manifest_from( - _jit_graph.constant_leaves(lowered_fns)) - else: - lowered_fns = None - constants_manifest, constants_subs_map = _extract_constants(callbacks.flat(), mesh) + # Each callback is lowered onto the shared graph of named quantities once + # (``_jit_graph``, #823); the manifest is the constant leaves of what was lowered. + lowered_fns = _jit_graph.lower_callbacks(callbacks.flat(), mesh) + constants_manifest, constants_subs_map = _manifest_from( + _jit_graph.constant_leaves(lowered_fns)) if debug and underworld3.mpi.rank == 0: if constants_manifest: @@ -1051,6 +875,157 @@ def getext( ) +class _Unspellable(Exception): + """A leaf the kernel has no C for.""" + + +class _Refusal: + """A leaf that has a meaning but no C in a weak form, with the reason.""" + + def __init__(self, message): + self.message = message + + +_IP_DERIVATIVE_REFUSAL = _Refusal( + "derivative of an integration-point " + "variable has no meaning (the field is defined only at the " + "quadrature points), so the gradient here would be a silent " + "zero. This is refused in a WEAK FORM only, where the " + "discretisation is yours to choose: build the variable with " + "proxy_location='cells' instead, whose level sets are a " + "least-squares polynomial per cell and differentiate directly. " + "uw.function.evaluate() of the same derivative does answer: as " + "a query it recovers the gradient from a per-cell fit for you." +) + + +def _leaf_spellings(mesh, primary_field_list): + """The C a kernel reads for each mesh-variable leaf, keyed by the leaf's class. + + For a 2-D velocity and pressure in the primary arrays: ``V_x -> petsc_u[0]``, + ``V_y -> petsc_u[1]``, ``P -> petsc_u[2]``, ``V_x_x -> petsc_u_x[0]``, ..., + ``P_y -> petsc_u_x[5]``. Every field of the mesh is entered first, from the + auxiliary arrays (``petsc_a``), at its own field's component offset in the DM + (``_aux_component_offsets``: a dropped variable's field stays in the DM and keeps + its slots); the primary fields then replace their entries with ``petsc_u``. + Gradients run to ``cdim``, the embedded dimension, so a manifold mesh's third + partial is wired too. The gradient of an integration-point variable is a + ``_Refusal``. + """ + from underworld3 import VarType + + spellings = {} + + def enter(varlist, prefix, component_offsets=None): + u_i = 0 # component + u_x_i = 0 # gradient component + for var in varlist: + if component_offsets is not None: + u_i = component_offsets[var.field_id] + u_x_i = u_i * mesh.cdim + if var.vtype == VarType.SCALAR: + components = [var.fn] + elif var.vtype in (VarType.VECTOR, VarType.TENSOR, VarType.SYM_TENSOR, + VarType.MATRIX): + components = list(var.sym_1d) + else: + raise RuntimeError( + f"Unsupported type {var.vtype} for code generation. " + f"Please contact developers.") + ip = getattr(var, "is_integration_point", False) + for component in components: + spellings[type(component)] = f"{prefix}[{u_i}]" + u_i += 1 + for ind in range(mesh.cdim): + # _diff[ind] is the gradient component's class + spellings[component._diff[ind]] = ( + _IP_DERIVATIVE_REFUSAL if ip else f"{prefix}_x[{u_x_i}]") + u_x_i += 1 + + enter(_stable_sorted(mesh.vars.values()), "petsc_a", + component_offsets=_aux_component_offsets(mesh)) + enter(primary_field_list, "petsc_u") + return spellings + + +def _spell_leaf(leaf, spellings, constants_subs_map): + """The C for one leaf: a constants[] slot, a field value or gradient + (``spellings``), a coordinate or boundary normal, or a symbol that names its own + C (the time, ``petsc_t``). Raises ``_Unspellable`` for anything else.""" + from sympy.core.function import AppliedUndef + from sympy.vector.scalar import BaseScalar + + placeholder = constants_subs_map.get(leaf) if constants_subs_map else None + if placeholder is not None: + return placeholder._ccodestr + if isinstance(leaf, BaseScalar): + # the mesh names its coordinates; a fresh instance, or a UWCoordinate that + # SymPy's cache returned for its equal base scalar, is named from its index + # and system + try: + text = leaf._ccodestr + except AttributeError: + text = None + if isinstance(text, str): + return text + idx, system = leaf._id[0], str(leaf._id[1]) + return f"petsc_n[{idx}]" if "Gamma" in system else f"petsc_x[{idx}]" + if isinstance(leaf, AppliedUndef): + entry = spellings.get(type(leaf)) + if entry is None: + raise _Unspellable(leaf) + if isinstance(entry, _Refusal): + raise RuntimeError(f"{type(leaf).__name__}: {entry.message}") + return entry + text = getattr(leaf, "_ccodestr", None) + if isinstance(text, str) and hasattr(leaf, "_ccode"): + return text + raise _Unspellable(leaf) + + +def _unconvertible_message(symbols, index, fn_original): + details = [] + for sym in symbols: + detail = f" - {sym} (type: {type(sym).__name__})" + if hasattr(sym, "units"): + detail += f" [has units: {sym.units}]" + if hasattr(sym, "value"): + detail += f" [value: {sym.value}]" + details.append(detail) + return ( + f"\n{'=' * 70}\n" + f"JIT COMPILATION ERROR: Expression contains unconvertible symbols\n" + f"{'=' * 70}\n\n" + f"The following symbols could not be converted to C code:\n" + + "\n".join(details) + "\n\n" + f"This usually means:\n" + f" 1. A UWexpression or UWQuantity was not properly expanded\n" + f" 2. An arithmetic operation failed (e.g., Matrix * UWexpression)\n" + f" 3. A symbolic function is missing from the expression tree\n" + f" 4. A field that belongs to another mesh, or to no mesh\n\n" + f"Expression index: {index}\n" + f"Original expression: {fn_original}\n\n" + f"TIP: Check that all expression operations (*, /, +, -) produce\n" + f"valid SymPy expressions. For example, ensure scalar * Matrix\n" + f"and not Matrix * scalar when using UWexpression objects.\n" + f"{'=' * 70}" + ) + + +def _print_kernel(printer, temporaries, outputs, out): + """The C body of one kernel: a ``const double`` per temporary, in order, then + the outputs. A SymPy function the printer cannot write returns its + ``// Not supported in C:`` text, which the caller refuses.""" + lines = [] + for t, body in temporaries: + code = printer.doprint(body) + if code.startswith("// Not supported in C:"): + return code + lines.append(f"const double {t._ccodestr} = {code};") + lines.append(printer.doprint(outputs, out)) + return "\n".join(lines) + + @timing.routine_timer_decorator def _aux_component_offsets(mesh): """Component offset of every field of the mesh DM, keyed by field id. @@ -1134,11 +1109,13 @@ def generate_c_source( primary_field_list : list Variables that map to PETSc primary variable arrays (``petsc_u[]``). constants_subs_map : dict, optional - Mapping from UWexpression → ``_JITConstant`` placeholder. + Mapping from UWexpression → ``_JITConstant`` placeholder; built from the + lowered callbacks when not given. lowered_fns : list of sympy.Matrix, optional The callbacks lowered onto the shared graph (``_jit_graph.lower_callbacks``), - one per entry of ``callbacks.flat()``. When given, each kernel is emitted as - temporaries and outputs instead of being unwrapped and printed whole. + one per entry of ``callbacks.flat()``; lowered here when not given. Each + kernel is emitted as one C temporary per distinct computation, then its + outputs. Returns ------- @@ -1158,154 +1135,14 @@ def generate_c_source( count_residual_sig, count_bc_sig, count_jacobian_sig, \ count_bd_residual_sig, count_bd_jacobian_sig = callbacks.counts - # `_ccode` patching - def ccode_patch_fns(varlist, prefix_str, component_offsets=None): - """ - This function patches uw functions with the necessary ccode - routines for the code printing. - - For a `varlist` consisting of 2d velocity & pressure variables, - for example, it'll generate routines which write the following, - where `prefix_str="petsc_u"`: - V_x : "petsc_u[0]" - V_y : "petsc_u[1]" - P : "petsc_u[2]" - V_x_x : "petsc_u_x[0]" - V_x_y : "petsc_u_x[1]" - V_y_x : "petsc_u_x[2]" - V_y_y : "petsc_u_x[3]" - P_x : "petsc_u_x[4]" - P_y : "petsc_u_x[5]" - - Params - ------ - varlist: list - The variables to patch. Note that *all* the variables in the - corresponding `PetscDM` must be included. They must also be - ordered according to their `field_id`. - prefix_str: str - The string prefix to write. - component_offsets: dict, optional - Component offset of every field in the DM, by ``field_id`` - (see ``_aux_component_offsets``). When given, each variable - is patched from ITS OWN field's offset instead of a running - count over ``varlist``: a field whose Python variable has - been dropped stays in the DM and still occupies its slots, - so a running count would shift every later variable onto - the wrong data. - """ - u_i = 0 # variable increment - u_x_i = 0 # variable gradient increment - lambdafunc = lambda self, printer: self._ccodestr - - def _no_derivative(self, printer): - # An integration-point variable has no gradient (its tabulated - # derivative is identically zero), so a derivative of its symbol - # in a weak form would be a silent zero. Refuse at code generation. - raise RuntimeError( - f"{self.__class__.__name__}: derivative of an integration-point " - "variable has no meaning (the field is defined only at the " - "quadrature points), so the gradient here would be a silent " - "zero. This is refused in a WEAK FORM only, where the " - "discretisation is yours to choose: build the variable with " - "proxy_location='cells' instead, whose level sets are a " - "least-squares polynomial per cell and differentiate directly. " - "uw.function.evaluate() of the same derivative does answer: as " - "a query it recovers the gradient from a per-cell fit for you." - ) - - for var in varlist: - is_ip = getattr(var, "is_integration_point", False) - dfunc = _no_derivative if is_ip else lambdafunc - if component_offsets is not None: - u_i = component_offsets[var.field_id] - u_x_i = u_i * mesh.cdim - if var.vtype == VarType.SCALAR: - # monkey patch this guy into the function - type(var.fn)._ccodestr = f"{prefix_str}[{u_i}]" - type(var.fn)._ccode = lambdafunc - u_i += 1 - # Now patch the gradient components. The gradient of a - # field on the mesh lives in the embedded coordinate - # space (cdim-dim), so iterate to cdim — not dim. For - # volume meshes ``dim == cdim`` so this is unchanged; - # for manifold meshes (e.g. SphericalManifold dim=2, - # cdim=3) the third partial ``f_{,2}`` exists and - # needs to be wired to ``u_x[2]``. - for ind in range(mesh.cdim): - # Note that var.fn._diff[ind] returns the class, so we don't need type(var.fn._diff[ind]) - var.fn._diff[ind]._ccodestr = f"{prefix_str}_x[{u_x_i}]" - var.fn._diff[ind]._ccode = dfunc - u_x_i += 1 - elif ( - var.vtype == VarType.VECTOR - or var.vtype == VarType.TENSOR - or var.vtype == VarType.SYM_TENSOR - or var.vtype == VarType.MATRIX - ): - # Pull out individual sub components - for comp in var.sym_1d: - # monkey patch - type(comp)._ccodestr = f"{prefix_str}[{u_i}]" - type(comp)._ccode = lambdafunc - u_i += 1 - # Iterate to cdim (embedded coord dim) — see the - # scalar branch above for the dim != cdim reason. - for ind in range(mesh.cdim): - # Note that var.fn._diff[ind] returns the class, so we don't need type(var.fn._diff[ind]) - comp._diff[ind]._ccodestr = f"{prefix_str}_x[{u_x_i}]" - comp._diff[ind]._ccode = dfunc - u_x_i += 1 - else: - raise RuntimeError( - f"Unsupported type {var.vtype} for code generation. Please contact developers." - ) - - # Patch in `_code` methods. Note that the order here - # is important, as the secondary call will overwrite - # those patched in the first call. - - ccode_patch_fns(_stable_sorted(mesh.vars.values()), "petsc_a", - component_offsets=_aux_component_offsets(mesh)) - ccode_patch_fns(primary_field_list, "petsc_u") - - # Also patch `BaseScalar` types. Nothing fancy - patch the overall type, - # make sure each component points to the correct PETSc data + if lowered_fns is None: + lowered_fns = _jit_graph.lower_callbacks(fns, mesh) + if constants_subs_map is None: + _, constants_subs_map = _manifest_from(_jit_graph.constant_leaves(lowered_fns)) - ## This is set up in the mesh at the moment but this does seem to be the wrong place - - # mesh.N.x._ccodestr = "petsc_x[0]" - # mesh.N.y._ccodestr = "petsc_x[1]" - # mesh.N.z._ccodestr = "petsc_x[2]" - - # # Surface integrals also have normal vector information as petsc_n - - # mesh.Gamma_N.x._ccodestr = "petsc_n[0]" - # mesh.Gamma_N.y._ccodestr = "petsc_n[1]" - # mesh.Gamma_N.z._ccodestr = "petsc_n[2]" - - def _basescalar_ccode(self, printer): - """C code for coordinate symbols, with fallback for new instances. - - sympy.simplify() may create new BaseScalar/UWCoordinate instances - that lack _ccodestr. We recover it from the coordinate's _id attribute - which stores (index, system_name). - """ - if hasattr(self, '_ccodestr'): - return self._ccodestr - # Fallback: compute from _id - idx = self._id[0] - system_name = str(self._id[1]) - if 'Gamma' in system_name: - return f"petsc_n[{idx}]" - else: - return f"petsc_x[{idx}]" - - type(mesh.N.x)._ccode = _basescalar_ccode - # Gamma base scalars (un-normalised face normal) — ensure ccode is registered - Gamma_scalars = mesh._Gamma.base_scalars() - if type(Gamma_scalars[0]) is not type(mesh.N.x): - type(Gamma_scalars[0])._ccode = _basescalar_ccode + # The C a kernel reads for each leaf of the graph: an explicit map, built for + # this compile, instead of C names patched onto the field classes (#823). + spellings = _leaf_spellings(mesh, primary_field_list) # Create a custom functions replacement dictionary. # Note that this dictionary is really just to appease Sympy, @@ -1374,246 +1211,28 @@ def _handle_UnevaluatedExpr(expr): underworld3._libdirs.clear() underworld3._libfiles.clear() - def _spell(leaf): - # the C a kernel reads for a leaf of the graph - placeholder = constants_subs_map.get(leaf) if constants_subs_map else None - if placeholder is not None: - return placeholder._ccodestr - if isinstance(leaf, sympy.vector.scalar.BaseScalar): - # a coordinate may be a fresh instance, or a UWCoordinate SymPy's cache - # handed back for its equal base scalar, without the name the mesh set: - # recover it from the coordinate's index and system, as the coordinate - # recovery below does for the printer - try: - return leaf._ccodestr - except AttributeError: - idx, system = leaf._id[0], str(leaf._id[1]) - leaf._ccodestr = (f"petsc_n[{idx}]" if "Gamma" in system - else f"petsc_x[{idx}]") - return leaf._ccodestr - return printer.doprint(leaf) - eqns = [] - for index, fn in enumerate(fns): - - # Save original for debugging - fn_original = fn - temporaries = () - if lowered_fns is not None: - temporaries, fn = _jit_graph.emit( - lowered_fns[index], _spell, constants_subs_map or {}) - - # --- Gate the UW lowering (issue #302 pipeline) on the presence of - # UW-expression atoms. Plain-sympy components — the derivative - # blocks, which dominate the expression size — have no UW atoms, so - # the reveal / validate / xreplace / unwrap pipeline (≈5 full - # traversals per component) would be pure overhead: skip it entirely - # when there is nothing to lower. - from underworld3.function.expressions import UWexpression as _UWexpr - from underworld3.function.expressions import UWDerivativeExpression as _UWderiv - - _needs_lowering = lowered_fns is None and ( - isinstance(fn, (_UWexpr, _UWderiv)) - # `has` is a bare traversal (no atom-set build) — the atoms() - # form built a set of every node, which cost seconds per 100k-node - # Jacobian component (measured ~20 s on a large collision model). - or (hasattr(fn, "has") and fn.has(_UWexpr)) - or not isinstance(fn, (sympy.MatrixBase, sympy.MatrixExpr)) - ) - if _needs_lowering: - # Phase 1: reveal constants nested inside other UWexpressions, so - # the substitution below can reach them. A top-level - # xreplace missed constants inside template-wrapped - # parameters and baked them as C literals while the - # manifest listed them. - fn = _reveal_constants(fn) - - # A truly-constant atom the manifest does NOT know about would be - # silently folded to a literal in phase 3 — the manifest and the - # C source must never disagree (issue #302). - if constants_subs_map is not None and hasattr(fn, 'atoms'): - unmanifested = [ - a.name for a in _stable_sorted(_unique_symbols(fn)) - if isinstance(a, _UWexpr) - and _is_truly_constant(a, _UWexpr) - and a not in constants_subs_map - ] - if unmanifested: - raise RuntimeError( - f"JIT constants manifest is incomplete: constant expression(s) " - f"{unmanifested} appear in a kernel but have no constants[] " - f"slot — they would be baked into the C source (issue #302)." - ) - - # Phase 2: Substitute constant UWexpressions with _JITConstant symbols - # These survive into C code as constants[i] - if constants_subs_map and fn is not None: - try: - fn = _xreplace_shared(fn, constants_subs_map) if hasattr(fn, 'xreplace') else fn - except Exception: - pass - - # Phase 3: Unwrap remaining non-constant UWexpressions to numerical values - fn = underworld3.function.expressions.unwrap(fn, keep_constants=False, return_self=False) - - # A manifested constant surviving to here bypassed its constants[] - # slot and is about to be baked — refuse rather than freeze the - # parameter silently (issue #302). - if constants_subs_map and hasattr(fn, 'atoms'): - baked = [a.name for a in _stable_sorted(_unique_symbols(fn)) - if a in constants_subs_map] - if baked: - raise RuntimeError( - f"Manifested constant(s) {baked} were not routed through " - f"constants[] and would be baked into the C source " - f"(issue #302)." - ) + for index, fn_original in enumerate(fns): + unspellable = [] - if lowered_fns is not None: - pass # shaped by _jit_graph.lower_callbacks - elif isinstance(fn, sympy.vector.Vector): - fn = fn.to_matrix(mesh.N)[0 : mesh.dim, 0] - elif isinstance(fn, sympy.vector.Dyadic): - fn = fn.to_matrix(mesh.N)[0 : mesh.dim, 0 : mesh.dim] - else: - fn = sympy.Matrix([fn]) - - # === COORDINATE SYMBOL RECOVERY === - # When sympy.simplify() manipulates expressions containing coordinate - # symbols (BaseScalar/UWCoordinate), it may create NEW instances that - # lack the _ccodestr attribute set during mesh initialization. - # This commonly occurs with coordinate-dependent constitutive models - # (e.g., TransverseIsotropicFlowModel with a radial director). - # We recover _ccodestr from the coordinate's _id attribute. - from sympy.vector.scalar import BaseScalar - - free_syms = fn.free_symbols - for _t, _body in temporaries: - free_syms = free_syms | _body.free_symbols - free_syms = tuple(_stable_sorted(free_syms)) - for sym in free_syms: - if isinstance(sym, BaseScalar) and not hasattr(sym, '_ccodestr'): - idx = sym._id[0] # 0, 1, or 2 for x, y, z - system_name = str(sym._id[1]) - if 'Gamma' in system_name: - sym._ccodestr = f"petsc_n[{idx}]" - else: - sym._ccodestr = f"petsc_x[{idx}]" - - # === JIT VALIDATION GATEWAY === - # Check for symbols that cannot be converted to C code. - # Expected symbols (coordinates) have _ccodestr attribute set. - # Unexpected symbols indicate malformed expressions from user code. - unconvertible_symbols = [] - for sym in free_syms: - # Check if this symbol can be converted to C code - if not hasattr(sym, '_ccodestr'): - unconvertible_symbols.append(sym) - - if unconvertible_symbols: - # Build a helpful error message - sym_details = [] - for sym in unconvertible_symbols: - detail = f" - {sym} (type: {type(sym).__name__})" - if hasattr(sym, 'units'): - detail += f" [has units: {sym.units}]" - if hasattr(sym, 'value'): - detail += f" [value: {sym.value}]" - sym_details.append(detail) + def spell(leaf): + try: + return _spell_leaf(leaf, spellings, constants_subs_map) + except _Unspellable: + unspellable.append(leaf) + return f"?{sympy.srepr(leaf)}" - raise RuntimeError( - f"\n{'='*70}\n" - f"JIT COMPILATION ERROR: Expression contains unconvertible symbols\n" - f"{'='*70}\n\n" - f"The following symbols could not be converted to C code:\n" - + "\n".join(sym_details) + "\n\n" - f"This usually means:\n" - f" 1. A UWexpression or UWQuantity was not properly expanded\n" - f" 2. An arithmetic operation failed (e.g., Matrix * UWexpression)\n" - f" 3. A symbolic function is missing from the expression tree\n\n" - f"Expression index: {index}\n" - f"Original expression: {fn_original}\n" - f"After unwrap: {fn}\n\n" - f"TIP: Check that all expression operations (*, /, +, -) produce\n" - f"valid SymPy expressions. For example, ensure scalar * Matrix\n" - f"and not Matrix * scalar when using UWexpression objects.\n" - f"{'='*70}" - ) + temporaries, fn = _jit_graph.emit(lowered_fns[index], spell) + if unspellable: + raise RuntimeError(_unconvertible_message( + _stable_sorted(set(unspellable)), index, fn_original)) if verbose: - print("Processing JIT {:4d} / {}".format(index, fn)) - # Enhanced debugging output for remaining (valid) free symbols - if free_syms: - print(" Free symbols (all convertible):") - for sym in free_syms: - print(f" - {sym} (type: {type(sym).__name__}, _ccodestr: {getattr(sym, '_ccodestr', 'N/A')})") + # the kernel as mathematics (named quantities as nodes); the C is in the header + print("Processing JIT {:4d} / {}".format(index, lowered_fns[index])) out = sympy.MatrixSymbol("out", *fn.shape) - - # CSE before printing: shared subexpressions become ``double xN = ...;`` - # temps evaluated in dependency order, so the generated C — and hence - # the codegen time, gcc memory/time, and .so size — collapses on large - # expressions (measured: monster Jacobian output ~460k nodes -> ~30k). - # Semantics-preserving: temps are exact aliases of repeated - # subexpressions, so the generated kernel evaluates identical values. - # Opt in with UW_JIT_CSE=1 (default is off to preserve original behavior). - if temporaries: - # the graph route: one temporary per distinct computation, then the - # outputs in terms of them (#823) - _temp_code = [] - for _t, _body in temporaries: - _code = printer.doprint(_body) - if _code.startswith("// Not supported in C:"): - _temp_code = None - eqn = ("eqn_" + str(index), _code) - break - _temp_code.append(f"const double {_t._ccodestr} = {_code};") - if _temp_code is not None: - eqn = ("eqn_" + str(index), - "\n".join(_temp_code) + "\n" + printer.doprint(fn, out)) - elif lowered_fns is None and os.environ.get("UW_JIT_CSE") in ("1", "true", "True", "yes", "YES"): - from sympy.simplify.cse_main import cse - from sympy.vector.scalar import BaseScalar - - _repl, _red = cse([fn]) - if _repl: - # cse may mint NEW coordinate instances (BaseScalar / - # UWCoordinate wrappers) that lack the mesh-set _ccodestr; - # recover it from their _id (same scheme as the - # COORDINATE SYMBOL RECOVERY above). - def _patch_coords(expr): - for _sym in set(expr.free_symbols): - _target = getattr(_sym, "_original_base_scalar", _sym) - if isinstance(_target, BaseScalar) and not hasattr( - _target, "_ccodestr" - ): - _idx = _target._id[0] - _sys = str(_target._id[1]) - _target._ccodestr = ( - f"petsc_n[{_idx}]" - if "Gamma" in _sys - else f"petsc_x[{_idx}]" - ) - - for _t_sym, _t_expr in _repl: - _patch_coords(_t_expr) - _patch_coords(_red[0]) - - _temp_code = "\n".join( - "double {} = {};".format( - printer.doprint(t_sym), printer.doprint(t_expr) - ) - for t_sym, t_expr in _repl - ) - _red_code = printer.doprint(_red[0], out) - if _red_code.startswith("// Not supported in C:"): - eqn = ("eqn_" + str(index), _red_code) - else: - eqn = ("eqn_" + str(index), _temp_code + "\n" + _red_code) - else: - eqn = ("eqn_" + str(index), printer.doprint(fn, out)) - else: - eqn = ("eqn_" + str(index), printer.doprint(fn, out)) + eqn = ("eqn_" + str(index), _print_kernel(printer, temporaries, fn, out)) if eqn[1].startswith("// Not supported in C:"): spliteqn = eqn[1].split("\n") @@ -1961,67 +1580,3 @@ def load_dynamic(name, path): ) return module, tmpdir - - -@timing.routine_timer_decorator -def _createext( - name, - mesh: underworld3.discretisation.Mesh, - callbacks: JITCallbackSet, - primary_field_list, - constants_subs_map: Optional[dict] = None, - verbose: Optional[bool] = False, - debug: Optional[bool] = False, - debug_name=None, -): - """Thin wrapper: generate source, compile, stash in ``_ext_dict[name]``. - - Retained for backwards compatibility with :func:`getext`. New code - should call :func:`generate_c_source` and :func:`compile_and_load` - directly — splitting the two phases is what makes cache keys on the - generated C source possible. - """ - modname, codeguys, diag = generate_c_source( - name, - mesh, - callbacks, - primary_field_list, - constants_subs_map=constants_subs_map, - verbose=verbose, - debug=debug, - debug_name=debug_name, - ) - module, tmpdir = compile_and_load(modname, codeguys, verbose=verbose) - _ext_dict[name] = module - - if underworld3.mpi.rank == 0 and verbose: - randstr = diag["randstr"] - print(f"Location of compiled module: {str(tmpdir)}") - print(f"{randstr} Equation count - {diag['eqn_count']}", flush=True) - print( - f"{randstr} {diag['count_residual_sig']:5d} residuals: " - f"{diag['residual_equations'][0]}:{diag['residual_equations'][1]}", - flush=True, - ) - print( - f"{randstr} {diag['count_bc_sig']:5d} boundaries: " - f"{diag['boundary_equations'][0]}:{diag['boundary_equations'][1]}", - flush=True, - ) - print( - f"{randstr} {diag['count_jacobian_sig']:5d} jacobians: " - f"{diag['jacobian_equations'][0]}:{diag['jacobian_equations'][1]}", - flush=True, - ) - print( - f"{randstr} {diag['count_bd_residual_sig']:5d} boundary_res: " - f"{diag['boundary_residual_equations'][0]}:{diag['boundary_residual_equations'][1]}", - flush=True, - ) - print( - f"{randstr} {diag['count_bd_jacobian_sig']:5d} boundary_jac: " - f"{diag['boundary_jacobian_equations'][0]}:{diag['boundary_jacobian_equations'][1]}", - flush=True, - ) - - return diff --git a/tests/test_0022_unwrap_memoised_matches_fixed_point.py b/tests/test_0022_unwrap_memoised_matches_fixed_point.py index be7477277..e351d3f09 100644 --- a/tests/test_0022_unwrap_memoised_matches_fixed_point.py +++ b/tests/test_0022_unwrap_memoised_matches_fixed_point.py @@ -1,5 +1,8 @@ -"""The memoised unwrap, the memoised sqrt guard and the identity-walk symbol scan (#823) -give exactly what the algorithms they replace gave. +"""The memoised unwrap (#823) gives exactly what the algorithm it replaces gave. + +(The memoised sqrt guard, the identity-walk symbol scan and the shared xreplace were +tested here too; they went with the expanded-tree JIT route, and the graph route's +guard is held to the guarded tree by test_0024.) The old ``unwrap_expression`` iterated ``subs`` passes to a fixed point, and in ``symbolic_keep_constants`` mode ran ``_is_truly_constant`` (itself a full unwrap) for every @@ -101,59 +104,6 @@ def test_memoised_unwrap_is_the_fixed_point(laws, mode): assert sympy.srepr(new) == sympy.srepr(old), (name, mode) -def test_the_jacobian_sqrt_guard_matches_replace(laws, monkeypatch): - """The expanded-tree route's guard. The graph route guards each node body - instead; test_0024 holds it to this one by value.""" - from underworld3.cython.generic_solvers import _jacobian_unwrap - - monkeypatch.setenv("UW_JIT_GRAPH", "0") - - eps2 = sympy.Float(1.0e-36) - - def old_guard(e): - return e.replace( - lambda n: (n.is_Pow and n.exp.is_Rational and n.exp.q == 2 - and n.args[0].free_symbols), - lambda n: sympy.Pow(n.args[0] + eps2, n.exp)) - - for name, expr in laws: - new = _jacobian_unwrap(expr) - old = _unwrap_each(expr, lambda e: old_guard( - _fixed_point_unwrap(e, "symbolic_keep_constants"))) - assert sympy.srepr(new) == sympy.srepr(old), name - - -def test_the_identity_walk_finds_the_same_symbols(laws): - from underworld3.utilities._jitextension import _unique_symbols - - for name, expr in laws: - for e in (expr, ex.unwrap_expression(expr, mode="nondimensional") - if not isinstance(expr, sympy.MatrixBase) - else expr.applyfunc(lambda a: ex.unwrap_expression(a, mode="nondimensional"))): - assert _unique_symbols(e) == set(e.atoms(sympy.Symbol)), name - - -def test_the_shared_xreplace_matches_xreplace(laws): - """The constants[] substitution in the JIT lowering visits each node object once; - it must build exactly what xreplace builds, the generated C and the cache key - both depend on it.""" - from underworld3.utilities._jitextension import _xreplace_shared, _unique_symbols - - swapped = 0 - for name, expr in laws: - unwrapped = _unwrap_each(expr, lambda e: ex.unwrap_expression( - e, mode="symbolic_keep_constants")) - held = sorted((a for a in _unique_symbols(unwrapped) - if isinstance(a, ex.UWexpression)), key=str) - if not held: - continue # k(u) = 1 + u**2 has no parameter to swap - rule = {a: sympy.Symbol(f"slot_{k}") for k, a in enumerate(held)} - assert sympy.srepr(_xreplace_shared(unwrapped, rule)) == \ - sympy.srepr(unwrapped.xreplace(rule)), name - swapped += 1 - assert swapped >= 7 # the VP (three laws, viscosity and flux) and VEP - - def test_each_atom_is_tested_for_constancy_once(monkeypatch): """The cost that #823 removed: the fixed-point passes asked whether each atom was a constant on every pass, and each answer was itself a full unwrap. On a chain of diff --git a/tests/test_0023_field_realness_and_derivatives.py b/tests/test_0023_field_realness_and_derivatives.py index 3a1401816..a7e7b4b55 100644 --- a/tests/test_0023_field_realness_and_derivatives.py +++ b/tests/test_0023_field_realness_and_derivatives.py @@ -251,7 +251,6 @@ def test_parameters_given_with_units_are_known_real(): source on the Spiegelman notch (#823). A UWQuantity is real when its value is.""" from underworld3.cython.generic_solvers import _jacobian_unwrap from underworld3.function.expressions import UWexpression - from underworld3.utilities._jitextension import _unique_symbols uw.reset_default_model() orchestration_model = uw.get_default_model() @@ -282,7 +281,7 @@ def test_parameters_given_with_units_are_known_real(): assert nan.is_extended_real is None and nan.is_finite is None flux = _jacobian_unwrap(stokes.constitutive_model.flux) - held = [a for a in _unique_symbols(flux) if isinstance(a, UWexpression)] + held = [a for a in flux.free_symbols if isinstance(a, UWexpression)] assert held, "the Newton flux keeps its constant parameters as atoms" unknown = [a for a in held if a.is_extended_real is not True] assert not unknown, unknown @@ -316,10 +315,9 @@ def yielding_box(name, delta): uw.reset_default_model() ramped, v_ramped = yielding_box("0023s", 1.0) from underworld3.cython.generic_solvers import _jacobian_unwrap - from underworld3.utilities._jitextension import _unique_symbols sharpness = ramped.constitutive_model._get_yield_sharpness() - assert sharpness in _unique_symbols(_jacobian_unwrap(ramped.constitutive_model.flux)) + assert sharpness in _jacobian_unwrap(ramped.constitutive_model.flux).free_symbols ramped.solve() assert ramped.snes.getConvergedReason() > 0 at_one = np.array(v_ramped.array) diff --git a/tests/test_0024_jit_graph_lowering.py b/tests/test_0024_jit_graph_lowering.py index ff59cb312..c93a9f278 100644 --- a/tests/test_0024_jit_graph_lowering.py +++ b/tests/test_0024_jit_graph_lowering.py @@ -145,7 +145,19 @@ def test_a_changed_body_is_lowered_afresh_without_clearing_the_cache(box): assert jg.expand_nodes(diff_wrt_field(second, u)) == 3 * c2 * u ** 2 + c1 -def test_the_guarded_lowering_is_the_guarded_tree(box, monkeypatch): +def _guarded_tree(e): + """The expanded-tree Newton source the graph replaced: every non-constant atom + expanded, then 1e-36 added to the base of every half-integer power with free + symbols, by SymPy's own replace.""" + eps2 = sympy.Float(1.0e-36) + tree = ex.unwrap_expression(e, mode="symbolic_keep_constants") + return tree.replace( + lambda n: (n.is_Pow and n.exp.is_Rational and n.exp.q == 2 + and n.args[0].free_symbols), + lambda n: sympy.Pow(n.args[0] + eps2, n.exp)) + + +def test_the_guarded_lowering_is_the_guarded_tree(box): """The Newton source on the graph (the sqrt guard in each node body) evaluates to the guarded tree, and so does its tangent, including at a state of rest where the unguarded tangent is 0/0.""" @@ -162,9 +174,7 @@ def test_the_guarded_lowering_is_the_guarded_tree(box, monkeypatch): cm.Parameters.yield_stress_min = uw.expression(r"\tau_{0024e}", 0.01) flux = sympy.Matrix(stokes.F1.sym) - monkeypatch.setenv("UW_JIT_GRAPH", "0") - tree = _jacobian_unwrap(flux) - monkeypatch.setenv("UW_JIT_GRAPH", "1") + tree = flux.applyfunc(_guarded_tree) graph = _jacobian_unwrap(flux) assert graph.atoms(jg._KernelNode), "the Newton source was not lowered" @@ -182,9 +192,9 @@ def test_the_guarded_lowering_is_the_guarded_tree(box, monkeypatch): assert abs(vg - vt) <= 1.0e-12 * max(abs(vt), 1.0), (state, i, j) -def _header(solver, monkeypatch, route): - """The generated header of ``solver`` on ``route``, with the module name and the - symbol prefix canonicalised as ``getext`` canonicalises them.""" +def _header(solver, monkeypatch): + """The generated header of ``solver``, with the module name and the symbol prefix + canonicalised as ``getext`` canonicalises them.""" import underworld3.utilities._jitextension as jx seen = {} @@ -197,7 +207,6 @@ def keep(*a, **k): return modname, codeguys, diag monkeypatch.setattr(jx, "generate_c_source", keep) - monkeypatch.setenv("UW_JIT_GRAPH", route) solver.is_setup = False solver._setup_pointwise_functions() return seen["h"] @@ -222,7 +231,7 @@ def test_the_emitted_source_does_not_depend_on_what_came_before(monkeypatch): uw.reset_default_model() mesh = uw.meshing.UnstructuredSimplexBox( minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) - first = _header(_poisson("0024f", mesh), monkeypatch, "1") + first = _header(_poisson("0024f", mesh), monkeypatch) for k in range(7): uw.expression(rf"junk_{{0024f,{k}}}", float(k), "preamble") other = uw.meshing.UnstructuredSimplexBox( @@ -230,12 +239,14 @@ def test_the_emitted_source_does_not_depend_on_what_came_before(monkeypatch): uw.discretisation.MeshVariable("J0024f", other, 1, degree=1) mesh2 = uw.meshing.UnstructuredSimplexBox( minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) - second = _header(_poisson("0024f", mesh2), monkeypatch, "1") + second = _header(_poisson("0024f", mesh2), monkeypatch) assert first == second assert "uwt_0" in first, "the kernel has no temporaries: nothing was lowered" -def test_a_law_with_no_named_quantity_emits_the_tree_route_source(monkeypatch): +def test_a_law_with_no_named_quantity_has_no_temporaries(monkeypatch): + """A constant law has no node to lower: its kernels are the outputs alone, as the + expanded tree printed them.""" uw.reset_default_model() mesh = uw.meshing.UnstructuredSimplexBox( minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) @@ -244,13 +255,14 @@ def test_a_law_with_no_named_quantity_emits_the_tree_route_source(monkeypatch): pois.constitutive_model = uw.constitutive_models.DiffusionModel pois.constitutive_model.Parameters.diffusivity = 1.0 pois.f = 2.0 - assert _header(pois, monkeypatch, "0") == _header(pois, monkeypatch, "1") + header = _header(pois, monkeypatch) + assert "out[0]" in header and "uwt_" not in header -def test_the_manifest_from_the_leaves_is_the_scanned_manifest(box): - """The constants the lowered kernels read are the constants the expressions hold, - in the same slots.""" - from underworld3.utilities._jitextension import _extract_constants, _manifest_from +def test_the_manifest_is_the_constants_the_kernels_read(box): + """Every constant atom a lowered kernel reads, at any depth, has a slot, ordered by + name; a constant inside a constant is folded into its holder's slot.""" + from underworld3.utilities._jitextension import _extract_constants mesh, T, v = box u = T.sym[0] @@ -261,10 +273,9 @@ def test_the_manifest_from_the_leaves_is_the_scanned_manifest(box): a = uw.expression(r"a_{0024h}", c1 * u + c4 * u ** 2, "inner") b = uw.expression(r"b_{0024h}", a / (c2 + a ** 2), "outer") fns = (b, sympy.Matrix([[b * u, a]])) - scanned, _ = _extract_constants(fns, mesh) - leaves, _ = _manifest_from(jg.constant_leaves(jg.lower_callbacks(fns, mesh))) - assert [e for _, e in leaves] == [e for _, e in scanned] - assert len(scanned) == 3 # c4 is a slot of its own; c3 is folded into it + manifest, placeholders = _extract_constants(fns, mesh) + assert [e for _, e in manifest] == [c1, c2, c4] + assert len(set(placeholders.values())) == 3 def test_a_repeated_condition_stays_a_condition(box): From 8df29f4dbc95acb69898e1707efb9e10bf9c46a2 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 17:49:46 +1100 Subject: [PATCH 09/14] Docs for the graph JIT (#823, tier 2, step 4) - jit-cache.md: what is compiled (one temporary per named quantity, canonical source); the MPI section was stale before this change (it said a cross-rank mismatch raises and every rank compiles): adopting rank 0's source, the collective compile decision. - expressions-functions.md: the JIT no longer unwraps whole kernels; the unwrappers expand graph nodes. - jacobian-consistent-tangent.md and the plasticity guide: the Newton source is built from nodes, the same tangent without the expanded tree. - The design note: steps 2-4 implemented; one change, UW_JIT_CSE retired (decisions); the prototype scripts' removal; the final line counts. - _jacobian_unwrap's docstring pointed at a design doc that does not exist. Underworld development team with AI support from Claude Code --- .../design/jacobian-consistent-tangent.md | 6 +++ .../design/jit-shared-graph-codegen.md | 50 ++++++++++++------- docs/developer/guides/plasticity-solvers.md | 5 ++ .../subsystems/expressions-functions.md | 7 ++- docs/developer/subsystems/jit-cache.md | 34 +++++++++---- .../cython/petsc_generic_snes_solvers.pyx | 2 +- 6 files changed, 72 insertions(+), 32 deletions(-) diff --git a/docs/developer/design/jacobian-consistent-tangent.md b/docs/developer/design/jacobian-consistent-tangent.md index eff27c366..343c30dff 100644 --- a/docs/developer/design/jacobian-consistent-tangent.md +++ b/docs/developer/design/jacobian-consistent-tangent.md @@ -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` diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index 368422810..d3af2f739 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -1,8 +1,8 @@ # Generating JIT kernels from the shared expression graph -**Status**: Proposed, 2026-10-07. Staging steps 2 and 3 implemented behind the private -switch `UW_JIT_GRAPH=1` on `feature/jit-graph-codegen`, 2026-10-08, and measured against -tier 1 end to end ({ref}`jit-graph-in-the-library`). +**Status**: Implemented on `feature/jit-graph-codegen`, 2026-10-08: the graph route is +the only JIT route (staging steps 2–4). Steps 2 and 3 were first measured against tier 1 +end to end while both routes existed behind a switch ({ref}`jit-graph-in-the-library`). Tier 2 of [#823](https://github.com/underworldcode/underworld3/issues/823). Tier 1 ([#830](https://github.com/underworldcode/underworld3/pull/830), merged 2026-10-07) is the fix for today: the memoised @@ -303,9 +303,10 @@ repair stays. ## The prototype agrees with the library to round-off -`scripts/sessions/jit_graph/kernel_graph.py` is the lowering; -`scripts/sessions/jit_graph/graph_vs_library.py` builds three kernels of a fixture twice -and compiles each into a C function of the same leaves: +`scripts/sessions/jit_graph/kernel_graph.py` was the lowering; +`scripts/sessions/jit_graph/graph_vs_library.py` built three kernels of a fixture twice +and compiled each into a C function of the same leaves. (Both scripts were removed in +staging step 4, when `_jit_graph.py` replaced them; they are in commit 5a0d2c46.) - the residual flux `F1`, lowered with plain nodes; the library route unwraps it as `getext()` does (keep-constants); @@ -381,11 +382,11 @@ exactly 0 or 1 — but a model with a projected or higher-degree material field (jit-graph-in-the-library)= ## In the library, against tier 1 -Steps 2 and 3 are implemented on `feature/jit-graph-codegen` +Steps 2 and 3 were implemented on `feature/jit-graph-codegen` (`src/underworld3/utilities/_jit_graph.py`, with branches in `getext()`, -`generate_c_source()`, `_jacobian_unwrap` and the two unwrappers), selected by -`UW_JIT_GRAPH=1`. Unset, every path is tier 1's. Both routes therefore run on one build, -and each fixture is run once per route in a fresh process with the JIT cache off +`generate_c_source()`, `_jacobian_unwrap` and the two unwrappers), selected by a private +switch, `UW_JIT_GRAPH=1`; unset, every path was tier 1's. Both routes therefore ran on +one build, and each fixture was run once per route in a fresh process with the JIT cache off (`scripts/sessions/jit_graph/route_ab.py`): the solver's whole pointwise setup, a Newton solve, then the residual and the Jacobian assembled repeatedly at the solution. The fixtures are those of the prototype, with the box ones solved as a lid-driven box @@ -596,9 +597,10 @@ multiplied by a name that holds its reciprocal. ## Benchmark plan -During development both routes run in one process on identical inputs, selected by a -private switch, so every comparison is A against B on one build. The switch and the -expanded route are removed before the change merges. The fixtures are built in +While steps 2 and 3 were developed, both routes ran on one build, selected by a private +switch, so every comparison was A against B on identical inputs. Step 4 removed the +switch and the expanded route; the session scripts now take the route from the build, +and the tier 1 JIT is measured in a build of `development`. The fixtures are built in `scripts/sessions/jit_graph/` so that the measurements can be repeated from the repository; the notch is the one exception, and its driver is named. @@ -646,13 +648,20 @@ repository; the notch is the one exception, and its driver is named. blocks; in the same change the two unwrappers learn to expand them and `test_0022`'s tree-shaped guard test is pinned to the tree route. - Steps 2 and 3 are implemented together behind `UW_JIT_GRAPH=1` (2026-10-08). + Steps 2 and 3 were implemented together behind `UW_JIT_GRAPH=1` (2026-10-08). 4. The expanded route and the switch are removed, together with the two functions that mirror it and have no callers (`prepare_for_cache_key`, `_createext`), and `jit-cache.md`, `expressions-functions.md` and `jacobian-consistent-tangent.md` are updated. `jit-cache.md` already describes the cross-rank check as an abort and every rank as compiling; both stopped being true before this change. + Done, 2026-10-08. The field classes' patched C names (`ccode_patch_fns`), the + coordinate recovery and the unconvertible-symbol scan were replaced by an explicit map + from leaf to C built for each compile (`_leaf_spellings`, `_spell_leaf`), and the + opt-in `UW_JIT_CSE` path, which served only the expanded tree, was retired. Against + `development`, the library source gains 692 lines and loses 639: `_jit_graph.py` is + 457 of the gain, and `_jitextension.py` is about 400 lines shorter. + Each step is benchmarked against the one before it and reviewed adversarially before the next begins. @@ -679,12 +688,15 @@ the next begins. C. After the change the C is canonical, bit-identical across hash seeds, preambles, re-declarations and processes by construction. +- **2026-10-08: steps 2–4 land as one change** (Louis: "crash through, then compare + with the existing, improved JIT"). The comparison is against `development` with tier 1 + and `-fno-math-errno` (#834, PR #835), the competitor at its best. +- **2026-10-08: `UW_JIT_CSE` is retired** with the expanded tree it served. + ## Questions for the maintainer -1. One PR for steps 2–4, or one PR per step. -2. Whether the generated C should carry each temporary's display name as a comment. +1. Whether the generated C should carry each temporary's display name as a comment. It makes kernels readable; it also puts display names into the cache key, so a `rename()` recompiles. -3. Whether the adjoint's own unwrap (`_peel_except`, on the adjoint branches) adopts the - nodes in this change or after it. -4. Whether `UW_JIT_CSE` is retired once this lands. +2. Whether the adjoint's own unwrap (`_peel_except`, on the adjoint branches) adopts the + nodes when those branches rebase onto this change. diff --git a/docs/developer/guides/plasticity-solvers.md b/docs/developer/guides/plasticity-solvers.md index 8f03f37c4..3693f847f 100644 --- a/docs/developer/guides/plasticity-solvers.md +++ b/docs/developer/guides/plasticity-solvers.md @@ -124,6 +124,11 @@ 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. + 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 diff --git a/docs/developer/subsystems/expressions-functions.md b/docs/developer/subsystems/expressions-functions.md index 9abad6646..2e7415677 100644 --- a/docs/developer/subsystems/expressions-functions.md +++ b/docs/developer/subsystems/expressions-functions.md @@ -147,9 +147,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) diff --git a/docs/developer/subsystems/jit-cache.md b/docs/developer/subsystems/jit-cache.md index d2fca7255..4fc68d5a6 100644 --- a/docs/developer/subsystems/jit-cache.md +++ b/docs/developer/subsystems/jit-cache.md @@ -33,6 +33,18 @@ 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 + +`getext()` 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 quadrature point, 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. + ## Cache key The key is the SHA-256 of the **canonical** generated C source plus an @@ -104,20 +116,18 @@ 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 | @@ -130,7 +140,9 @@ performs the cold compile but only rank 0 publishes the result. ## 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 diff --git a/src/underworld3/cython/petsc_generic_snes_solvers.pyx b/src/underworld3/cython/petsc_generic_snes_solvers.pyx index 47c6790b2..df0868925 100644 --- a/src/underworld3/cython/petsc_generic_snes_solvers.pyx +++ b/src/underworld3/cython/petsc_generic_snes_solvers.pyx @@ -96,7 +96,7 @@ def _jacobian_unwrap(expr): tangent never calls this function. See ``docs/developer/design/jit-shared-graph-codegen.md`` and - ``docs/developer/design/jacobian-unwrap-constants-bug.md``. + ``docs/developer/design/jacobian-consistent-tangent.md``. """ from underworld3.utilities import _jit_graph From be26dd84e886dede1db2a91d362aa12eb2d4f96f Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 18:52:40 +1100 Subject: [PATCH 10/14] Adversarial review fixes for the graph JIT (#823, tier 2) Correctness (each with a test in test_0024, shown failing first): - a mesh.X coordinate beside a field in one named quantity lost its explicit derivative: the per-body cse rebuilt the UWCoordinate with a cloned coordinate system. Leaves and child nodes are hidden behind placeholders while cse runs, and lowering finds UWCoordinates by type; - a matrix- or vector-valued atom became a scalar node and failed to compile; such bodies are expanded in place (_is_scalar); - constancy was decided by a complete unwrap of each atom, exponential in the nesting depth; it is decided bottom-up on the graph with the same rule; - a number symbol (EulerGamma, Catalan) in a temporary wrote its declaration into the initialiser; declarations are hoisted, once each. test_0024 also finds that the expanded tree escaped its own sqrt guard on a power law on a named invariant (a NaN Newton flux at rest) where the graph does not. Tests and docs: - test_0105 sweeps hash seeds over a Newton viscoplastic law that emits temporaries; test_0022's guard identity test is restored against the graph's guard; test_0024 adds a power law, a twelve-layer law, coordinate spelling, matrix atoms, a constant law with named constants. - Nodes print as their quantity's name (str, latex); srepr and the C are unchanged. - One manifest helper (_manifest_of); generate_c_source takes the lowered callbacks and the substitution map; dead fallbacks, imports and the _ccodestr constancy branch removed; the stale allowlist entry dropped. - The solver's consistent_jacobian docstring no longer promises bit-identical Picard kernels; comments and doc pointers describe the graph route; expressions and mathematical-objects docs no longer describe the two-phase unwrap; the design note is consistent on #752, lists what was not measured, and records the review findings. Underworld development team with AI support from Claude Code --- .../UW3_Developers_MathematicalObjects.md | 17 +- .../design/jit-shared-graph-codegen.md | 95 +++++----- .../subsystems/expressions-functions.md | 27 +-- docs/developer/subsystems/jit-cache.md | 3 +- scripts/deprecated_pattern_allowlist.txt | 1 - .../sessions/jit_graph/graph_size_probe.py | 4 +- .../jit_graph/library_setup_profile.py | 4 +- .../sessions/jit_graph/rank_agreement_752.py | 3 +- .../sessions/jit_graph/route_ab_compare.py | 5 +- .../jit_graph/route_assemble_compare.py | 6 +- src/underworld3/constitutive_models.py | 8 +- .../cython/petsc_generic_snes_solvers.pyx | 51 ++--- src/underworld3/utilities/_jit_graph.py | 122 +++++++++--- src/underworld3/utilities/_jitextension.py | 115 +++++------ ...022_unwrap_memoised_matches_fixed_point.py | 27 ++- tests/test_0024_jit_graph_lowering.py | 178 +++++++++++++++++- tests/test_0103_jit_rampable_constants.py | 5 + .../test_0105_jit_source_seed_independence.py | 66 +++++++ 18 files changed, 543 insertions(+), 194 deletions(-) diff --git a/docs/developer/UW3_Developers_MathematicalObjects.md b/docs/developer/UW3_Developers_MathematicalObjects.md index 95551a1f2..ccdccf238 100644 --- a/docs/developer/UW3_Developers_MathematicalObjects.md +++ b/docs/developer/UW3_Developers_MathematicalObjects.md @@ -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]; diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index d3af2f739..1c801c898 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -54,10 +54,10 @@ erased. Five properties follow: 5. **The generated C is readable.** One line per named quantity, so a kernel can be checked by eye and a wrong one can be found. -Today's JIT has none of these. Tier 1 makes it fast enough and correct on the cases we -have found, with patches placed where each failure surfaced. +The JIT before tier 2 had none of these. Tier 1 made it fast enough and correct on the +cases we had found, with patches placed where each failure surfaced. -## The current pipeline expands every named sub-expression +## The tier 1 pipeline expanded every named sub-expression A solver hands the JIT a `JITCallbackSet` of residual and Jacobian expressions. `getext()` turns them into one C function per callback. For a Newton tangent: @@ -216,12 +216,14 @@ when it emits the kernel. The continuation blend contains both, and each lowers own temporaries. The guard moves from the expanded tree to each node body, under the same rule. The two -placements could differ only where SymPy merges powers across a sub-expression -boundary: a bare `sqrt(g)**(-2/3)` becomes `g**(-1/3)`, which is no longer a -half-integer power. A real law raises a quotient — $(\dot\varepsilon_{II}/\dot\varepsilon_0)^{1/n-1}$ — -which SymPy does not merge, and on a power law over a named invariant, with $n$ as a -constant atom or as the number 3, both routes give a finite Newton tangent at a state of -rest, equal entry for entry. +placements differ where SymPy merges powers across a sub-expression boundary: a bare +`sqrt(g)**(-2/3)` becomes `g**(-1/3)`, which is no longer a half-integer power, and the +tree's guard misses it. A law that raises a quotient, +$(\dot\varepsilon_{II}/\dot\varepsilon_0)^{1/n-1}$, is not merged, and both routes give +a finite Newton tangent at a state of rest, equal entry for entry. A power law on the +bare named invariant, $\eta = \dot\varepsilon_{II}^{\,1/n-1}$, is merged into +$g^{(1/n-1)/2}$: the tree's Newton flux is NaN at rest and the graph's is finite, because +$\dot\varepsilon_{II}$ stays a node whose body carries the guard (`test_0024`). ### Each lowering is read back by `getext()` from its nodes @@ -235,8 +237,8 @@ temporaries by the C they compute. The solvers are therefore unchanged apart fro `_jacobian_unwrap`; one context per setup would only save the constancy test of an atom that both lowerings meet. -`getext()` no longer runs `unwrap_expression` over a whole kernel. Its present phases — -reveal the constants, substitute them, unwrap the rest — become: lower atoms to nodes, +`getext()` no longer runs `unwrap_expression` over a whole kernel. Its tier 1 phases — +reveal the constants, substitute them, unwrap the rest — became: lower atoms to nodes, collect the manifest from the leaves, emit. ### A node expands to its body on request @@ -267,9 +269,9 @@ Applications with equal keys compute the same C and share one temporary. The temporaries are written so that each follows the ones it uses, ties broken by key: ```c -const double t0 = ...; /* one line per distinct computation, evaluated once */ -const double t1 = ...; -out[0] = ...; /* outputs in terms of t0, t1, ... and the leaves */ +const double uwt_0 = ...; /* one line per distinct computation, evaluated once */ +const double uwt_1 = ...; +out[0] = ...; /* outputs in terms of uwt_0, uwt_1, ... and the leaves */ ``` The generated source is therefore a function of the mathematics and of the kernel's data @@ -280,7 +282,7 @@ unconvertible symbols and integration-point derivatives. ### The constants manifest is built from the leaves The manifest is the set of constant atoms among the leaves of the emitted kernels, -tested with the same predicate (`_is_truly_constant`, memoised once per setup) and +tested with the same predicate (`_is_truly_constant`, memoised once per lowering) and ordered by the same key (name, then creation order). It contains every slot today's manifest contains. It can contain more: where the tree cancels a constant across a name boundary ($A = c\,x$, then $A/c$), the graph keeps it, and the constant keeps its slot. @@ -297,9 +299,10 @@ Underworld version. The prototype's generated C is byte-identical under `PYTHONHASHSEED` 0, 1 and 2, after a preamble that creates extra objects before the law (shifting every creation counter), and when the law is declared twice in one process, as a re-run notebook cell does. -Whether canonical emission also removes the cross-rank disagreement of #752 is a -hypothesis we test (np ≥ 3, counting how often `_agree_source_across_ranks` repairs); the -repair stays. +Canonical emission makes the source the same on every rank by construction, which is +#752's defect. The #752 fixture no longer disagrees on either route (0 of 10 runs at +np = 2, 3 and 4), so the fixture cannot show the difference; the repair stays as a +guard. ## The prototype agrees with the library to round-off @@ -401,7 +404,7 @@ session's seven solver processes, so the times are indicative. Tree first, graph | generated C, all modules of the solve | 4.0 MB / 22 KB | 380 KB / 40 KB | 104 KB / 16 KB | 42 KB / 31 KB | 28 KB / 14 KB | 10.7 KB, byte-identical | | Jacobian assembly, ms | 180 / 132 | 2.40 / 2.19 | 1.58 / 1.45 | 2.23 / 2.22 | 1.48 / 1.40 | 1.39 / 1.44 | | residual assembly, ms | 21.4 / 18.7 | 0.62 / 0.59 | 0.35 / 0.33 | 0.61 / 0.60 | 0.32 / 0.32 | 0.33 / 0.33 | -| Newton iterations, nonlinear / linear | 75 / 496 and 57 / 386 (see below) | 6 / 6, both | 30 / 30, both (limit) | 2 / 2, both | 16 / 16, both | 1 / 1, both | +| Newton iterations, nonlinear / linear | 75 / 496 and 57 / 386 (a property of the problem: 60–111 and 43–86 under round-off-sized perturbations) | 6 / 6, both | 30 / 30, both (limit) | 2 / 2, both | 16 / 16, both | 1 / 1, both | The compile time is not shown separately: it counts every module the solve builds (the VEP's history projections among them), and on these fixtures it is 1–4 s on either @@ -464,9 +467,11 @@ problem, not of the route. **The test suite passes on the graph route.** Every serial batch of `scripts/test.sh` (tier A and B, levels 1 to 3) run with `UW_JIT_GRAPH=1`: 3,023 passed, none failed. The -first run found two defects, both in the fault-network laws and both fixed with a test: -a repeated Piecewise condition shared as a node (a value, which Piecewise refuses as a -condition), and a coordinate leaf without the C name the mesh sets. +first run found two defects, both in the fault-network laws, each fixed with a unit test +in `test_0024`: a repeated Piecewise condition shared as a node (a value, which +Piecewise refuses as a condition), and a coordinate leaf without the C name the mesh +sets. After step 4 removed the tree route: 3,082 passed, none failed; and again after the +second review's fixes: SUITE_FINAL. **The source is canonical.** On the notch, box and VEP fixtures the graph route's C is byte-identical under `PYTHONHASHSEED` 0, 1 and 2; `test_0024` checks that a law @@ -492,9 +497,11 @@ faster: the notch's pointwise setup takes 2.5 s against 15.9 s, its C is 22 KB a emits byte-identical C, because a constant law has no node to lower. The assembly gain is smaller than the prototype's kernel timings suggested (an eighth of the time per call) because the pointwise kernel is one part of the assembly, beside quadrature and -the element loop. On Linux with gcc the ratios are the same (below). Still to measure: -an idle machine with repeated runs, a three-dimensional fixture, and the small kernels -of `Integral` and `BdIntegral`. +the element loop. On Linux with gcc the ratios are the same. Not measured: +an idle machine with repeated runs, a three-dimensional fixture, the small kernels +of `Integral` and `BdIntegral`, a Newton source that differs from its residual +(`set_jacobian_F1_source`), and SolCx with Nitsche free-slip; the test suite exercises +the last three but does not time them. **More robust.** Each failure class below needed a patch in tier 1, placed where it surfaced; in the graph it cannot arise, or arises only in one body: @@ -503,24 +510,25 @@ surfaced; in the graph it cannot arise, or arises only in one body: |---|---|---|---| | SymPy's automatic algebra on an expanded base is slow (`im()` inside `Pow.__new__`, 40 s on the notch) | a constant atom added to the power-mean law; realness for unit-carrying parameters | bases are bodies, not trees | graph lowering took 0.04 s on every tier 1 commit, including those where the library took 40 s | | realness turns `sqrt(x**2)` into `Abs`, whose derivative with respect to a field SymPy leaves unevaluated | a stand-in derivative at 49 call sites | body partials are taken against real dummies | with plain `sympy.diff` at every solver site, the Drucker–Prager floor law fails to compile on the tree route and solves on the graph route | +| SymPy merges powers across names and the merged power escapes the sqrt guard: a power law on a named invariant has a NaN Newton flux at rest | none; found by this change's tests | the guard sits in each node body, which SymPy does not merge into its parent | `test_0024`: tree NaN at rest, graph finite | | manifest and C built by two walks (#302) | two consistency guards and `_reveal_constants` | one walk | by construction; manifests identical on all fixtures | -| generated C that differs between ranks (#752, open) | rank 0's source adopted | canonical emission | byte-identical under hash seeds 0–2 (notch, box, VEP), after a preamble and when re-declared (`test_0024`). Across ranks untested: the #752 fixture (`rank_agreement_752.py`, np = 2) disagreed in 0 of 10 runs on either route, against about 2 in 10 before tier 1, so it no longer reproduces the defect | -| field symbols given their C names by mutating their classes, in an order that matters | `ccode_patch_fns`, the coordinate recovery block | an explicit map from leaf to C | not yet: steps 2–3 still spell leaves through the patched printer; the map belongs to step 4 | +| generated C that differs between ranks (#752, open) | rank 0's source adopted | canonical emission | byte-identical under hash seeds 0–2 (notch, box, VEP; `test_0105` for a Newton viscoplastic law), after a preamble and when re-declared (`test_0024`). The #752 fixture (`rank_agreement_752.py`) disagreed in 0 of 10 runs at np = 2, 3 and 4 on either route, against about 2 in 10 at np = 2 before tier 1, so it no longer reproduces the defect | +| field symbols given their C names by mutating their classes, in an order that matters | `ccode_patch_fns`, the coordinate recovery block | an explicit map from leaf to C | `_leaf_spellings`, built for each compile (step 4); a field the compile does not own is now an error, where a patched class could print another compile's slot | | generated C too large to read or to compile (#547) | opt-in CSE, lower optimisation flags | one line per named quantity | 22 KB against 4.0 MB for the notch's whole solve | The graph brings one failure class tier 1 does not have: it cannot cancel a quantity against its own reciprocal across a name. **Less global state, not less code.** The earlier estimate here (530 lines replaced by -360) does not survive the implementation. The lowering module is 442 lines, about a -third of them docstrings, and the branches it needs elsewhere about 60. Step 4 deletes -about 390 lines of the tree route: `_reveal_constants`, the scanning half of -`_extract_constants`, `_collect_constant_atoms`, `_xreplace_shared`, `_unique_symbols`, -the tree lowering and its two consistency guards in `generate_c_source`, the coordinate -recovery, the opt-in CSE path, the tree guard in `_jacobian_unwrap`, and the unused -`prepare_for_cache_key` and `_createext`. Replacing the class patching of -`ccode_patch_fns` (148 lines with its comments) by a map from leaf to C would remove -about 100 more. The line count comes out about even. What changes is where the +360) did not survive the implementation. Against `development`, the library source +gains 692 lines and loses 639 (step 4 counts): the lowering module is about 460 lines, a +third of them docstrings, and step 4 deleted the tree route — `_reveal_constants`, the +scanning half of `_extract_constants`, `_collect_constant_atoms`, `_xreplace_shared`, +`_unique_symbols`, the tree lowering and its two consistency guards in +`generate_c_source`, the coordinate recovery, the opt-in CSE path, the tree guard in +`_jacobian_unwrap`, the unused `prepare_for_cache_key` and `_createext` — and replaced +the class patching of `ccode_patch_fns` by a map from leaf to C. The line count comes out +about even. What changes is where the correctness lives: in one module whose every risk has a test that fails when the mechanism is broken (`test_0024`, each test checked against a mutation of the lowering), instead of in patches at the places each failure surfaced. The stand-in derivative at @@ -570,7 +578,7 @@ multiplied by a name that holds its reciprocal. - The cache protocol: memory, disk, rank-0 compile behind a collective decision. - The residual is never guarded, and the Picard tangent is frozen exactly as now. - The slot dictionaries `getext()` returns are keyed by the callback objects the solver - passed in, so the solvers' lookups (`ext_dict.jac[self._uu_G3]`) are unaffected. + passed in, so the solvers' lookups (`i_jac[self._uu_G3]`) are unaffected. - `describe()`, the transcript and the model fingerprints record the named, unexpanded forms (`str(template.sym)`) and the packed constants; none reads the generated C. - A model's `flux_jacobian` is still differentiated as given. @@ -586,12 +594,15 @@ multiplied by a name that holds its reciprocal. | a node class from an earlier compile reused | a changed `.sym` is ignored | the compile serial in each class's identity; a test that changes a body and rebuilds in one process without clearing SymPy's cache | | source that depends on the hash seed, a creation counter or a re-declaration | every parallel run repairs; a new process or a re-run notebook cell misses the cache | canonical emission; `test_0105` extended with a Newton fixture, a preamble of extra objects and a re-declared law | | a cancellation across a name lost | NaN at a state where the tree gave a finite limit, in an unguarded residual or Picard kernel | every residual and Picard kernel compared with the tree at a state of rest on every fixture | -| `test_0022` (tier 1) asserts the expanded, guarded tree of `_jacobian_unwrap`, and checks a block for `Derivative` with `has()`, which does not see node bodies | one test fails by construction, the other cannot fail | both rewritten against expanded nodes in the change that alters `_jacobian_unwrap` | +| `test_0022` (tier 1) asserts the expanded, guarded tree of `_jacobian_unwrap`, and checks a block for `Derivative` with `has()`, which does not see node bodies | one test fails by construction, the other cannot fail | the guard test now holds `_jit_graph.guard_half_integer_powers` to SymPy's `replace` on eight laws; the tests of deleted helpers went with them | | an entry that is zero today becomes a non-zero expression | assembly of an entry that evaluates to zero | numbers and leaves inlined; zero patterns compared on every fixture | -| guard placement changes cold-start behaviour | NaN at a state of rest | `test_1067`, extended to a power-law viscosity (the prototype finds both routes finite and equal) | -| `getext()` or another walker expands nodes back into the tree | the cost returns, silently | `getext()` lowers atoms itself; a size check on the emitted source of the notch fixture | -| an atom whose `.sym` is a matrix | it cannot be a scalar temporary | such atoms are expanded in place, as now | -| verbose-output assertions in `test_0004` | a test fails on wording, not on a defect | rewrite those assertions against the kernel contract | +| guard placement changes cold-start behaviour | NaN at a state of rest | `test_0024`: the guarded lowering against a guarded tree, at a random state and at rest, for a viscoplastic law and a power law on a named invariant; `test_1067` (cold Newton start) | +| `getext()` or another walker expands nodes back into the tree | the cost returns, silently | `getext()` lowers atoms itself; `test_0024` emits a twelve-layer law whose tree doubles with each layer and bounds its C | +| an atom whose `.sym` is a matrix or a vector | it cannot be a scalar temporary | such atoms are expanded in place (`_is_scalar`; `test_0024`) | +| a `mesh.X` coordinate rebuilt by the per-body cse | a cloned coordinate system, unequal to `mesh.N.x`: derivatives by the coordinate lose the explicit term | every leaf and child node hidden behind a placeholder while cse runs; `test_0024` | +| deciding constancy by a complete unwrap of each atom | setup exponential in the nesting depth while the C stays small | constancy decided bottom-up on the graph, with the rule of `_is_truly_constant`; `test_0024` counts no complete unwrap on a 16-layer law | +| a temporary hoists a quantity out of a Piecewise branch | it is computed even where the branch is not taken; values are unaffected, but under PETSc `-fp_trap` a `log` of a negative or a division at rest can trap | accepted: nothing in the repository traps floating-point exceptions | +| verbose-output assertions in `test_0004` | a test fails on wording, not on a defect | the verbose line prints the lowered kernel, so the assertions hold unchanged | | overhead on small kernels | constant-viscosity problems get slower to set up | measured on the small fixtures; budget: no slower than today | | a law singular at a state a model reaches (a zero denominator) | NaN where the tree gave rounding noise | the Newton solve of the acceptance criteria; such a law is the model's to fix | diff --git a/docs/developer/subsystems/expressions-functions.md b/docs/developer/subsystems/expressions-functions.md index 2e7415677..ea434835b 100644 --- a/docs/developer/subsystems/expressions-functions.md +++ b/docs/developer/subsystems/expressions-functions.md @@ -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). - -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. - -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) +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 generated C, + in which each constant is written as its slot. Changing a constant value + produces the same hash → cache hit → no recompilation. + +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 diff --git a/docs/developer/subsystems/jit-cache.md b/docs/developer/subsystems/jit-cache.md index 4fc68d5a6..ba0ae250f 100644 --- a/docs/developer/subsystems/jit-cache.md +++ b/docs/developer/subsystems/jit-cache.md @@ -38,7 +38,8 @@ processes. `getext()` 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 quadrature point, and the Newton tangent is +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, diff --git a/scripts/deprecated_pattern_allowlist.txt b/scripts/deprecated_pattern_allowlist.txt index fa55e67f7..42f2c6235 100644 --- a/scripts/deprecated_pattern_allowlist.txt +++ b/scripts/deprecated_pattern_allowlist.txt @@ -55,7 +55,6 @@ src/underworld3/units.py:except-pass src/underworld3/utilities/_api_tools.py:except-pass src/underworld3/utilities/_interrupt.py:except-pass src/underworld3/utilities/_jit_cache.py:except-pass -src/underworld3/utilities/_jitextension.py:except-pass src/underworld3/utilities/_params.py:except-pass src/underworld3/utilities/diagnostics.py:except-pass src/underworld3/utilities/mathematical_mixin.py:except-pass diff --git a/scripts/sessions/jit_graph/graph_size_probe.py b/scripts/sessions/jit_graph/graph_size_probe.py index 123f66eda..4193c745a 100644 --- a/scripts/sessions/jit_graph/graph_size_probe.py +++ b/scripts/sessions/jit_graph/graph_size_probe.py @@ -1,5 +1,5 @@ -"""Size of a kernel's shared graph of named sub-expressions against the tree the -current JIT expands it into (#823). Builds the campaign notch model, does NOT solve, +"""Size of a kernel's shared graph of named sub-expressions against the tree the tier 1 +JIT expanded it into (#823). Builds the campaign notch model, does NOT solve, and never expands: the expanded size is counted by dynamic programming over the graph. """ import os, sys, time diff --git a/scripts/sessions/jit_graph/library_setup_profile.py b/scripts/sessions/jit_graph/library_setup_profile.py index 96557af34..b9213cc4e 100644 --- a/scripts/sessions/jit_graph/library_setup_profile.py +++ b/scripts/sessions/jit_graph/library_setup_profile.py @@ -81,7 +81,9 @@ def timed_compile(*a, **k): "subs": ("sympy/core/basic.py", "subs"), "as_real_imag (Pow)": ("sympy/core/power.py", "as_real_imag"), "_jacobian_unwrap": ("underworld3/cython/petsc_generic_snes_solvers", "_jacobian_unwrap"), - "_extract_constants": ("underworld3/utilities/_jitextension.py", "_extract_constants"), + "_extract_constants (tier 1 build)": ("underworld3/utilities/_jitextension.py", "_extract_constants"), + "lower_callbacks (graph build)": ("underworld3/utilities/_jit_graph.py", "lower_callbacks"), + "emit (graph build)": ("underworld3/utilities/_jit_graph.py", "emit"), "compile_and_load": ("underworld3/utilities/_jitextension.py", "compile_and_load"), } for label, (path, func) in groups.items(): diff --git a/scripts/sessions/jit_graph/rank_agreement_752.py b/scripts/sessions/jit_graph/rank_agreement_752.py index cb849747b..18e63975b 100644 --- a/scripts/sessions/jit_graph/rank_agreement_752.py +++ b/scripts/sessions/jit_graph/rank_agreement_752.py @@ -5,7 +5,8 @@ differed across ranks in about one np=2 run in five on development. Builds the kernels and reports, on rank 0, whether ``_agree_source_across_ranks`` saw different hashes. The route is chosen by the build: run it in a build of this branch -(graph) and in a build of development (tree). Run under ``mpirun -n 2``. +(graph) and in a build of development (tree). Run under ``mpirun -n N``; it was +measured at N = 2, 3 and 4. """ import sympy import underworld3 as uw diff --git a/scripts/sessions/jit_graph/route_ab_compare.py b/scripts/sessions/jit_graph/route_ab_compare.py index 95b79130d..93f30e4cb 100644 --- a/scripts/sessions/jit_graph/route_ab_compare.py +++ b/scripts/sessions/jit_graph/route_ab_compare.py @@ -1,10 +1,11 @@ """Compare the two routes' records written by ``route_ab.py`` for one fixture.""" import os -import sys import numpy as np +import underworld3 as uw -out = os.path.expanduser(sys.argv[1] if len(sys.argv) > 1 else "~/+Simulations/jit_graph/tier2") +params = uw.Params(uw_dir=uw.Param("~/+Simulations/jit_graph/tier2", description="directory of the .npz records")) +out = os.path.expanduser(str(params.uw_dir)) cases = sorted({f.rsplit("_", 1)[0] for f in os.listdir(out) if f.endswith(".npz")}) for case in cases: paths = [os.path.join(out, f"{case}_{r}.npz") for r in ("tree", "graph")] diff --git a/scripts/sessions/jit_graph/route_assemble_compare.py b/scripts/sessions/jit_graph/route_assemble_compare.py index 606318791..63bbac98c 100644 --- a/scripts/sessions/jit_graph/route_assemble_compare.py +++ b/scripts/sessions/jit_graph/route_assemble_compare.py @@ -1,11 +1,11 @@ """Compare the residuals and Jacobians ``route_assemble.py`` saved for the two routes.""" import os -import sys import numpy as np +import underworld3 as uw -out = os.path.expanduser(sys.argv[1] if len(sys.argv) > 1 else - "~/+Simulations/jit_graph/tier2/assemble") +params = uw.Params(uw_dir=uw.Param("~/+Simulations/jit_graph/tier2/assemble", description="directory of the .npz records")) +out = os.path.expanduser(str(params.uw_dir)) cases = sorted({f.rsplit("_", 1)[0] for f in os.listdir(out) if f.endswith(".npz")}) for case in cases: paths = [os.path.join(out, f"{case}_{r}.npz") for r in ("tree", "graph")] diff --git a/src/underworld3/constitutive_models.py b/src/underworld3/constitutive_models.py index 4beb479d8..94b8b3595 100644 --- a/src/underworld3/constitutive_models.py +++ b/src/underworld3/constitutive_models.py @@ -2644,7 +2644,7 @@ def viscosity(self): # consistent tangent of the *harmonic* problem) and converges WORSE than # Picard on hard-yield VEP; the robust route is problem-space homotopy # (ramp the softmin softness δ→0), not a smooth tangent. See the design doc - # docs/developer/design/jacobian-unwrap-constants-bug.md. The generic + # docs/developer/design/jacobian-consistent-tangent.md. The generic # Constitutive_Model.flux_jacobian hook (default None) remains available. @property @@ -4838,9 +4838,9 @@ def __init__( Default: False (assumes IndexSwarmVariable maintains partition of unity) """ # Constituents that share a parameter name (every ViscousFlowModel - # calls its viscosity \eta) rely on _JITConstant keeping its - # constants[] slots distinct; see utilities/_jitextension.py and the - # regression in tests/test_0103_jit_rampable_constants.py. + # calls its viscosity \eta) rely on the constants manifest keeping their + # constants[] slots distinct (by object, ordered by creation); see + # utilities/_jitextension.py and tests/test_0103_jit_rampable_constants.py. # Validate compatibility before initialization self._validate_model_compatibility(constitutive_models) if any(getattr(m, "_stress_history", "stress") != "stress" for m in constitutive_models): diff --git a/src/underworld3/cython/petsc_generic_snes_solvers.pyx b/src/underworld3/cython/petsc_generic_snes_solvers.pyx index df0868925..fbf46b422 100644 --- a/src/underworld3/cython/petsc_generic_snes_solvers.pyx +++ b/src/underworld3/cython/petsc_generic_snes_solvers.pyx @@ -403,12 +403,13 @@ class SolverBaseClass(uw_object): ``False`` (default) Differentiate the residual flux *as wrapped* — the effective viscosity is frozen, giving a Picard / defect-correction tangent. - Bit-identical to the long-standing behaviour. Globally robust; - load-bearing for the tuned hard-yield viscoplastic paths. + Globally robust; load-bearing for the tuned hard-yield viscoplastic + paths. ``True`` - Unwrap the flux before differentiation so the tangent captures - :math:`\partial\eta/\partial(\nabla v)` (full Newton). Fast near - the solution; its yield kink can stall the line search far from it. + Differentiate through the named quantities of the flux, so the + tangent captures :math:`\partial\eta/\partial(\nabla v)` (full + Newton). Fast near the solution; its yield kink can stall the line + search far from it. ``"continuation"`` Picard :math:`\rightarrow` Newton. Blend :math:`J(\alpha) = J_{\mathrm{picard}} + \alpha\,(J_{\mathrm{newton}} @@ -421,8 +422,9 @@ class SolverBaseClass(uw_object): The Newton flux for a model whose flux has a non-smooth yield kink is the model's own smooth law (``constitutive_model.flux_jacobian``) when - it provides one; otherwise the exact unwrapped flux. See - ``docs/developer/design/jacobian-unwrap-constants-bug.md``. + it provides one; otherwise the exact flux with its named quantities as + graph nodes (``_jacobian_unwrap``). See + ``docs/developer/design/jacobian-consistent-tangent.md``. Raises ------ @@ -4160,8 +4162,8 @@ class SNES_Scalar(SolverBaseClass): sympy.core.cache.clear_cache() - # RESIDUAL: don't unwrap here — let getext()'s two-phase unwrap handle - # it (preserves constant UWexpressions as symbols for constants[]). + # RESIDUAL: don't unwrap here — getext() lowers it onto the graph of + # named quantities (constant UWexpressions stay constants[] slots). f0 = sympy.Array(self.F0.sym).reshape(1).as_immutable() # F1 is the flux vector, which lives in the embedded coordinate # space (cdim components). For volume meshes dim==cdim so this @@ -5126,8 +5128,8 @@ class SNES_Vector(SolverBaseClass): ## The jacobians are determined from the above (assuming we ## do not concern ourselves with the zeros) # Residual piece shapes: f0 is (cdim,) per-component, F1 is (cdim, cdim). - # RESIDUAL: don't unwrap here — let getext()'s two-phase unwrap handle - # it (preserves constant UWexpressions as symbols for constants[]). The + # RESIDUAL: don't unwrap here — getext() lowers it onto the graph of + # named quantities (constant UWexpressions stay constants[] slots). The # Jacobian sources (f0_jac_list / F1_user_jac) are derived below. F0_user = sympy.Matrix(self.F0.sym) F1_user = sympy.Matrix(self.F1.sym) @@ -8189,8 +8191,8 @@ class SNES_Stokes_SaddlePt(SolverBaseClass): sympy.core.cache.clear_cache() - # RESIDUAL: don't unwrap here — let getext()'s two-phase unwrap handle - # it (preserves constant UWexpressions as symbols for constants[]). The + # RESIDUAL: don't unwrap here — getext() lowers it onto the graph of + # named quantities (constant UWexpressions stay constants[] slots). The # JACOBIAN sources are unwrapped separately below (see _jac_source) so # the derivative sees through the viscosity — that is the Newton fix. F0 = sympy.Array(self.F0.sym) @@ -8219,19 +8221,18 @@ class SNES_Stokes_SaddlePt(SolverBaseClass): U = sympy.Array(self.u.sym).reshape(dim) P = sympy.Array(self.p.sym).reshape(1) - # Expand UWexpressions down to (but NOT including) constant atoms, - # element-wise, for the Jacobian derivative ONLY. This exposes the - # field / grad-v dependence of the (effective) viscosity so that - # derive_by_array forms the full Newton tangent (e.g. Min -> Heaviside - # yield switch), instead of freezing eta_eff as an opaque atom and - # silently running a Picard / defect-correction tangent. Truly-constant - # atoms (eta0, tau_y, ...) survive as symbols so the constants[] - # runtime-update mechanism is preserved (the keep-constants predicate - # is shared with getext()'s _extract_constants, so they cannot drift). + # For the Jacobian derivative ONLY, each non-constant UWexpression + # becomes a graph node (_jacobian_unwrap, #823), whose partial + # derivatives the chain rule composes. This exposes the field / grad-v + # dependence of the (effective) viscosity so that the derivative forms + # the full Newton tangent (e.g. Min -> Heaviside yield switch), instead + # of freezing eta_eff as an opaque atom and silently running a Picard / + # defect-correction tangent. Truly-constant atoms (eta0, tau_y, ...) stay + # as symbols, so the constants[] runtime-update mechanism is preserved. # The residual fns above (self._u_F0/_u_F1/_p_F0) are left untouched — - # getext() unwraps those itself. For constant-viscosity problems this - # is a no-op (eta has no grad-v dependence) so the Jacobian is - # bit-identical. See docs/developer/design/jacobian-unwrap-constants-bug.md + # getext() lowers those itself. For constant-viscosity problems this is + # a no-op (eta has no grad-v dependence). See + # docs/developer/design/jacobian-consistent-tangent.md # # (see consistent_jacobian / _jacobian_source: default Picard, bit- # identical; True -> Newton; "continuation" -> alpha-blended.) diff --git a/src/underworld3/utilities/_jit_graph.py b/src/underworld3/utilities/_jit_graph.py index 63db6af39..cbc2b7bd5 100644 --- a/src/underworld3/utilities/_jit_graph.py +++ b/src/underworld3/utilities/_jit_graph.py @@ -21,19 +21,16 @@ import sympy from sympy.core.function import AppliedUndef, UndefinedFunction from sympy.tensor.array import NDimArray +from sympy.vector.basisdependent import BasisDependent from sympy.vector.scalar import BaseScalar +# TODO(Charter S6): a literal regulariser, the one the expanded-tree guard used; a +# change of value changes every Newton kernel's C (see uw.maths.functions.vanishing). _EPS2 = sympy.Float(1.0e-36) _serial = itertools.count(1) _nodes_made = False -def nodes_exist(): - """Whether any node has been made in this process; when not, no expression can - hold one and the unwrappers skip the search.""" - return _nodes_made - - def guard_half_integer_powers(e): r"""``e`` with :math:`10^{-36}` added to the base of every half-integer power whose base has free symbols: the sqrt guard of ``_jacobian_unwrap``, the same rule @@ -66,11 +63,22 @@ def guard(n): class _KernelNode(AppliedUndef): - """A named quantity of a kernel, applied to the leaves its value depends on.""" + """A named quantity of a kernel, applied to the leaves its value depends on. + + Printed (``str``, ``latex``) as the name of the quantity it was made from, so a + lowered block reads as the law; ``srepr``, which the class name and the emission + keys are built from, is unchanged.""" def fdiff(self, argindex=1): return self._graph.slot_derivative(self, argindex - 1) + def _sympystr(self, printer): + return self._label + + def _latex(self, printer, exp=None): + text = self._label + return f"{text}^{{{exp}}}" if exp is not None else text + def _ccode(self, printer): raise RuntimeError( f"JIT: node {self.func.__name__} reached the C printer; every node must be " @@ -102,6 +110,13 @@ def _ccode(self, printer): return self._ccodestr +def _is_scalar(body): + """Whether ``body`` can be one C temporary: a scalar value, not a condition, a + matrix, an array or a vector.""" + return (isinstance(body, sympy.Expr) and not getattr(body, "is_Matrix", False) + and not isinstance(body, (NDimArray, BasisDependent))) + + def body_of(app): """A node application's body, with the application's arguments substituted.""" cls = app.func @@ -150,6 +165,7 @@ class KernelGraph: def __init__(self): self.serial = next(_serial) self._const = {} # id(atom) -> (atom, bool) + self._const_busy = set() self._node = {} # (id(atom), guarded) -> (atom, replacement) self._by_body = {} # body -> node class self._deriv = {} # (node class, slot) -> derivative in the class's deps @@ -158,14 +174,55 @@ def __init__(self): # ------------------------------------------------------------------ leaves def is_constant(self, atom): + """Whether a UWexpression is a ``constants[]`` leaf: the rule of + ``_is_truly_constant`` (its complete non-dimensional value reads no + coordinate, so no field either), decided bottom-up over the atoms it reads, + once per atom. ``_is_truly_constant`` unwraps each atom completely, which costs + the size of the expanded tree under it, for every atom.""" hit = self._const.get(id(atom)) - if hit is None: - from underworld3.function.expressions import UWexpression - from underworld3.utilities._jitextension import _is_truly_constant - hit = self._const[id(atom)] = (atom, _is_truly_constant(atom, UWexpression)) - return hit[1] + if hit is not None: + return hit[1] + if id(atom) in self._const_busy: + return False # cyclic: lowering the atom refuses it + self._const_busy.add(id(atom)) + try: + result = self._decide_constant(atom) + finally: + self._const_busy.discard(id(atom)) + self._const[id(atom)] = (atom, result) + return result + + def _decide_constant(self, atom): + import underworld3 + from underworld3.function.expressions import UWexpression + from underworld3.function.quantities import UWQuantity + + # as the non-dimensional unwrap resolves an atom (_unwrap_atom) + if atom.has_units and underworld3._is_scaling_active(): + try: + float(atom.data) + return True + except Exception: # not a number: decided by its content below + pass + inner = atom.sym + if isinstance(inner, UWQuantity) and not isinstance(inner, UWexpression): + return True + if not hasattr(inner, "free_symbols"): + try: + float(inner) + return True + except (TypeError, ValueError): + return False + for s in inner.free_symbols: + if isinstance(s, BaseScalar): + return False + if isinstance(s, UWexpression) and not self.is_constant(s): + return False + return True def is_leaf(self, s): + """Whether ``s`` is read by a kernel as it is: a field value or gradient, a + coordinate, a constant atom, or another symbol.""" from underworld3.function.expressions import UWexpression if isinstance(s, _KernelNode): @@ -197,6 +254,7 @@ def order_key(s): return (KernelGraph.display_key(s), getattr(s, "instance_number", 0)) def leaves(self, e): + """The leaves ``e`` reads, through the arguments of the nodes it holds.""" out, seen, stack = set(), set(), [e] while stack: a = stack.pop() @@ -232,8 +290,11 @@ def lower(self, e, guarded=False): if not isinstance(e, sympy.Basic): return sympy.sympify(e) uw_types = (UWexpression, UWQuantity, UWCoordinate) - atoms = sorted((s for s in e.free_symbols if isinstance(s, uw_types)), - key=self.order_key) + # a UWCoordinate equals the base scalar it wraps, so free_symbols can hold the + # base scalar in its place: they are found by type + found = {s for s in e.free_symbols if isinstance(s, uw_types)} + found |= e.atoms(UWCoordinate) + atoms = sorted(found, key=self.order_key) rule = {} for s in atoms: r = self._replacement(s, guarded) @@ -254,6 +315,8 @@ def _replacement(self, atom, guarded): return self.node_of(atom, guarded) def node_of(self, atom, guarded): + """The node of ``atom`` (or its expansion in place, for a matrix-valued + atom), its body lowered with the same ``guarded`` variant.""" key = (id(atom), guarded) hit = self._node.get(key) if hit is not None: @@ -267,7 +330,7 @@ def node_of(self, atom, guarded): body = guard_half_integer_powers(body) # a matrix-valued atom cannot be one scalar temporary: it is expanded # in place, as the tree route expands it - out = self.make_node(body) if isinstance(body, sympy.Expr) else body + out = self.make_node(body, label=atom.name) if _is_scalar(body) else body finally: self._busy.discard(key) self._node[key] = (atom, out) @@ -280,14 +343,15 @@ def _display_name(self, body): text = sympy.srepr(body.xreplace(canon)) return "N" + hashlib.sha1(text.encode()).hexdigest()[:12] - def make_node(self, body): - """The node for ``body``, one per distinct body. A body that is a number, a - single leaf, a single node, or reads no leaf is returned as itself, and so is a - condition (a node is a value; a Piecewise refuses one as its condition).""" + def make_node(self, body, label=None): + """The node for ``body``, one per distinct body, printed as ``label`` (the + first label it is made with). A body that is a number, a single leaf, a single + node, or reads no leaf is returned as itself, and so is a condition (a node is + a value; a Piecewise refuses one as its condition).""" global _nodes_made body = sympy.sympify(body) - if not isinstance(body, sympy.Expr): + if not _is_scalar(body): return body if not self._splitting: body = self._split_shared(body) @@ -302,6 +366,7 @@ def make_node(self, body): real=True, _ctx=self.serial, _n=len(self._by_body), __dict__={"_graph": self}) cls._body, cls._deps = body, deps + cls._label = label if label is not None else "_shared" self._by_body[body] = cls _nodes_made = True return cls(*cls._deps) @@ -311,14 +376,22 @@ def _split_shared(self, body): Canonical order, so the split does not depend on the hash seed.""" if body.is_Atom: return body - repl, (reduced,) = sympy.cse([body], symbols=sympy.numbered_symbols("_cse", real=True)) + # every leaf and child node becomes a placeholder while cse runs: cse rebuilds + # what it holds, and a rebuilt UWCoordinate gets a cloned coordinate system and + # is no longer equal to the base scalar it stood for + opaque = body.atoms(AppliedUndef) | { + s for s in body.free_symbols if self.is_leaf(s)} + hide = {s: sympy.Dummy(real=True) for s in sorted(opaque, key=self.order_key)} + show = {d: s for s, d in hide.items()} + repl, (reduced,) = sympy.cse([body.xreplace(hide)], + symbols=sympy.numbered_symbols("_cse", real=True)) if not repl: return body self._splitting = True try: - rule = {} + rule = dict(show) for sym, e in repl: - rule[sym] = self.make_node(e.xreplace(rule)) + rule[sym] = self.make_node(e.xreplace(rule), label="_shared") return reduced.xreplace(rule) finally: self._splitting = False @@ -336,7 +409,8 @@ def slot_derivative(self, app, i): dummies = [sympy.Dummy(real=True) for _ in deps] b = cls._body.xreplace(dict(zip(deps, dummies))) d = sympy.diff(b, dummies[i]).xreplace(dict(zip(dummies, deps))) - self._deriv[key] = self.make_node(d) + self._deriv[key] = self.make_node( + d, label=f"\\partial_{{{i}}}{{{cls._label}}}") d = self._deriv[key] if app.args == cls._deps: return d diff --git a/src/underworld3/utilities/_jitextension.py b/src/underworld3/utilities/_jitextension.py index 0a0ece09c..db2bf940c 100644 --- a/src/underworld3/utilities/_jitextension.py +++ b/src/underworld3/utilities/_jitextension.py @@ -1,5 +1,6 @@ from typing import Optional import os +import re import shutil import subprocess from xmlrpc.client import boolean @@ -290,16 +291,20 @@ def counts(self): # ============================================================================ class _JITConstant(sympy.Symbol): - r"""Symbol subclass that renders as ``constants[i]`` in generated C code. + r"""A ``constants[]`` slot of the manifest: the placeholder each manifested + constant maps to, carrying its C, ``constants[i]``. - Used by the JIT compiler to route constant UWexpressions through PETSc's - ``PetscDSSetConstants()`` mechanism instead of baking values as C literals. + Constant UWexpressions are routed through PETSc's ``PetscDSSetConstants()`` + instead of being baked as C literals. The graph lowering (#823) writes each + constant leaf as ``constants[i]`` from this placeholder; the placeholder itself + no longer appears in the expressions that are printed, so its identity and + ordering below now matter to code that substitutes it into an expression (as + ``test_0103`` does), not to the generated C, whose order is canonical. Two constants may legitimately share a display name — every ``ViscousFlowModel`` calls its viscosity :math:`\eta`, so a two-material model has two of them — and each needs its own ``constants[]`` slot. Two - separate SymPy properties have to hold for that to work, and they are not - the same property: + separate SymPy properties hold for that, and they are not the same property: **Identity** — the slot index is in ``_hashable_content``, and the symbol is built with ``Symbol.__xnew__`` to bypass SymPy's ``(cls, name)`` @@ -311,13 +316,10 @@ class _JITConstant(sympy.Symbol): **Ordering** — the slot index is also in the NAME. ``_hashable_content`` does nothing for ``Symbol.sort_key()``, which is derived from the name, so two same-named placeholders sort equal; term order inside an ``Add`` then - falls back to hash order, which is randomised per process. The generated C - then differs between MPI ranks and ``getext``'s cross-rank hash check - aborts the run — intermittently, since it depends on the hash seed. + falls back to hash order, which is randomised per process. (When the + placeholders were printed, that made the generated C differ between MPI ranks.) - Identity without ordering is a parallel abort; ordering without identity is - a silently wrong answer. Keep both. ``tests/test_0103_jit_rampable_constants.py`` - pins each one separately. + ``tests/test_0103_jit_rampable_constants.py`` pins each one separately. A slot holds a C double, so it is built ``real`` (real and finite) for SymPy's simplification (#823). Declared at construction, not by a class handler: SymPy @@ -365,6 +367,11 @@ def _extract_constants(all_fns, mesh): return _manifest_from(_jit_graph.constant_leaves(lowered)) +def _manifest_of(lowered): + """``(manifest, subs_map)`` of callbacks already lowered (``_jit_graph``).""" + return _manifest_from(_jit_graph.constant_leaves(lowered)) + + def _manifest_from(constant_exprs): """The ``constants[]`` manifest of a set of constant atoms: ``(manifest, subs_map)``, slots ordered by name, then creation order.""" @@ -497,11 +504,6 @@ def _is_truly_constant(expr, UWexpression): return False if isinstance(sym, sympy.Function): return False - # UnderworldFunction symbols have _ccodestr pointing to petsc arrays - if hasattr(sym, '_ccodestr') and not isinstance(sym, _JITConstant): - ccode = sym._ccodestr - if 'petsc_u' in ccode or 'petsc_a' in ccode or 'petsc_x' in ccode or 'petsc_n' in ccode: - return False # Other UWexpressions that didn't fully unwrap — not constant if isinstance(sym, UWexpression): return False @@ -659,8 +661,7 @@ def getext( # Each callback is lowered onto the shared graph of named quantities once # (``_jit_graph``, #823); the manifest is the constant leaves of what was lowered. lowered_fns = _jit_graph.lower_callbacks(callbacks.flat(), mesh) - constants_manifest, constants_subs_map = _manifest_from( - _jit_graph.constant_leaves(lowered_fns)) + constants_manifest, constants_subs_map = _manifest_of(lowered_fns) if debug and underworld3.mpi.rank == 0: if constants_manifest: @@ -677,11 +678,11 @@ def getext( mesh, callbacks, primary_field_list, - constants_subs_map=constants_subs_map, + lowered_fns, + constants_subs_map, verbose=verbose, debug=debug, debug_name=debug_name, - lowered_fns=lowered_fns, ) gen_randstr = diag["randstr"] @@ -710,10 +711,12 @@ def getext( # rank-0-compiles/others-load protocol would break. # # Agreement used to be REQUIRED here, and a mismatch was a hard error. It - # fires in practice: the lowering above is not yet deterministic across - # ranks (#752), and a Stokes solve with a power-law transversely isotropic - # viscosity trips it in roughly half of np=2 runs. What we measured there - # matters for why this is safe to repair rather than refuse: + # fired in practice before the graph lowering (#752): a Stokes solve with a + # power-law transversely isotropic viscosity tripped it in about one np=2 run + # in five. The graph's emission is canonical by construction (temporaries + # ordered by a hash of their C, leaves written as the C they read), and the + # #752 fixture now agrees at np = 2, 3 and 4; the check stays as a guard. What + # was measured on the disagreeing runs is why repairing is safe: # # * the sources differ only in the ORDER of factors in commutative # products — identical token multisets, identical length, identical @@ -728,9 +731,8 @@ def getext( # rank rehashes from it, which restores the one invariant that matters: one # source, one hash, one module. # - # This is a REPAIR, not a fix. The non-determinism upstream is still a bug - # and still worth finding, which is why it is said out loud rather than - # papered over silently. + # A disagreement is a defect in the lowering, which is why it is reported + # rather than papered over silently. canonical_codeguys, canonical_source, source_hash = _agree_source_across_ranks( canonical_codeguys, canonical_source, source_hash ) @@ -875,6 +877,9 @@ def getext( ) +_NUMBER_DECLARATION = re.compile(r"const double [A-Za-z_]\w* = [^;]*;") + + class _Unspellable(Exception): """A leaf the kernel has no C for.""" @@ -962,10 +967,7 @@ def _spell_leaf(leaf, spellings, constants_subs_map): # the mesh names its coordinates; a fresh instance, or a UWCoordinate that # SymPy's cache returned for its equal base scalar, is named from its index # and system - try: - text = leaf._ccodestr - except AttributeError: - text = None + text = getattr(leaf, "_ccodestr", None) if isinstance(text, str): return text idx, system = leaf._id[0], str(leaf._id[1]) @@ -984,6 +986,7 @@ def _spell_leaf(leaf, spellings, constants_subs_map): def _unconvertible_message(symbols, index, fn_original): + """The message for leaves of kernel ``index`` that have no C.""" details = [] for sym in symbols: detail = f" - {sym} (type: {type(sym).__name__})" @@ -1014,16 +1017,29 @@ def _unconvertible_message(symbols, index, fn_original): def _print_kernel(printer, temporaries, outputs, out): """The C body of one kernel: a ``const double`` per temporary, in order, then - the outputs. A SymPy function the printer cannot write returns its - ``// Not supported in C:`` text, which the caller refuses.""" - lines = [] + the outputs. The printer declares a number symbol it reads (``EulerGamma``, + ``Catalan``) before the code that reads it; each declaration is written once, at + the top. A SymPy function the printer cannot write raises + ``PrintMethodNotImplementedError`` (SymPy 1.14), or, in a SymPy that returns + ``// Not supported in C:`` text instead, the text comes back for the caller to + refuse.""" + declarations, lines = [], [] + + def without_declarations(code): + rest = code.split("\n") + while rest and _NUMBER_DECLARATION.fullmatch(rest[0]): + if rest[0] not in declarations: + declarations.append(rest[0]) + rest = rest[1:] + return "\n".join(rest) + for t, body in temporaries: code = printer.doprint(body) if code.startswith("// Not supported in C:"): return code - lines.append(f"const double {t._ccodestr} = {code};") - lines.append(printer.doprint(outputs, out)) - return "\n".join(lines) + lines.append(f"const double {t._ccodestr} = {without_declarations(code)};") + lines.append(without_declarations(printer.doprint(outputs, out))) + return "\n".join(declarations + lines) @timing.routine_timer_decorator @@ -1085,11 +1101,11 @@ def generate_c_source( mesh: underworld3.discretisation.Mesh, callbacks: JITCallbackSet, primary_field_list, - constants_subs_map: Optional[dict] = None, + lowered_fns, + constants_subs_map, verbose: Optional[bool] = False, debug: Optional[bool] = False, debug_name=None, - lowered_fns=None, ): """Generate the setup.py / C header / Cython wrapper for a JIT bundle. @@ -1108,14 +1124,13 @@ def generate_c_source( callbacks : JITCallbackSet primary_field_list : list Variables that map to PETSc primary variable arrays (``petsc_u[]``). - constants_subs_map : dict, optional - Mapping from UWexpression → ``_JITConstant`` placeholder; built from the - lowered callbacks when not given. - lowered_fns : list of sympy.Matrix, optional + lowered_fns : list of sympy.Matrix The callbacks lowered onto the shared graph (``_jit_graph.lower_callbacks``), - one per entry of ``callbacks.flat()``; lowered here when not given. Each - kernel is emitted as one C temporary per distinct computation, then its - outputs. + one per entry of ``callbacks.flat()``. Each kernel is emitted as one C + temporary per distinct computation, then its outputs. + constants_subs_map : dict + Mapping from UWexpression to its ``_JITConstant`` placeholder + (``_manifest_of(lowered_fns)``). Returns ------- @@ -1128,18 +1143,10 @@ def generate_c_source( Equation-range counts and the random symbol prefix, used by the caller for verbose printing and for building the fn-layout manifest. """ - from sympy import symbols, Eq, MatrixSymbol - from underworld3 import VarType - fns = callbacks.flat() count_residual_sig, count_bc_sig, count_jacobian_sig, \ count_bd_residual_sig, count_bd_jacobian_sig = callbacks.counts - if lowered_fns is None: - lowered_fns = _jit_graph.lower_callbacks(fns, mesh) - if constants_subs_map is None: - _, constants_subs_map = _manifest_from(_jit_graph.constant_leaves(lowered_fns)) - # The C a kernel reads for each leaf of the graph: an explicit map, built for # this compile, instead of C names patched onto the field classes (#823). spellings = _leaf_spellings(mesh, primary_field_list) diff --git a/tests/test_0022_unwrap_memoised_matches_fixed_point.py b/tests/test_0022_unwrap_memoised_matches_fixed_point.py index e351d3f09..be03342f5 100644 --- a/tests/test_0022_unwrap_memoised_matches_fixed_point.py +++ b/tests/test_0022_unwrap_memoised_matches_fixed_point.py @@ -1,8 +1,5 @@ -"""The memoised unwrap (#823) gives exactly what the algorithm it replaces gave. - -(The memoised sqrt guard, the identity-walk symbol scan and the shared xreplace were -tested here too; they went with the expanded-tree JIT route, and the graph route's -guard is held to the guarded tree by test_0024.) +"""The memoised unwrap and the memoised sqrt guard (#823) give exactly what the +algorithms they replace gave. The old ``unwrap_expression`` iterated ``subs`` passes to a fixed point, and in ``symbolic_keep_constants`` mode ran ``_is_truly_constant`` (itself a full unwrap) for every @@ -104,6 +101,26 @@ def test_memoised_unwrap_is_the_fixed_point(laws, mode): assert sympy.srepr(new) == sympy.srepr(old), (name, mode) +def test_the_jacobian_sqrt_guard_matches_replace(laws): + """The guard every Newton node body gets (``_jit_graph.guard_half_integer_powers``), + a memoised rebuild, against SymPy's own ``replace`` with the same rule.""" + from underworld3.utilities._jit_graph import guard_half_integer_powers + + eps2 = sympy.Float(1.0e-36) + + def old_guard(e): + return e.replace( + lambda n: (n.is_Pow and n.exp.is_Rational and n.exp.q == 2 + and n.args[0].free_symbols), + lambda n: sympy.Pow(n.args[0] + eps2, n.exp)) + + for name, expr in laws: + tree = _unwrap_each(expr, lambda e: _fixed_point_unwrap(e, "symbolic_keep_constants")) + new = _unwrap_each(tree, guard_half_integer_powers) + old = _unwrap_each(tree, old_guard) + assert sympy.srepr(new) == sympy.srepr(old), name + + def test_each_atom_is_tested_for_constancy_once(monkeypatch): """The cost that #823 removed: the fixed-point passes asked whether each atom was a constant on every pass, and each answer was itself a full unwrap. On a chain of diff --git a/tests/test_0024_jit_graph_lowering.py b/tests/test_0024_jit_graph_lowering.py index c93a9f278..84e819315 100644 --- a/tests/test_0024_jit_graph_lowering.py +++ b/tests/test_0024_jit_graph_lowering.py @@ -81,6 +81,8 @@ def test_a_derivative_through_nodes_is_the_derivative_of_the_tree(box): def test_second_derivatives_and_a_constant_sensitivity_pass_through_nodes(box): + """A node's derivative is a node, so it differentiates again; and a derivative + with respect to a constant atom (a sensitivity) reaches every body that reads it.""" mesh, T, v = box u = T.sym[0] c = uw.expression(r"c_{0024b}", 1.3, "constant") @@ -160,7 +162,12 @@ def _guarded_tree(e): def test_the_guarded_lowering_is_the_guarded_tree(box): """The Newton source on the graph (the sqrt guard in each node body) evaluates to the guarded tree, and so does its tangent, including at a state of rest where the - unguarded tangent is 0/0.""" + unguarded tangent is 0/0. + + Except where the tree escapes its own guard: in a power law on a named invariant, + eta = edot**(1/n - 1) with edot = sqrt(g), the tree merges the powers into + g**((1/n - 1)/2), not a half-integer power, so its Newton flux is NaN at rest. The + graph keeps edot a node with a guarded body and stays finite there.""" from underworld3.cython.generic_solvers import _jacobian_unwrap mesh, T, v = box @@ -172,14 +179,29 @@ def test_the_guarded_lowering_is_the_guarded_tree(box): cm.Parameters.shear_viscosity_0 = 1.0 cm.Parameters.yield_stress = uw.expression(r"C_{0024e}", 0.5) + 0.3 * p.sym[0] cm.Parameters.yield_stress_min = uw.expression(r"\tau_{0024e}", 0.01) - flux = sympy.Matrix(stokes.F1.sym) + viscoplastic = sympy.Matrix(stokes.F1.sym) - tree = flux.applyfunc(_guarded_tree) - graph = _jacobian_unwrap(flux) - assert graph.atoms(jg._KernelNode), "the Newton source was not lowered" + # a power law on a named strain-rate invariant, singular at rest unless guarded + stokes.constitutive_model = uw.constitutive_models.ViscousFlowModel + edot = uw.expression(r"\dot\varepsilon_{0024e}", stokes.Unknowns.Einv2, "invariant") + n = uw.expression(r"n_{0024e}", 3, "stress exponent") + stokes.constitutive_model.Parameters.shear_viscosity_0 = edot ** (1 / n - 1) + power_law = sympy.Matrix(stokes.F1.sym) L = stokes.Unknowns.L rest = {L[i, j] for i in range(2) for j in range(2)} + for flux, tree_escapes_at_rest in ((viscoplastic, False), (power_law, True)): + tree = flux.applyfunc(_guarded_tree) + graph = _jacobian_unwrap(flux) + assert graph.atoms(jg._KernelNode), "the Newton source was not lowered" + escaped = _same_guarded_values(tree, graph, L, rest) + assert escaped == tree_escapes_at_rest, escaped + + +def _same_guarded_values(tree, graph, L, rest): + """Graph against tree at a random state and at rest; returns whether the tree was + non-finite anywhere (only ever at rest), where the graph must be finite.""" + escaped = False for state in (dict(seed=1), dict(seed=2, rest=rest)): for i in range(2): for j in range(2): @@ -188,8 +210,13 @@ def test_the_guarded_lowering_is_the_guarded_tree(box): n = _numbers([tree[i, j], d_tree], **state) for t, g in ((tree[i, j], graph[i, j]), (d_tree, d_graph)): vt, vg = _value(t, n), _value(g, n) - assert np.isfinite(vt) and np.isfinite(vg), (state, t) + assert np.isfinite(vg), (state, g) + if not np.isfinite(vt): + assert "rest" in state, (state, t) + escaped = True + continue assert abs(vg - vt) <= 1.0e-12 * max(abs(vt), 1.0), (state, i, j) + return escaped def _header(solver, monkeypatch): @@ -253,8 +280,9 @@ def test_a_law_with_no_named_quantity_has_no_temporaries(monkeypatch): u = uw.discretisation.MeshVariable("U0024g", mesh, 1, degree=1) pois = uw.systems.Poisson(mesh, u_Field=u) pois.constitutive_model = uw.constitutive_models.DiffusionModel - pois.constitutive_model.Parameters.diffusivity = 1.0 - pois.f = 2.0 + # named constants: leaves, not nodes + pois.constitutive_model.Parameters.diffusivity = uw.expression(r"k_{0024g}", 2.0) + pois.f = uw.expression(r"f_{0024g}", 1.0) header = _header(pois, monkeypatch) assert "out[0]" in header and "uwt_" not in header @@ -297,3 +325,137 @@ def test_a_repeated_condition_stays_a_condition(box): n[x] = sympy.Float(xv) for t, g in ((tree, lowered), (diff_wrt_field(tree, u), diff_wrt_field(lowered, u))): assert abs(_value(g, n) - _value(t, n)) <= 1.0e-12 * abs(_value(t, n)) + + +def test_a_deep_law_is_compiled_as_one_temporary_per_layer(monkeypatch): + """A law of twelve named layers, each using the one below twice: the expanded tree + doubles with every layer (4096 copies of the bottom), the graph has one temporary + per layer. The cost of the tree must not come back by any route that expands + nodes.""" + uw.reset_default_model() + mesh = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) + u = uw.discretisation.MeshVariable("U0024j", mesh, 1, degree=1) + c = uw.expression(r"c_{0024j}", 0.5, "constant") + layer = uw.expression(r"e_{0024j,0}", 1 + u.sym[0] ** 2, "bottom") + for k in range(1, 13): + layer = uw.expression(rf"e_{{0024j,{k}}}", c * layer + sympy.sqrt(layer), "layer") + pois = uw.systems.Poisson(mesh, u_Field=u) + pois.constitutive_model = uw.constitutive_models.DiffusionModel + pois.constitutive_model.Parameters.diffusivity = layer + pois.f = 1.0 + pois.consistent_jacobian = True + header = _header(pois, monkeypatch) + assert header.count("const double uwt_") <= 12 * 8, header.count("const double uwt_") + assert len(header) < 60_000, len(header) + + +def test_a_coordinate_without_its_c_name_is_spelled_from_its_index(): + """A coordinate leaf can be a UWCoordinate that SymPy's cache returned for its + equal base scalar, or a fresh base scalar, without the C name the mesh set: it is + written from its index and system (test_0850 and test_0851 met it).""" + from underworld3.utilities._jitextension import _spell_leaf + + mesh = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) + x = mesh.X[1] # the UWCoordinate wrapping mesh.N.y + base = x._original_base_scalar + saved = base.__dict__.pop("_ccodestr") + try: + assert _spell_leaf(x, {}, {}) == "petsc_x[1]" + assert _spell_leaf(base, {}, {}) == "petsc_x[1]" + finally: + base._ccodestr = saved + normal = mesh._Gamma.base_scalars()[0] + assert _spell_leaf(normal, {}, {}) == "petsc_n[0]" + + +def test_a_mesh_coordinate_beside_a_field_keeps_its_partial_derivative(box): + """``mesh.X`` coordinates are UWCoordinates, equal to the base scalars they wrap. In + a named quantity that reads one beside a field, the per-body common sub-expression + split rebuilt the coordinate with a cloned coordinate system, a symbol no longer + equal to ``mesh.N.x``: the derivative by the coordinate lost its explicit term.""" + mesh, T, v = box + X = mesh.X + u = T.sym[0] + q = uw.expression(r"q_{0024k}", X[0] * u + sympy.sin(X[1]) * u ** 2, + "coordinates beside a field") + lowered = jg.KernelGraph().lower(q * u) + tree = ex.unwrap_expression(q * u, mode="symbolic_keep_constants") + for w in (mesh.N.x, mesh.N.y): + d_tree = sympy.diff(tree, w) + n = _numbers([d_tree, tree]) + reference = _value(d_tree, n) + assert abs(reference) > 1.0e-3 + assert abs(_value(sympy.diff(lowered, w), n) - reference) <= 1.0e-12 * abs(reference) + + +def test_a_matrix_or_vector_valued_atom_is_expanded_in_place(box): + """A node is one scalar temporary, so an atom whose value is a matrix or a vector is + expanded in place, as the tree expanded it.""" + mesh, T, v = box + u = T.sym[0] + row = uw.expression(r"M_{0024l}", sympy.Matrix([[u, u ** 2]]), "a row") + lowered = jg.lower_callbacks([row], mesh)[0] + assert lowered.shape == (1, 2) + assert jg.expand_nodes(lowered) == sympy.Matrix([[u, u ** 2]]) + arrow = uw.expression(r"A_{0024l}", u * mesh.N.i + u ** 2 * mesh.N.j, "a vector") + lowered = jg.lower_callbacks([arrow], mesh)[0] + assert lowered.shape == (2, 1) + assert jg.expand_nodes(lowered) == sympy.Matrix([[u], [u ** 2]]) + + +def test_constancy_is_decided_on_the_graph(box, monkeypatch): + """Whether an atom is a constant is decided bottom-up on the graph, with the rule + of ``_is_truly_constant``, not by that function: it unwraps an atom completely, a + full expansion of everything under it, for every atom, which costs 2**depth on a + law whose every layer reads the one below twice.""" + import underworld3.utilities._jitextension as jx + + mesh, T, v = box + u = T.sym[0] + x = mesh.X[0] + c = uw.expression(r"c_{0024n}", 0.5, "constant") + nested = uw.expression(r"d_{0024n}", 2 * c + 1, "constant of a constant") + atoms = [ + c, nested, + uw.expression(r"e_{0024n}", nested * u, "reads a field"), + uw.expression(r"g_{0024n}", c * x, "reads a coordinate"), + uw.expression(r"h_{0024n}", uw.quantity(3.0, "m/s"), "a quantity"), + uw.expression(r"k_{0024n}", sympy.Matrix([[c, 2 * c]]), "a constant row"), + ] + expected = [jx._is_truly_constant(a, ex.UWexpression) for a in atoms] + assert expected == [True, True, False, False, True, True], expected + + calls = [] + original = jx._is_truly_constant + monkeypatch.setattr(jx, "_is_truly_constant", + lambda *args: calls.append(args) or original(*args)) + graph = jg.KernelGraph() + assert [graph.is_constant(a) for a in atoms] == expected + layer = uw.expression(r"e_{0024n,0}", 1 + u ** 2, "bottom") + for k in range(1, 17): + layer = uw.expression(rf"e_{{0024n,{k}}}", c * layer + sympy.sqrt(layer), "layer") + graph.lower(layer) + assert not calls, f"{len(calls)} complete unwraps" + + +def test_a_number_symbol_in_a_temporary_is_declared_once(monkeypatch): + """The C99 printer declares a number symbol (EulerGamma, Catalan) before the + expression that reads it; inside a temporary that declaration was written into + the temporary's own initialiser, which does not compile.""" + uw.reset_default_model() + mesh = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.5) + u = uw.discretisation.MeshVariable("U0024m", mesh, 1, degree=1) + pois = uw.systems.Poisson(mesh, u_Field=u) + pois.constitutive_model = uw.constitutive_models.DiffusionModel + k = uw.expression(r"k_{0024m}", 1 + sympy.EulerGamma * u.sym[0] ** 2 + sympy.Catalan, + "number symbols") + pois.constitutive_model.Parameters.diffusivity = k + pois.f = 1.0 + pois.consistent_jacobian = True + header = _header(pois, monkeypatch) # compiles + assert "uwt_0" in header + for body in header.split("\nvoid ")[1:]: + assert body.count("const double EulerGamma =") <= 1 diff --git a/tests/test_0103_jit_rampable_constants.py b/tests/test_0103_jit_rampable_constants.py index 5691da513..aec70a572 100644 --- a/tests/test_0103_jit_rampable_constants.py +++ b/tests/test_0103_jit_rampable_constants.py @@ -145,6 +145,11 @@ def test_same_named_placeholders_order_deterministically(): between MPI ranks and the cross-rank hash check aborts the run, on roughly half of launches. This test is the deterministic proxy: distinct sort keys, and a sum that canonicalises the same however it is written. + + Since #823 the placeholders no longer reach the generated C: the graph writes + each constant leaf as ``constants[i]`` directly, in canonical order, and + ``test_0105`` holds the C to one hash under every seed. This test keeps the + placeholders' own semantics, which code that substitutes them relies on. """ import sympy from underworld3.utilities._jitextension import _JITConstant diff --git a/tests/test_0105_jit_source_seed_independence.py b/tests/test_0105_jit_source_seed_independence.py index 0e5463162..abd410da6 100644 --- a/tests/test_0105_jit_source_seed_independence.py +++ b/tests/test_0105_jit_source_seed_independence.py @@ -104,3 +104,69 @@ def test_the_generated_module_is_the_same_under_every_hash_seed(seeds, tmp_path) "the emitted C depends on the Python hash seed, so MPI ranks will " f"disagree and getext will abort: {hashes}" ) + + +# A Newton viscoplastic Stokes (Drucker-Prager yield with floors, a temperature- +# dependent viscosity), whose kernels carry the graph's temporaries (#823): the +# emission order of named quantities is part of what must not depend on the seed. +_CHILD_NEWTON = textwrap.dedent( + """ + import pathlib, sys + import sympy, underworld3 as uw + import underworld3.utilities._jitextension as jx + + temporaries = [] + generate = jx.generate_c_source + + def counted(*args, **kwargs): + modname, codeguys, diag = generate(*args, **kwargs) + temporaries.append(dict(codeguys)["cy_ext.h"].count("const double uwt_")) + return modname, codeguys, diag + + jx.generate_c_source = counted + + mesh = uw.meshing.UnstructuredSimplexBox(cellSize=0.25, qdegree=2, regular=True) + v = uw.discretisation.MeshVariable("vn", mesh, mesh.dim, degree=2) + p = uw.discretisation.MeshVariable("pn", mesh, 1, degree=1) + T = uw.discretisation.MeshVariable("Tn", mesh, 1, degree=1) + stokes = uw.systems.Stokes(mesh, velocityField=v, pressureField=p) + stokes.constitutive_model = uw.constitutive_models.ViscoPlasticFlowModel + P = stokes.constitutive_model.Parameters + P.shear_viscosity_0 = uw.expression(r"\\eta_0", 1.0) * sympy.exp( + -uw.expression(r"\\theta", 3.0) * T.sym[0]) + P.yield_stress = uw.expression(r"C", 0.5) + uw.expression(r"\\mu", 0.6) * p.sym[0] + P.yield_stress_min = uw.expression(r"\\tau_{\\min}", 0.01) + P.shear_viscosity_min = uw.expression(r"\\eta_{\\min}", 1.0e-3) + stokes.consistent_jacobian = True + stokes.add_dirichlet_bc((1.0, 0.0), "Top") + stokes.add_dirichlet_bc((0.0, 0.0), "Bottom") + stokes._build() + + cache = pathlib.Path(sys.argv[1]) + names = sorted({q.name.split(".")[0] for q in cache.glob("*.so")}) + print("TEMPORARIES:" + str(sum(temporaries))) + print("MODULES:" + ",".join(names)) + """ +) + + +@pytest.mark.parametrize("seeds", [(0, 1, 2)]) +def test_a_newton_law_with_named_quantities_is_the_same_under_every_hash_seed( + seeds, tmp_path): + hashes = {} + for seed in seeds: + cache = tmp_path / f"cache_newton_{seed}" + cache.mkdir() + child = tmp_path / "child_newton.py" + child.write_text(_CHILD_NEWTON) + env = dict(os.environ, PYTHONHASHSEED=str(seed), UW_JIT_CACHE_DIR=str(cache)) + result = subprocess.run([sys.executable, str(child), str(cache)], + capture_output=True, text=True, env=env, timeout=900) + lines = dict(ln.split(":", 1) for ln in result.stdout.splitlines() + if ln.startswith(("MODULES:", "TEMPORARIES:"))) + assert "MODULES" in lines, ( + f"child failed under PYTHONHASHSEED={seed}\n" + f"stderr tail:\n{result.stderr[-2000:]}") + assert int(lines["TEMPORARIES"]) > 0, "no named quantity reached the C" + hashes[seed] = lines["MODULES"] + assert len(set(hashes.values())) == 1, hashes From 2819932ddaa13652d590ae6a27270a0d8b943665 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 19:56:04 +1100 Subject: [PATCH 11/14] Constancy by structure: a constant in exponent position ramps (#823, tier 2) The bottom-up constancy decision is structural. Where tier 1 decided by value, a collapsing expression such as (1 + T**2)**(-m) + 1 at m = 0 banked as one constants[] slot and raised when m ramped (test_0104, the "rampable constant in exponent position does not ramp" report). Now the expression is compiled as a node reading T and m's slot, and m ramps with no recompile: test_0104 pins that against a fresh build at each value, and keeps the slot-stops-being-constant error tested through a constant whose content is replaced by one that reads a field. test_0024 pins where the structural rule and _is_truly_constant differ. Design note: the final suite (2,910 passed), the final comparison against development with -fno-math-errno, and the manifest's one exception. Underworld development team with AI support from Claude Code --- .../design/jit-shared-graph-codegen.md | 18 ++++- src/underworld3/utilities/_jit_graph.py | 15 ++-- src/underworld3/utilities/_jitextension.py | 13 ++-- tests/test_0024_jit_graph_lowering.py | 14 ++-- .../test_0104_constant_slot_still_constant.py | 68 +++++++++++++------ 5 files changed, 89 insertions(+), 39 deletions(-) diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index 1c801c898..a7be06c19 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -283,8 +283,11 @@ unconvertible symbols and integration-point derivatives. The manifest is the set of constant atoms among the leaves of the emitted kernels, tested with the same predicate (`_is_truly_constant`, memoised once per lowering) and -ordered by the same key (name, then creation order). It contains every slot today's -manifest contains. It can contain more: where the tree cancels a constant across a name +ordered by the same key (name, then creation order). It contains every slot the tier 1 +manifest contained, with one exception: where a constant's value collapsed an enclosing +expression to a number (`(1 + T**2)**(-m) + 1` at `m = 0`), tier 1 gave the expression +one slot, which stopped being constant when `m` ramped; the graph gives `m` its slot +and compiles the expression, so `m` ramps (`test_0104`). It can contain more: where the tree cancels a constant across a name boundary ($A = c\,x$, then $A/c$), the graph keeps it, and the constant keeps its slot. An extra slot is harmless — it is packed and read — but each one must be traced to such a cancellation. The prototype finds none on its two fixtures, and the benchmark adds a @@ -440,6 +443,15 @@ the Mac does not show the effect. On the same server the #752 fixture agreed acr ranks in all 10 runs at np = 3 and at np = 4 on both routes, and `tests/parallel/ptest_jit_cache.py` passed at np = 4 on both. +**After step 4, against the improved JIT.** With the tree route gone, the comparison is a +build of this branch against a build of `development` with tier 1 and `-fno-math-errno` +(#834), on the Mac, where clang sets no `errno` anyway. The notch's pointwise setup is +1.95 s against 16.6 s, its C 21.8 KB against 4.0 MB, its Jacobian assembly 122 ms +against 184 ms; per quadrature point its Jacobian callbacks cost 0.49 µs against 2.67 µs +and its residual callbacks 107 ns against 188 ns. Linear Stokes emits byte-identical C; +the other fixtures take the same solver path, with setup and assembly equal or faster +within a few per cent. + **The operators agree to round-off.** `route_assemble.py` assembles the residual and the Jacobian on each route at the same state — rest (zero velocity, boundary values imposed), the tree's final state and the graph's final state — and compares them entry @@ -471,7 +483,7 @@ first run found two defects, both in the fault-network laws, each fixed with a u in `test_0024`: a repeated Piecewise condition shared as a node (a value, which Piecewise refuses as a condition), and a coordinate leaf without the C name the mesh sets. After step 4 removed the tree route: 3,082 passed, none failed; and again after the -second review's fixes: SUITE_FINAL. +second review's fixes: 2,910 passed, none failed. **The source is canonical.** On the notch, box and VEP fixtures the graph route's C is byte-identical under `PYTHONHASHSEED` 0, 1 and 2; `test_0024` checks that a law diff --git a/src/underworld3/utilities/_jit_graph.py b/src/underworld3/utilities/_jit_graph.py index cbc2b7bd5..bb619975a 100644 --- a/src/underworld3/utilities/_jit_graph.py +++ b/src/underworld3/utilities/_jit_graph.py @@ -174,11 +174,16 @@ def __init__(self): # ------------------------------------------------------------------ leaves def is_constant(self, atom): - """Whether a UWexpression is a ``constants[]`` leaf: the rule of - ``_is_truly_constant`` (its complete non-dimensional value reads no - coordinate, so no field either), decided bottom-up over the atoms it reads, - once per atom. ``_is_truly_constant`` unwraps each atom completely, which costs - the size of the expanded tree under it, for every atom.""" + """Whether a UWexpression is a ``constants[]`` leaf: its content reads no + coordinate (so no field either) and no non-constant atom, decided bottom-up + over the atoms it reads, once per atom. + + This is decided by structure. ``_is_truly_constant`` decides by value, after + a complete unwrap: it also calls ``(1 + T**2)**(-m) + 1`` a constant while + ``m`` is zero, and that slot stops being one when ``m`` ramps. Here such an + atom is a node reading ``T`` and ``m``'s slot, and ``m`` ramps without a + recompile. A complete unwrap per atom also costs the size of the expanded + tree under it, for every atom.""" hit = self._const.get(id(atom)) if hit is not None: return hit[1] diff --git a/src/underworld3/utilities/_jitextension.py b/src/underworld3/utilities/_jitextension.py index db2bf940c..1d40b5e69 100644 --- a/src/underworld3/utilities/_jitextension.py +++ b/src/underworld3/utilities/_jitextension.py @@ -551,17 +551,18 @@ def _pack_constants(manifest): # catastrophic: a zero diffusivity or viscosity diverges, and # nothing says why. # - # The usual cause is a nested atom that has been ramped. - # `(1 + T**2)**(-m) + 1` is the NUMBER 2 while m is zero, so it - # banks as one constant; ramp m and it depends on T again, but - # the kernel still expects a scalar. + # The usual cause is a constant atom whose content has been + # replaced by one that reads a field (c.sym = 1 + T**2) without + # a rebuild: the kernel still expects a scalar. (Constancy is + # decided by structure, so an atom ramped inside an expression, + # (1 + T**2)**(-m), keeps its own slot and ramps; #823.) raise RuntimeError( f"constants[] slot {idx} ({uw_expr.name!r}) no longer " f"reduces to a number, so the compiled kernel — which " f"treats it as a scalar constant — is out of date.\n" f" current content: {str(getattr(uw_expr, '_sym', uw_expr))[:160]}\n" - f"This usually means an atom nested inside it has been " - f"ramped, and the expression has stopped being constant. " + f"This usually means its content was replaced by an " + f"expression that reads a field or a coordinate. " f"Force a rebuild before solving again:\n" f" solver.is_setup = False\n" f" solver._needs_function_rewire = True\n" diff --git a/tests/test_0024_jit_graph_lowering.py b/tests/test_0024_jit_graph_lowering.py index 84e819315..5146c4193 100644 --- a/tests/test_0024_jit_graph_lowering.py +++ b/tests/test_0024_jit_graph_lowering.py @@ -406,10 +406,12 @@ def test_a_matrix_or_vector_valued_atom_is_expanded_in_place(box): def test_constancy_is_decided_on_the_graph(box, monkeypatch): - """Whether an atom is a constant is decided bottom-up on the graph, with the rule - of ``_is_truly_constant``, not by that function: it unwraps an atom completely, a - full expansion of everything under it, for every atom, which costs 2**depth on a - law whose every layer reads the one below twice.""" + """Whether an atom is a constant is decided bottom-up on the graph, by structure, + not by ``_is_truly_constant``: that unwraps an atom completely, a full expansion of + everything under it, for every atom, which costs 2**depth on a law whose every + layer reads the one below twice. The two agree except where a constant's current + value collapses an expression to a number: ``(1 + T**2)**(-m) + 1`` at ``m = 0`` + is a node reading ``m``, so ``m`` ramps (test_0104).""" import underworld3.utilities._jitextension as jx mesh, T, v = box @@ -426,6 +428,10 @@ def test_constancy_is_decided_on_the_graph(box, monkeypatch): ] expected = [jx._is_truly_constant(a, ex.UWexpression) for a in atoms] assert expected == [True, True, False, False, True, True], expected + m = uw.expression(r"m_{0024n}", 0, "zero for now") + collapsing = uw.expression(r"p_{0024n}", (1 + u ** 2) ** (-m) + 1, "2 while m is 0") + assert jx._is_truly_constant(collapsing, ex.UWexpression) + assert not jg.KernelGraph().is_constant(collapsing) calls = [] original = jx._is_truly_constant diff --git a/tests/test_0104_constant_slot_still_constant.py b/tests/test_0104_constant_slot_still_constant.py index 0b425f00e..4d1a7f24b 100644 --- a/tests/test_0104_constant_slot_still_constant.py +++ b/tests/test_0104_constant_slot_still_constant.py @@ -1,15 +1,17 @@ -"""A constants[] slot that stops being constant must say so, not pack a zero. - -An expression is given a ``constants[]`` slot because it resolved to a single -number when the kernel was compiled. Ramping an atom nested inside it can make -it depend on position again — the compiled kernel still reads a scalar, and -packing a zero into that slot hands the solve a zero coefficient. Silent, and -catastrophic: a zero diffusivity diverges and nothing says why. - -This is the failure behind the "rampable constant in exponent position does not -ramp" report. The atom is not compiled out; the ENCLOSING expression collapses -to a number while the atom is zero, banks as one constant, and then stops being -one. +"""A rampable constant in exponent position ramps; a constants[] slot that stops +being constant says so, not pack a zero. + +The "rampable constant in exponent position does not ramp" report: with the atom at +zero, the ENCLOSING expression ``(1 + T**2)**(-m) + 1`` is the number 2. When the JIT +decided constancy by the atom's current value, that expression banked as one +constants[] slot, and ramping ``m`` left the kernel reading a scalar where the law +depends on position again. The JIT now decides constancy by structure (#823, tier 2): +the expression reads a field, so it is compiled as a quantity reading ``T`` and the +``m`` slot, and ``m`` ramps without a recompile. + +A slot can still stop being constant: a constant atom whose content is replaced by one +that reads a field, without a rebuild. Packing a zero into it would hand the solve a +zero coefficient, silently; it must raise. """ import numpy as np @@ -42,25 +44,49 @@ def _build(initial): return uw, poisson, T, m -def test_a_slot_that_stops_being_constant_raises(): +def _mean_after_fresh_build(value): + uw, poisson, T, m = _build(value) + poisson.solve(zero_init_guess=True) + return float(np.asarray(T.data)[:, 0].mean()) + + +def test_an_atom_in_exponent_position_ramps_without_a_rebuild(): uw, poisson, T, m = _build(0.0) poisson.solve(zero_init_guess=True) + # m has a slot of its own: the collapsing expression is not banked as one + assert m in [expr for _, expr in poisson.constants_manifest] + compiled = poisson._current_jit_cache_key + + ramped = {} + for value in (0.5, 1.0): + m.sym = sympy.sympify(value) + poisson.solve(zero_init_guess=True) + assert poisson._current_jit_cache_key == compiled, "ramping m recompiled" + ramped[value] = float(np.asarray(T.data)[:, 0].mean()) + assert ramped[0.5] != pytest.approx(ramped[1.0]) + for value, mean in ramped.items(): + assert mean == pytest.approx(_mean_after_fresh_build(value), rel=1.0e-10), value - # Compiled while the whole expression was the number 2, so the diffusivity - # banks as a single scalar slot. (It is named for the parameter wrapper, - # not for the inner expression — the collector stops at the outermost thing - # that is truly constant and does not recurse past it.) + +def _slot_that_stops_being_constant(): + uw, poisson, T, m = _build(0.0) + c = uw.expression(r"c_slot", 1.0, "a constant coefficient") + poisson.constitutive_model.Parameters.diffusivity = c + poisson.solve(zero_init_guess=True) + # one slot: the diffusivity parameter, whose content is c assert len(poisson.constants_manifest) == 1 + c.sym = 1.0 + T.sym[0] ** 2 # now reads the field; the kernel reads a slot + return poisson - m.sym = sympy.sympify(0.5) # now depends on T again + +def test_a_slot_that_stops_being_constant_raises(): + poisson = _slot_that_stops_being_constant() with pytest.raises(RuntimeError, match="no longer.*reduces to a number"): poisson.solve(zero_init_guess=True) def test_the_message_names_the_slot_and_says_how_to_recover(): - uw, poisson, T, m = _build(0.0) - poisson.solve(zero_init_guess=True) - m.sym = sympy.sympify(0.5) + poisson = _slot_that_stops_being_constant() with pytest.raises(RuntimeError) as excinfo: poisson.solve(zero_init_guess=True) message = str(excinfo.value) From 7a63b708155ae28874f3f5e4cc542fb526f55ccd Mon Sep 17 00:00:00 2001 From: lmoresi Date: Thu, 8 Oct 2026 10:58:29 -1000 Subject: [PATCH 12/14] Keep the expanded JIT route beside the graph, switchable (#823, tier 2) Louis, 2026-10-09: a regression, or a case nobody considered, can only be told from a fault in the model if the old way can still be tried; if both routes are wrong, the model likely is. The graph route stays the default; the expanded route (the JIT before tier 2) is restored beside it, as development's code moved and not edited: - uw.use_jit_route("graph" | "expanded" | None) and UW_JIT_ROUTE choose the process default; solver.jit_route overrides it for one solver, which rebuilds its kernels at the next solve and keeps its state. getext and _jacobian_unwrap take the route. - generate_c_source builds its kernels with _graph_equations or _expanded_equations; the latter is development's class patching and per-kernel loop, verbatim, as are _reveal_constants, the scanning _extract_constants, _xreplace_shared, _unique_symbols, _collect_constant_atoms and the UW_JIT_CSE option. The expanded route's C is byte-identical to development's on all six fixtures (linear, power law, box, VEP, TI, notch). Tests: test_0026 (the setting; both routes give the same iteration counts and solutions on a viscoplastic box, Newton and Picard; switching one solver). The expanded route keeps its own tests: test_0022's tree guard test, test_0104's collapse-and-raise behaviour, test_0105's seed sweep on both routes. Graph tests pin route="graph". Docs: the design note's decision, jit-cache.md, and a "rule the JIT out" paragraph in the plasticity guide. Underworld development team with AI support from Claude Code --- .../design/jit-shared-graph-codegen.md | 32 +- docs/developer/guides/plasticity-solvers.md | 10 + docs/developer/subsystems/jit-cache.md | 12 +- src/underworld3/__init__.py | 40 + .../cython/petsc_generic_snes_solvers.pyx | 113 ++- src/underworld3/utilities/_jitextension.py | 699 ++++++++++++++++-- ...022_unwrap_memoised_matches_fixed_point.py | 55 +- tests/test_0024_jit_graph_lowering.py | 15 +- tests/test_0026_jit_route_switch.py | 124 ++++ tests/test_0103_jit_rampable_constants.py | 9 +- .../test_0104_constant_slot_still_constant.py | 18 +- .../test_0105_jit_source_seed_independence.py | 7 +- 12 files changed, 1046 insertions(+), 88 deletions(-) create mode 100644 tests/test_0026_jit_route_switch.py diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index a7be06c19..c75c34b0c 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -1,7 +1,9 @@ # Generating JIT kernels from the shared expression graph **Status**: Implemented on `feature/jit-graph-codegen`, 2026-10-08: the graph route is -the only JIT route (staging steps 2–4). Steps 2 and 3 were first measured against tier 1 +the default JIT route (staging steps 2–4), and the expanded route — the JIT before tier +2, generating byte-identical C — is kept beside it, selectable per process and per +solver (decision of 2026-10-09). Steps 2 and 3 were first measured against tier 1 end to end while both routes existed behind a switch ({ref}`jit-graph-in-the-library`). Tier 2 of [#823](https://github.com/underworldcode/underworld3/issues/823). Tier 1 ([#830](https://github.com/underworldcode/underworld3/pull/830), merged @@ -678,12 +680,13 @@ repository; the notch is the one exception, and its driver is named. updated. `jit-cache.md` already describes the cross-rank check as an abort and every rank as compiling; both stopped being true before this change. - Done, 2026-10-08. The field classes' patched C names (`ccode_patch_fns`), the - coordinate recovery and the unconvertible-symbol scan were replaced by an explicit map - from leaf to C built for each compile (`_leaf_spellings`, `_spell_leaf`), and the - opt-in `UW_JIT_CSE` path, which served only the expanded tree, was retired. Against - `development`, the library source gains 692 lines and loses 639: `_jit_graph.py` is - 457 of the gain, and `_jitextension.py` is about 400 lines shorter. + Done, 2026-10-08: on the graph route, an explicit map from leaf to C built for each + compile (`_leaf_spellings`, `_spell_leaf`) replaced the field classes' patched C + names, the coordinate recovery and the unconvertible-symbol scan. Amended + 2026-10-09: the expanded route was restored beside the graph, unchanged, as a + selectable fallback (see Decisions), so the step's deletions and the retirement of + `UW_JIT_CSE` were undone; only the unused `prepare_for_cache_key` and `_createext` + stay removed. Each step is benchmarked against the one before it and reviewed adversarially before the next begins. @@ -714,7 +717,20 @@ the next begins. - **2026-10-08: steps 2–4 land as one change** (Louis: "crash through, then compare with the existing, improved JIT"). The comparison is against `development` with tier 1 and `-fno-math-errno` (#834, PR #835), the competitor at its best. -- **2026-10-08: `UW_JIT_CSE` is retired** with the expanded tree it served. +- **2026-10-08: `UW_JIT_CSE` is retired** with the expanded tree it served + (superseded on 2026-10-09: it remains an option of the restored expanded route). +- **2026-10-09: the expanded route stays, as a fallback and a reference** (Louis: "we + will not be sure whether or not there is a regression, or more likely a case we did + not consider, unless we have the capacity to try it the old way. If both are wrong, + the model is likely wrong"). The graph is the default. `uw.use_jit_route("expanded")`, + or `UW_JIT_ROUTE=expanded` in the environment, selects the old route for the process, + and `solver.jit_route = "expanded"` for one solver, rebuilt at its next solve with its + state kept. The expanded route is `development`'s code, moved and not edited: its C is + byte-identical to `development`'s on the six fixtures, the notch included. + `UW_JIT_CSE` remains an option of that route. `test_0026` holds the two routes to + the same iteration counts and solutions on a viscoplastic box, Newton and Picard, and + the expanded route keeps its own tests (`test_0022`, `test_0103`, `test_0104`, + `test_0105`). Removing it is a later decision, once the graph has run in production. ## Questions for the maintainer diff --git a/docs/developer/guides/plasticity-solvers.md b/docs/developer/guides/plasticity-solvers.md index 3693f847f..eab59b7e7 100644 --- a/docs/developer/guides/plasticity-solvers.md +++ b/docs/developer/guides/plasticity-solvers.md @@ -129,6 +129,16 @@ node whose partial derivatives the chain rule composes (#823), so its compiled k 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 diff --git a/docs/developer/subsystems/jit-cache.md b/docs/developer/subsystems/jit-cache.md index ba0ae250f..20f35f8f8 100644 --- a/docs/developer/subsystems/jit-cache.md +++ b/docs/developer/subsystems/jit-cache.md @@ -35,7 +35,8 @@ processes. ## What is compiled -`getext()` lowers each callback onto the shared graph of named quantities +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 @@ -46,6 +47,14 @@ by a hash of the C each computes, with every leaf written as the C the kernel re 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 @@ -137,6 +146,7 @@ When `mpi.size > 1`: | `UW_JIT_CACHE_DIR` | Override the cache directory location | | `XDG_CACHE_HOME` | Used when `UW_JIT_CACHE_DIR` is unset | | `UW3_JIT_CFLAGS` | Replaces the kernels' default compile flags, `-O3 -g0 -fno-math-errno`; `-std=c99` is always kept. Use a lower level for a huge expression whose `-O3` compile is slow or runs out of memory (`-O1 -g0 -fno-math-errno`), or drop a flag the compiler rejects (nvc rejects `-g0` and `-fno-math-errno`). Without `-fno-math-errno`, gcc and clang on Linux cannot merge repeated `sqrt`, `pow` and `exp` calls (#834). The flags are part of the cache key. | +| `UW_JIT_ROUTE` | `graph` (default) or `expanded`, the JIT route; `uw.use_jit_route()` and `solver.jit_route` override it | ## Code references diff --git a/src/underworld3/__init__.py b/src/underworld3/__init__.py index ab31c3a5e..05dc2ca97 100644 --- a/src/underworld3/__init__.py +++ b/src/underworld3/__init__.py @@ -425,6 +425,46 @@ def use_nondimensional_scaling(enabled=True): _USE_NONDIMENSIONAL_SCALING = bool(enabled) +def use_jit_route(route): + """ + Set the JIT route solvers compile their kernels with, for this process. + + Parameters + ---------- + route : {"graph", "expanded"} or None + ``"graph"`` (the default) compiles each named quantity of a law once, as + one C temporary, and forms the Newton tangent through it by the chain rule. + ``"expanded"`` is the JIT before #823 tier 2: every named quantity is + expanded into one expression, differentiated and printed whole. ``None`` + returns to the ``UW_JIT_ROUTE`` environment variable (default ``"graph"``). + + Notes + ----- + A solver's own ``jit_route`` overrides this. A solver already set up keeps its + kernels until it is rebuilt. The two routes agree to round-off; if a model + misbehaves on one and not the other, the JIT is at fault. + + Examples + -------- + >>> uw.use_jit_route("expanded") # every solver built after this + >>> stokes.jit_route = "expanded" # or one solver, rebuilt at its next solve + + See Also + -------- + jit_route : The route in use + """ + from underworld3.utilities._jitextension import use_jit_route as _use + + _use(route) + + +def jit_route(): + """The process's JIT route, ``"graph"`` or ``"expanded"`` (see ``use_jit_route``).""" + from underworld3.utilities._jitextension import resolve_jit_route + + return resolve_jit_route() + + def is_nondimensional_scaling_active(): """ Check if non-dimensional scaling is currently enabled. diff --git a/src/underworld3/cython/petsc_generic_snes_solvers.pyx b/src/underworld3/cython/petsc_generic_snes_solvers.pyx index fbf46b422..2e397319d 100644 --- a/src/underworld3/cython/petsc_generic_snes_solvers.pyx +++ b/src/underworld3/cython/petsc_generic_snes_solvers.pyx @@ -51,6 +51,7 @@ from underworld3.function import expression as public_expression expression = lambda *x, **X: public_expression(*x, _unique_name_generation=True, **X) from underworld3.function._function import diff_wrt_field, derive_by_array_wrt_field +from underworld3.function.expressions import unwrap_expression as _unwrap_expression def _public_names(cls): @@ -63,10 +64,13 @@ def _public_names(cls): if obj is cls and not name.startswith("SNES_")) -def _jacobian_unwrap(expr): - r"""The Newton source of a residual flux: each non-constant UWexpression - replaced by its GUARDED node (``underworld3.utilities._jit_graph``), so that the - Jacobian derivative passes through it by the chain rule. +def _jacobian_unwrap(expr, route=None): + r"""The Newton source of a residual flux. On the graph route (the default) each + non-constant UWexpression is replaced by its GUARDED node + (``underworld3.utilities._jit_graph``), so that the Jacobian derivative passes + through it by the chain rule. On the expanded route (``route="expanded"``, the + JIT before #823 tier 2) every non-constant UWexpression is expanded down to the + constant atoms and the guard is applied to the expanded tree. Applied element-wise over a sympy ``Matrix``/``Array``. A node is an applied function of the leaves its value depends on (field values and gradients, @@ -98,10 +102,58 @@ def _jacobian_unwrap(expr): See ``docs/developer/design/jit-shared-graph-codegen.md`` and ``docs/developer/design/jacobian-consistent-tangent.md``. """ - from underworld3.utilities import _jit_graph + from underworld3.utilities._jitextension import resolve_jit_route + + if resolve_jit_route(route) == "graph": + from underworld3.utilities import _jit_graph + + graph = _jit_graph.KernelGraph() + f = lambda e: _jit_graph.guard_half_integer_powers(graph.lower(e, guarded=True)) + if isinstance(expr, sympy.MatrixBase): + return expr.applyfunc(f) + if isinstance(expr, sympy.NDimArray): + return sympy.Array([f(e) for e in expr], expr.shape) + return f(expr) # scalar expression + + # the expanded route: development's body, unchanged + eps2 = sympy.Float(1.0e-36) + + def _guard_sqrts(e): + # every HALF-INTEGER power: +1/2 (the invariant itself), -1/2 + # (its reciprocal in eta_pl = tau_y/(2 edot_II)), -3/2 (their + # derivatives), ... — all singular in value or derivative at a + # zero-argument state. + # The same bottom-up rebuild as `e.replace(query, value)`, memoised on node + # identity: the unwrapped flux repeats its shared sub-expressions as the same + # object, and `replace` walked every occurrence (measured 16 s of a 108 s + # notch compile, #823). + memo = {} + + def guard(n): + hit = memo.get(id(n)) + if hit is not None: + return hit[1] + out = n + args = getattr(n, "args", None) + if args: + new_args = tuple(guard(a) for a in args) + if any(a is not b for a, b in zip(args, new_args)) and args != new_args: + out = n.func(*new_args) + # replace(simultaneous=True): a rebuild that collapses to one of + # the changed arguments is not matched again + if any(out == a and a != b for a, b in zip(args, new_args)): + memo[id(n)] = (n, out) + return out + if (out.is_Pow and out.exp.is_Rational and out.exp.q == 2 + and out.args[0].free_symbols): + out = sympy.Pow(out.args[0] + eps2, out.exp) + memo[id(n)] = (n, out) + return out + + return guard(e) - graph = _jit_graph.KernelGraph() - f = lambda e: _jit_graph.guard_half_integer_powers(graph.lower(e, guarded=True)) + f = lambda e: _guard_sqrts( + _unwrap_expression(e, mode="symbolic_keep_constants")) if isinstance(expr, sympy.MatrixBase): return expr.applyfunc(f) if isinstance(expr, sympy.NDimArray): @@ -448,6 +500,47 @@ class SolverBaseClass(uw_object): f"consistent_jacobian must be False, True or 'continuation'; " f"got {mode!r}") + @property + def jit_route(self): + r"""The JIT route this solver compiles its kernels with: ``"graph"``, + ``"expanded"``, or ``None`` (default) for the process default + (``uw.use_jit_route``, else the ``UW_JIT_ROUTE`` environment variable, else + ``"graph"``). + + ``"graph"`` compiles each named quantity once, as one C temporary, and forms + the Newton tangent through it by the chain rule. ``"expanded"`` is the JIT + before #823 tier 2: every named quantity is expanded into one expression, + differentiated and printed whole. Use it as a fallback and as a reference: if + a model misbehaves on one route and not the other, the JIT is at fault; if on + both, look at the model. + + Setting it rebuilds this solver's kernels at the next solve; the solution and + the warm start are kept. + + Raises + ------ + ValueError + On assignment of anything other than ``None``, ``"graph"`` or + ``"expanded"``. + """ + return getattr(self, "_jit_route", None) + + @jit_route.setter + def jit_route(self, route): + from underworld3.utilities._jitextension import resolve_jit_route + + if route is not None: + route = resolve_jit_route(route) + if route != self.jit_route: + self._jit_route = route + self._needs_function_rewire = True + + def _jit_route_in_use(self): + """The route this solver's next build compiles with.""" + from underworld3.utilities._jitextension import resolve_jit_route + + return resolve_jit_route(self.jit_route) + def _jacobian_source(self, expr, newton_expr=None): """Prepare a residual flux for Jacobian differentiation. @@ -471,7 +564,7 @@ class SolverBaseClass(uw_object): if not mode: return expr if newton_expr is None: - newton_expr = _jacobian_unwrap(expr) + newton_expr = _jacobian_unwrap(expr, route=self._jit_route_in_use()) if mode == "continuation": a = self._get_newton_alpha() if isinstance(expr, sympy.MatrixBase): @@ -4295,6 +4388,7 @@ class SNES_Scalar(SolverBaseClass): prim_field_list, verbose=verbose, debug=debug, + route=self._jit_route_in_use(), ) self.compiled_extensions = _getext_result.ptrobj self.ext_dict = _getext_result.fn_dicts @@ -5316,6 +5410,7 @@ class SNES_Vector(SolverBaseClass): prim_field_list, verbose=verbose, debug=debug, + route=self._jit_route_in_use(), ) self.compiled_extensions = _getext_result.ptrobj self.ext_dict = _getext_result.fn_dicts @@ -6047,6 +6142,7 @@ class SNES_MultiComponent(SolverBaseClass): prim_field_list, verbose=verbose, debug=debug, + route=self._jit_route_in_use(), ) self.compiled_extensions = _getext_result.ptrobj self.ext_dict = _getext_result.fn_dicts @@ -8628,6 +8724,7 @@ class SNES_Stokes_SaddlePt(SolverBaseClass): verbose=verbose, debug=debug, debug_name=debug_name, + route=self._jit_route_in_use(), # Disk cache + rank-0-only compile: under MPI only rank 0 invokes # cc and publishes to the shared cache dir; the other ranks load # the compiled module. Without this, every rank compiles its own diff --git a/src/underworld3/utilities/_jitextension.py b/src/underworld3/utilities/_jitextension.py index 1d40b5e69..8b88f85ba 100644 --- a/src/underworld3/utilities/_jitextension.py +++ b/src/underworld3/utilities/_jitextension.py @@ -201,6 +201,46 @@ def _abi_salt(): return f"petsc={petsc_ver}|uw={uw_ver}" +# ============================================================================ +# The two JIT routes +# ============================================================================ +# +# "graph" (default, #823 tier 2): each named non-constant quantity is one C +# temporary, and the Newton tangent passes through it by the chain rule +# (_jit_graph). "expanded" (the JIT before tier 2): every named quantity is expanded +# into one expression tree, differentiated and printed whole. The expanded route is +# kept as a fallback and a reference: a model that misbehaves can be solved the old +# way, and if both routes agree the cause is the model, not the JIT. +# ============================================================================ + +JIT_ROUTES = ("graph", "expanded") +_jit_route_default = None # set by uw.use_jit_route(); None reads UW_JIT_ROUTE + + +def resolve_jit_route(route=None): + """The JIT route to compile with: ``route`` if given, else the process default + (``uw.use_jit_route``), else the ``UW_JIT_ROUTE`` environment variable, else + ``"graph"``.""" + if route is None: + route = _jit_route_default + if route is None: + route = os.environ.get("UW_JIT_ROUTE", "graph").strip().lower() or "graph" + if route not in JIT_ROUTES: + raise ValueError(f"JIT route must be one of {JIT_ROUTES}; got {route!r}") + return route + + +def use_jit_route(route): + """Set the process default JIT route, ``"graph"`` or ``"expanded"``; ``None`` + returns to the ``UW_JIT_ROUTE`` environment variable (default ``"graph"``). + Solvers built afterwards use it unless their own ``jit_route`` is set; a solver + already set up keeps its kernels until it is rebuilt.""" + global _jit_route_default + if route is not None: + resolve_jit_route(route) # validates + _jit_route_default = route + + # ============================================================================ # JIT Callback Set # ============================================================================ @@ -281,6 +321,34 @@ def counts(self): len(self.bd_residual), len(self.bd_jacobian)) +def _reveal_constants(fn): + """Expand non-constant UWexpressions to fixpoint, KEEPING truly-constant + atoms symbolic — so constants nested at ANY depth surface as atoms. + + This must run BEFORE the ``constants[]`` substitution: a top-level + ``xreplace`` cannot see a constant hidden inside a nested UWexpression + (every ``Parameters.*`` value is template-wrapped in one), so nested + constants were silently folded to C literals while the manifest still + listed them as live — issue #302. The keep-constants predicate here is + the same ``_is_truly_constant`` used to build the manifest, so the set + of atoms revealed is exactly the set the manifest routes to + ``constants[]``. + """ + from underworld3.function.expressions import ( + unwrap_expression, + UWDerivativeExpression, + ) + + if fn is None: + return fn + if isinstance(fn, UWDerivativeExpression): + fn = fn.doit() + if isinstance(fn, sympy.MatrixBase): + return fn.applyfunc( + lambda e: unwrap_expression(e, mode='symbolic_keep_constants')) + return unwrap_expression(fn, mode='symbolic_keep_constants') + + # ============================================================================ # JIT Constants Support # ============================================================================ @@ -352,23 +420,52 @@ def _ccode(self, printer): def _extract_constants(all_fns, mesh): - """The ``constants[]`` manifest of a list of callback expressions: the constant - atoms their lowered kernels read (``_jit_graph``), ordered as ``_manifest_from`` - orders them. + """Extract constant UWexpressions from a list of pre-unwrap functions: the + ``constants[]`` manifest of the expanded route. + + Scans all expressions for UWexpression atoms where is_constant_expr() + is True (no spatial/field dependencies). Assigns deterministic indices + sorted by expression name for MPI consistency. + + Parameters + ---------- + all_fns : tuple of sympy expressions + The raw (pre-unwrap) function list. + mesh : underworld3.discretisation.Mesh + The mesh (currently unused, reserved for future mesh.t support). Returns ------- list of (int, UWexpression) Ordered mapping from constants[] index to UWexpression reference. dict - Mapping from UWexpression to _JITConstant symbol. + Mapping from UWexpression to _JITConstant symbol for substitution. """ - lowered = _jit_graph.lower_callbacks([fn for fn in all_fns if fn is not None], mesh) - return _manifest_from(_jit_graph.constant_leaves(lowered)) + from underworld3.function.expressions import ( + is_constant_expr, + extract_expressions, + UWexpression, + ) + + constant_exprs = set() + + for fn in all_fns: + if fn is None: + continue + + # Handle Matrix expressions + if isinstance(fn, sympy.MatrixBase): + for elem in fn: + _collect_constant_atoms(elem, constant_exprs, is_constant_expr, UWexpression) + else: + _collect_constant_atoms(fn, constant_exprs, is_constant_expr, UWexpression) + + return _manifest_from(constant_exprs) def _manifest_of(lowered): - """``(manifest, subs_map)`` of callbacks already lowered (``_jit_graph``).""" + """``(manifest, subs_map)`` of callbacks lowered onto the graph (``_jit_graph``): + the manifest of the graph route.""" return _manifest_from(_jit_graph.constant_leaves(lowered)) @@ -407,6 +504,37 @@ def _manifest_from(constant_exprs): return manifest, subs_map +def _xreplace_shared(expr, rule): + """``expr.xreplace(rule)`` (a sympy expression or Matrix), visiting each node + OBJECT once. + + A compiled kernel repeats its shared sub-expressions as the same object (the + memoised unwrap inserts one object wherever an atom occurs, #823), and + ``xreplace`` rebuilds every occurrence: 4.3 s of the notch C generation for the + constants[] substitution alone. The result is the same expression. + """ + memo = {} + + def walk(e): + hit = memo.get(id(e)) + if hit is not None: + return hit[1] + if e in rule: + out = rule[e] + elif isinstance(e, sympy.Basic) and e.args: + new_args = tuple(walk(a) if isinstance(a, sympy.Basic) else a for a in e.args) + changed = any(n is not a for n, a in zip(new_args, e.args)) + out = e.func(*new_args) if changed else e + else: + out = e + memo[id(e)] = (e, out) + return out + + if isinstance(expr, sympy.MatrixBase): + return expr.applyfunc(walk) + return walk(expr) + + def _holds_instance(expr, types): """Whether ``expr`` holds a node of ``types``, visiting each node object once.""" from sympy.tensor.array import NDimArray @@ -470,6 +598,36 @@ def _without_dirac_deltas(expr, where): return expr.xreplace({d: sympy.S.Zero for d in deltas}) +def _unique_symbols(expr): + """The Symbol atoms of ``expr`` (a sympy expression, Matrix or Array): the same set + as ``expr.atoms(sympy.Symbol)``, found by visiting each node OBJECT once. + + ``atoms`` walks every occurrence of every node. An unwrapped constitutive law + repeats its shared sub-expressions as the SAME Python object (the memoised unwrap + inserts one object wherever an atom occurs, #823), so an identity walk is + proportional to the shared graph rather than the expanded tree: measured on the + Spiegelman notch kernels, ``atoms`` was 35 s of a 108 s compile. + """ + # keyed by id, holding the object so that no id is reused while the walk runs + seen = {} + found = set() + stack = [expr] + while stack: + e = stack.pop() + if id(e) in seen: + continue + seen[id(e)] = e + if isinstance(e, (sympy.MatrixBase, sympy.NDimArray)): + stack.extend(e) + continue + if isinstance(e, sympy.Symbol): + found.add(e) + continue + if isinstance(e, sympy.Basic): + stack.extend(e.args) + return found + + def _is_truly_constant(expr, UWexpression): """Check if a UWexpression resolves to a pure constant (no spatial deps). @@ -504,6 +662,11 @@ def _is_truly_constant(expr, UWexpression): return False if isinstance(sym, sympy.Function): return False + # UnderworldFunction symbols have _ccodestr pointing to petsc arrays + if hasattr(sym, '_ccodestr') and not isinstance(sym, _JITConstant): + ccode = sym._ccodestr + if 'petsc_u' in ccode or 'petsc_a' in ccode or 'petsc_x' in ccode or 'petsc_n' in ccode: + return False # Other UWexpressions that didn't fully unwrap — not constant if isinstance(sym, UWexpression): return False @@ -511,6 +674,31 @@ def _is_truly_constant(expr, UWexpression): return True +def _collect_constant_atoms(expr, result_set, is_constant_expr, UWexpression): + """Recursively collect constant UWexpression atoms from an expression.""" + + if isinstance(expr, UWexpression): + if _is_truly_constant(expr, UWexpression): + result_set.add(expr) + return # Don't recurse into constant expressions + # Non-constant UWexpression: check its inner sym for nested constants + if hasattr(expr, '_sym') and expr._sym is not None: + _collect_constant_atoms(expr._sym, result_set, is_constant_expr, UWexpression) + return + + if not hasattr(expr, 'atoms'): + return + + # Check all UWexpression atoms + for atom in _stable_sorted(_unique_symbols(expr)): + if isinstance(atom, UWexpression) and _is_truly_constant(atom, UWexpression): + result_set.add(atom) + elif isinstance(atom, UWexpression): + # Non-constant UWexpression: recurse into its sym + if hasattr(atom, '_sym') and atom._sym is not None: + _collect_constant_atoms(atom._sym, result_set, is_constant_expr, UWexpression) + + def _pack_constants(manifest): """Pack current values from a constants manifest into a flat array. @@ -631,6 +819,7 @@ def getext( debug=False, debug_name=None, cache=True, + route=None, ): """Compile (or retrieve cached) JIT extension for PETSc pointwise functions. @@ -644,6 +833,8 @@ def getext( primary_field_list : iterable Variables that map to PETSc primary arrays (``petsc_u[]``). All others map to auxiliary arrays (``petsc_a[]``). + route : {"graph", "expanded"}, optional + The JIT route (``resolve_jit_route``); the process default when omitted. Returns ------- @@ -659,10 +850,15 @@ def getext( # Extract constant UWexpressions that are routed through PETSc's # constants[] array. Value changes don't affect the C source — they # only alter what we pass to PetscDSSetConstants at solve time. - # Each callback is lowered onto the shared graph of named quantities once - # (``_jit_graph``, #823); the manifest is the constant leaves of what was lowered. - lowered_fns = _jit_graph.lower_callbacks(callbacks.flat(), mesh) - constants_manifest, constants_subs_map = _manifest_of(lowered_fns) + route = resolve_jit_route(route) + if route == "graph": + # Each callback is lowered onto the shared graph of named quantities once + # (``_jit_graph``, #823); the manifest is the constant leaves of the result. + lowered_fns = _jit_graph.lower_callbacks(callbacks.flat(), mesh) + constants_manifest, constants_subs_map = _manifest_of(lowered_fns) + else: + lowered_fns = None + constants_manifest, constants_subs_map = _extract_constants(callbacks.flat(), mesh) if debug and underworld3.mpi.rank == 0: if constants_manifest: @@ -684,6 +880,7 @@ def getext( verbose=verbose, debug=debug, debug_name=debug_name, + route=route, ) gen_randstr = diag["randstr"] @@ -881,6 +1078,425 @@ def getext( _NUMBER_DECLARATION = re.compile(r"const double [A-Za-z_]\w* = [^;]*;") +def _graph_equations(fns, lowered_fns, constants_subs_map, mesh, primary_field_list, + printer, verbose): + """The ``(name, C body)`` of each kernel on the graph route: one temporary per + distinct computation, then the outputs, every leaf written by an explicit map + built for this compile instead of C names patched onto the field classes.""" + spellings = _leaf_spellings(mesh, primary_field_list) + + eqns = [] + for index, fn_original in enumerate(fns): + unspellable = [] + + def spell(leaf): + try: + return _spell_leaf(leaf, spellings, constants_subs_map) + except _Unspellable: + unspellable.append(leaf) + return f"?{sympy.srepr(leaf)}" + + temporaries, fn = _jit_graph.emit(lowered_fns[index], spell) + if unspellable: + raise RuntimeError(_unconvertible_message( + _stable_sorted(set(unspellable)), index, fn_original)) + + if verbose: + # the kernel as mathematics (named quantities as nodes); the C is in the header + print("Processing JIT {:4d} / {}".format(index, lowered_fns[index])) + + out = sympy.MatrixSymbol("out", *fn.shape) + eqn = ("eqn_" + str(index), _print_kernel(printer, temporaries, fn, out)) + + if eqn[1].startswith("// Not supported in C:"): + spliteqn = eqn[1].split("\n") + raise RuntimeError( + f"Error encountered generating JIT extension:\n" + f"{spliteqn[0]}\n" + f"{spliteqn[1]}\n" + f"This is usually because code generation for a Sympy function (or its derivative) is not supported.\n" + f"Please contact the developers." + f"---" + f"The ID of the JIT component that failed is {index}" + f"The decription of the JIT component that failed:\n {fn}" + ) + eqns.append(eqn) + + return eqns + + +def _expanded_equations(fns, constants_subs_map, mesh, primary_field_list, printer, + verbose): + """The ``(name, C body)`` of each kernel on the expanded route: the JIT before + #823 tier 2, unchanged. Every named quantity is unwrapped into one expression and + printed whole; field symbols are given their C names by patching their classes.""" + from underworld3 import VarType + + # `_ccode` patching + def ccode_patch_fns(varlist, prefix_str, component_offsets=None): + """ + This function patches uw functions with the necessary ccode + routines for the code printing. + + For a `varlist` consisting of 2d velocity & pressure variables, + for example, it'll generate routines which write the following, + where `prefix_str="petsc_u"`: + V_x : "petsc_u[0]" + V_y : "petsc_u[1]" + P : "petsc_u[2]" + V_x_x : "petsc_u_x[0]" + V_x_y : "petsc_u_x[1]" + V_y_x : "petsc_u_x[2]" + V_y_y : "petsc_u_x[3]" + P_x : "petsc_u_x[4]" + P_y : "petsc_u_x[5]" + + Params + ------ + varlist: list + The variables to patch. Note that *all* the variables in the + corresponding `PetscDM` must be included. They must also be + ordered according to their `field_id`. + prefix_str: str + The string prefix to write. + component_offsets: dict, optional + Component offset of every field in the DM, by ``field_id`` + (see ``_aux_component_offsets``). When given, each variable + is patched from ITS OWN field's offset instead of a running + count over ``varlist``: a field whose Python variable has + been dropped stays in the DM and still occupies its slots, + so a running count would shift every later variable onto + the wrong data. + """ + u_i = 0 # variable increment + u_x_i = 0 # variable gradient increment + lambdafunc = lambda self, printer: self._ccodestr + + def _no_derivative(self, printer): + # An integration-point variable has no gradient (its tabulated + # derivative is identically zero), so a derivative of its symbol + # in a weak form would be a silent zero. Refuse at code generation. + raise RuntimeError( + f"{self.__class__.__name__}: derivative of an integration-point " + "variable has no meaning (the field is defined only at the " + "quadrature points), so the gradient here would be a silent " + "zero. This is refused in a WEAK FORM only, where the " + "discretisation is yours to choose: build the variable with " + "proxy_location='cells' instead, whose level sets are a " + "least-squares polynomial per cell and differentiate directly. " + "uw.function.evaluate() of the same derivative does answer: as " + "a query it recovers the gradient from a per-cell fit for you." + ) + + for var in varlist: + is_ip = getattr(var, "is_integration_point", False) + dfunc = _no_derivative if is_ip else lambdafunc + if component_offsets is not None: + u_i = component_offsets[var.field_id] + u_x_i = u_i * mesh.cdim + if var.vtype == VarType.SCALAR: + # monkey patch this guy into the function + type(var.fn)._ccodestr = f"{prefix_str}[{u_i}]" + type(var.fn)._ccode = lambdafunc + u_i += 1 + # Now patch the gradient components. The gradient of a + # field on the mesh lives in the embedded coordinate + # space (cdim-dim), so iterate to cdim — not dim. For + # volume meshes ``dim == cdim`` so this is unchanged; + # for manifold meshes (e.g. SphericalManifold dim=2, + # cdim=3) the third partial ``f_{,2}`` exists and + # needs to be wired to ``u_x[2]``. + for ind in range(mesh.cdim): + # Note that var.fn._diff[ind] returns the class, so we don't need type(var.fn._diff[ind]) + var.fn._diff[ind]._ccodestr = f"{prefix_str}_x[{u_x_i}]" + var.fn._diff[ind]._ccode = dfunc + u_x_i += 1 + elif ( + var.vtype == VarType.VECTOR + or var.vtype == VarType.TENSOR + or var.vtype == VarType.SYM_TENSOR + or var.vtype == VarType.MATRIX + ): + # Pull out individual sub components + for comp in var.sym_1d: + # monkey patch + type(comp)._ccodestr = f"{prefix_str}[{u_i}]" + type(comp)._ccode = lambdafunc + u_i += 1 + # Iterate to cdim (embedded coord dim) — see the + # scalar branch above for the dim != cdim reason. + for ind in range(mesh.cdim): + # Note that var.fn._diff[ind] returns the class, so we don't need type(var.fn._diff[ind]) + comp._diff[ind]._ccodestr = f"{prefix_str}_x[{u_x_i}]" + comp._diff[ind]._ccode = dfunc + u_x_i += 1 + else: + raise RuntimeError( + f"Unsupported type {var.vtype} for code generation. Please contact developers." + ) + + # Patch in `_code` methods. Note that the order here + # is important, as the secondary call will overwrite + # those patched in the first call. + + ccode_patch_fns(_stable_sorted(mesh.vars.values()), "petsc_a", + component_offsets=_aux_component_offsets(mesh)) + ccode_patch_fns(primary_field_list, "petsc_u") + + # Also patch `BaseScalar` types. Nothing fancy - patch the overall type, + # make sure each component points to the correct PETSc data + + ## This is set up in the mesh at the moment but this does seem to be the wrong place + + # mesh.N.x._ccodestr = "petsc_x[0]" + # mesh.N.y._ccodestr = "petsc_x[1]" + # mesh.N.z._ccodestr = "petsc_x[2]" + + # # Surface integrals also have normal vector information as petsc_n + + # mesh.Gamma_N.x._ccodestr = "petsc_n[0]" + # mesh.Gamma_N.y._ccodestr = "petsc_n[1]" + # mesh.Gamma_N.z._ccodestr = "petsc_n[2]" + + def _basescalar_ccode(self, printer): + """C code for coordinate symbols, with fallback for new instances. + + sympy.simplify() may create new BaseScalar/UWCoordinate instances + that lack _ccodestr. We recover it from the coordinate's _id attribute + which stores (index, system_name). + """ + if hasattr(self, '_ccodestr'): + return self._ccodestr + # Fallback: compute from _id + idx = self._id[0] + system_name = str(self._id[1]) + if 'Gamma' in system_name: + return f"petsc_n[{idx}]" + else: + return f"petsc_x[{idx}]" + + type(mesh.N.x)._ccode = _basescalar_ccode + # Gamma base scalars (un-normalised face normal) — ensure ccode is registered + Gamma_scalars = mesh._Gamma.base_scalars() + if type(Gamma_scalars[0]) is not type(mesh.N.x): + type(Gamma_scalars[0])._ccode = _basescalar_ccode + + eqns = [] + for index, fn in enumerate(fns): + + # Save original for debugging + fn_original = fn + + # --- Gate the UW lowering (issue #302 pipeline) on the presence of + # UW-expression atoms. Plain-sympy components — the derivative + # blocks, which dominate the expression size — have no UW atoms, so + # the reveal / validate / xreplace / unwrap pipeline (≈5 full + # traversals per component) would be pure overhead: skip it entirely + # when there is nothing to lower. + from underworld3.function.expressions import UWexpression as _UWexpr + from underworld3.function.expressions import UWDerivativeExpression as _UWderiv + + _needs_lowering = ( + isinstance(fn, (_UWexpr, _UWderiv)) + # `has` is a bare traversal (no atom-set build) — the atoms() + # form built a set of every node, which cost seconds per 100k-node + # Jacobian component (measured ~20 s on a large collision model). + or (hasattr(fn, "has") and fn.has(_UWexpr)) + or not isinstance(fn, (sympy.MatrixBase, sympy.MatrixExpr)) + ) + if _needs_lowering: + # Phase 1: reveal constants nested inside other UWexpressions, so + # the substitution below can reach them. A top-level + # xreplace missed constants inside template-wrapped + # parameters and baked them as C literals while the + # manifest listed them. + fn = _reveal_constants(fn) + + # A truly-constant atom the manifest does NOT know about would be + # silently folded to a literal in phase 3 — the manifest and the + # C source must never disagree (issue #302). + if constants_subs_map is not None and hasattr(fn, 'atoms'): + unmanifested = [ + a.name for a in _stable_sorted(_unique_symbols(fn)) + if isinstance(a, _UWexpr) + and _is_truly_constant(a, _UWexpr) + and a not in constants_subs_map + ] + if unmanifested: + raise RuntimeError( + f"JIT constants manifest is incomplete: constant expression(s) " + f"{unmanifested} appear in a kernel but have no constants[] " + f"slot — they would be baked into the C source (issue #302)." + ) + + # Phase 2: Substitute constant UWexpressions with _JITConstant symbols + # These survive into C code as constants[i] + if constants_subs_map and fn is not None: + try: + fn = _xreplace_shared(fn, constants_subs_map) if hasattr(fn, 'xreplace') else fn + except Exception: + pass + + # Phase 3: Unwrap remaining non-constant UWexpressions to numerical values + fn = underworld3.function.expressions.unwrap(fn, keep_constants=False, return_self=False) + + # A manifested constant surviving to here bypassed its constants[] + # slot and is about to be baked — refuse rather than freeze the + # parameter silently (issue #302). + if constants_subs_map and hasattr(fn, 'atoms'): + baked = [a.name for a in _stable_sorted(_unique_symbols(fn)) + if a in constants_subs_map] + if baked: + raise RuntimeError( + f"Manifested constant(s) {baked} were not routed through " + f"constants[] and would be baked into the C source " + f"(issue #302)." + ) + + if isinstance(fn, sympy.vector.Vector): + fn = fn.to_matrix(mesh.N)[0 : mesh.dim, 0] + elif isinstance(fn, sympy.vector.Dyadic): + fn = fn.to_matrix(mesh.N)[0 : mesh.dim, 0 : mesh.dim] + else: + fn = sympy.Matrix([fn]) + + # === COORDINATE SYMBOL RECOVERY === + # When sympy.simplify() manipulates expressions containing coordinate + # symbols (BaseScalar/UWCoordinate), it may create NEW instances that + # lack the _ccodestr attribute set during mesh initialization. + # This commonly occurs with coordinate-dependent constitutive models + # (e.g., TransverseIsotropicFlowModel with a radial director). + # We recover _ccodestr from the coordinate's _id attribute. + from sympy.vector.scalar import BaseScalar + + free_syms = tuple(_stable_sorted(fn.free_symbols)) + for sym in free_syms: + if isinstance(sym, BaseScalar) and not hasattr(sym, '_ccodestr'): + idx = sym._id[0] # 0, 1, or 2 for x, y, z + system_name = str(sym._id[1]) + if 'Gamma' in system_name: + sym._ccodestr = f"petsc_n[{idx}]" + else: + sym._ccodestr = f"petsc_x[{idx}]" + + # === JIT VALIDATION GATEWAY === + # Check for symbols that cannot be converted to C code. + # Expected symbols (coordinates) have _ccodestr attribute set. + # Unexpected symbols indicate malformed expressions from user code. + unconvertible_symbols = [] + for sym in free_syms: + # Check if this symbol can be converted to C code + if not hasattr(sym, '_ccodestr'): + unconvertible_symbols.append(sym) + + if unconvertible_symbols: + # Build a helpful error message + sym_details = [] + for sym in unconvertible_symbols: + detail = f" - {sym} (type: {type(sym).__name__})" + if hasattr(sym, 'units'): + detail += f" [has units: {sym.units}]" + if hasattr(sym, 'value'): + detail += f" [value: {sym.value}]" + sym_details.append(detail) + + raise RuntimeError( + f"\n{'='*70}\n" + f"JIT COMPILATION ERROR: Expression contains unconvertible symbols\n" + f"{'='*70}\n\n" + f"The following symbols could not be converted to C code:\n" + + "\n".join(sym_details) + "\n\n" + f"This usually means:\n" + f" 1. A UWexpression or UWQuantity was not properly expanded\n" + f" 2. An arithmetic operation failed (e.g., Matrix * UWexpression)\n" + f" 3. A symbolic function is missing from the expression tree\n\n" + f"Expression index: {index}\n" + f"Original expression: {fn_original}\n" + f"After unwrap: {fn}\n\n" + f"TIP: Check that all expression operations (*, /, +, -) produce\n" + f"valid SymPy expressions. For example, ensure scalar * Matrix\n" + f"and not Matrix * scalar when using UWexpression objects.\n" + f"{'='*70}" + ) + + if verbose: + print("Processing JIT {:4d} / {}".format(index, fn)) + # Enhanced debugging output for remaining (valid) free symbols + if free_syms: + print(" Free symbols (all convertible):") + for sym in free_syms: + print(f" - {sym} (type: {type(sym).__name__}, _ccodestr: {getattr(sym, '_ccodestr', 'N/A')})") + + out = sympy.MatrixSymbol("out", *fn.shape) + + # CSE before printing: shared subexpressions become ``double xN = ...;`` + # temps evaluated in dependency order, so the generated C — and hence + # the codegen time, gcc memory/time, and .so size — collapses on large + # expressions (measured: monster Jacobian output ~460k nodes -> ~30k). + # Semantics-preserving: temps are exact aliases of repeated + # subexpressions, so the generated kernel evaluates identical values. + # Opt in with UW_JIT_CSE=1 (default is off to preserve original behavior). + if os.environ.get("UW_JIT_CSE") in ("1", "true", "True", "yes", "YES"): + from sympy.simplify.cse_main import cse + from sympy.vector.scalar import BaseScalar + + _repl, _red = cse([fn]) + if _repl: + # cse may mint NEW coordinate instances (BaseScalar / + # UWCoordinate wrappers) that lack the mesh-set _ccodestr; + # recover it from their _id (same scheme as the + # COORDINATE SYMBOL RECOVERY above). + def _patch_coords(expr): + for _sym in set(expr.free_symbols): + _target = getattr(_sym, "_original_base_scalar", _sym) + if isinstance(_target, BaseScalar) and not hasattr( + _target, "_ccodestr" + ): + _idx = _target._id[0] + _sys = str(_target._id[1]) + _target._ccodestr = ( + f"petsc_n[{_idx}]" + if "Gamma" in _sys + else f"petsc_x[{_idx}]" + ) + + for _t_sym, _t_expr in _repl: + _patch_coords(_t_expr) + _patch_coords(_red[0]) + + _temp_code = "\n".join( + "double {} = {};".format( + printer.doprint(t_sym), printer.doprint(t_expr) + ) + for t_sym, t_expr in _repl + ) + _red_code = printer.doprint(_red[0], out) + if _red_code.startswith("// Not supported in C:"): + eqn = ("eqn_" + str(index), _red_code) + else: + eqn = ("eqn_" + str(index), _temp_code + "\n" + _red_code) + else: + eqn = ("eqn_" + str(index), printer.doprint(fn, out)) + else: + eqn = ("eqn_" + str(index), printer.doprint(fn, out)) + + if eqn[1].startswith("// Not supported in C:"): + spliteqn = eqn[1].split("\n") + raise RuntimeError( + f"Error encountered generating JIT extension:\n" + f"{spliteqn[0]}\n" + f"{spliteqn[1]}\n" + f"This is usually because code generation for a Sympy function (or its derivative) is not supported.\n" + f"Please contact the developers." + f"---" + f"The ID of the JIT component that failed is {index}" + f"The decription of the JIT component that failed:\n {fn}" + ) + eqns.append(eqn) + + return eqns + + class _Unspellable(Exception): """A leaf the kernel has no C for.""" @@ -1107,6 +1723,7 @@ def generate_c_source( verbose: Optional[bool] = False, debug: Optional[bool] = False, debug_name=None, + route="graph", ): """Generate the setup.py / C header / Cython wrapper for a JIT bundle. @@ -1125,13 +1742,15 @@ def generate_c_source( callbacks : JITCallbackSet primary_field_list : list Variables that map to PETSc primary variable arrays (``petsc_u[]``). - lowered_fns : list of sympy.Matrix - The callbacks lowered onto the shared graph (``_jit_graph.lower_callbacks``), - one per entry of ``callbacks.flat()``. Each kernel is emitted as one C - temporary per distinct computation, then its outputs. + lowered_fns : list of sympy.Matrix or None + The graph route: the callbacks lowered onto the shared graph + (``_jit_graph.lower_callbacks``), one per entry of ``callbacks.flat()``; + each kernel is emitted as one C temporary per distinct computation, then + its outputs. ``None`` on the expanded route. constants_subs_map : dict - Mapping from UWexpression to its ``_JITConstant`` placeholder - (``_manifest_of(lowered_fns)``). + Mapping from UWexpression to its ``_JITConstant`` placeholder. + route : {"graph", "expanded"} + Which JIT route generates the kernels. Returns ------- @@ -1148,10 +1767,6 @@ def generate_c_source( count_residual_sig, count_bc_sig, count_jacobian_sig, \ count_bd_residual_sig, count_bd_jacobian_sig = callbacks.counts - # The C a kernel reads for each leaf of the graph: an explicit map, built for - # this compile, instead of C names patched onto the field classes (#823). - spellings = _leaf_spellings(mesh, primary_field_list) - # Create a custom functions replacement dictionary. # Note that this dictionary is really just to appease Sympy, # and the actual implementation is printed directly into the @@ -1219,42 +1834,12 @@ def _handle_UnevaluatedExpr(expr): underworld3._libdirs.clear() underworld3._libfiles.clear() - eqns = [] - for index, fn_original in enumerate(fns): - unspellable = [] - - def spell(leaf): - try: - return _spell_leaf(leaf, spellings, constants_subs_map) - except _Unspellable: - unspellable.append(leaf) - return f"?{sympy.srepr(leaf)}" - - temporaries, fn = _jit_graph.emit(lowered_fns[index], spell) - if unspellable: - raise RuntimeError(_unconvertible_message( - _stable_sorted(set(unspellable)), index, fn_original)) - - if verbose: - # the kernel as mathematics (named quantities as nodes); the C is in the header - print("Processing JIT {:4d} / {}".format(index, lowered_fns[index])) - - out = sympy.MatrixSymbol("out", *fn.shape) - eqn = ("eqn_" + str(index), _print_kernel(printer, temporaries, fn, out)) - - if eqn[1].startswith("// Not supported in C:"): - spliteqn = eqn[1].split("\n") - raise RuntimeError( - f"Error encountered generating JIT extension:\n" - f"{spliteqn[0]}\n" - f"{spliteqn[1]}\n" - f"This is usually because code generation for a Sympy function (or its derivative) is not supported.\n" - f"Please contact the developers." - f"---" - f"The ID of the JIT component that failed is {index}" - f"The decription of the JIT component that failed:\n {fn}" - ) - eqns.append(eqn) + if route == "graph": + eqns = _graph_equations(fns, lowered_fns, constants_subs_map, mesh, + primary_field_list, printer, verbose) + else: + eqns = _expanded_equations(fns, constants_subs_map, mesh, primary_field_list, + printer, verbose) _warn_dirac_deltas_dropped(dropped_deltas, "JIT", collective=True) diff --git a/tests/test_0022_unwrap_memoised_matches_fixed_point.py b/tests/test_0022_unwrap_memoised_matches_fixed_point.py index be03342f5..ef02864cb 100644 --- a/tests/test_0022_unwrap_memoised_matches_fixed_point.py +++ b/tests/test_0022_unwrap_memoised_matches_fixed_point.py @@ -1,5 +1,5 @@ -"""The memoised unwrap and the memoised sqrt guard (#823) give exactly what the -algorithms they replace gave. +"""The memoised unwrap, the memoised sqrt guard and the identity-walk symbol scan (#823) +give exactly what the algorithms they replace gave. The old ``unwrap_expression`` iterated ``subs`` passes to a fixed point, and in ``symbolic_keep_constants`` mode ran ``_is_truly_constant`` (itself a full unwrap) for every @@ -102,6 +102,57 @@ def test_memoised_unwrap_is_the_fixed_point(laws, mode): def test_the_jacobian_sqrt_guard_matches_replace(laws): + """The expanded JIT route's Newton source: its guard, a memoised rebuild, against + SymPy's own replace.""" + from underworld3.cython.generic_solvers import _jacobian_unwrap + + eps2 = sympy.Float(1.0e-36) + + def old_guard(e): + return e.replace( + lambda n: (n.is_Pow and n.exp.is_Rational and n.exp.q == 2 + and n.args[0].free_symbols), + lambda n: sympy.Pow(n.args[0] + eps2, n.exp)) + + for name, expr in laws: + new = _jacobian_unwrap(expr, route="expanded") + old = _unwrap_each(expr, lambda e: old_guard( + _fixed_point_unwrap(e, "symbolic_keep_constants"))) + assert sympy.srepr(new) == sympy.srepr(old), name + + +def test_the_identity_walk_finds_the_same_symbols(laws): + from underworld3.utilities._jitextension import _unique_symbols + + for name, expr in laws: + for e in (expr, ex.unwrap_expression(expr, mode="nondimensional") + if not isinstance(expr, sympy.MatrixBase) + else expr.applyfunc(lambda a: ex.unwrap_expression(a, mode="nondimensional"))): + assert _unique_symbols(e) == set(e.atoms(sympy.Symbol)), name + + +def test_the_shared_xreplace_matches_xreplace(laws): + """The constants[] substitution in the JIT lowering visits each node object once; + it must build exactly what xreplace builds, the generated C and the cache key + both depend on it.""" + from underworld3.utilities._jitextension import _xreplace_shared, _unique_symbols + + swapped = 0 + for name, expr in laws: + unwrapped = _unwrap_each(expr, lambda e: ex.unwrap_expression( + e, mode="symbolic_keep_constants")) + held = sorted((a for a in _unique_symbols(unwrapped) + if isinstance(a, ex.UWexpression)), key=str) + if not held: + continue # k(u) = 1 + u**2 has no parameter to swap + rule = {a: sympy.Symbol(f"slot_{k}") for k, a in enumerate(held)} + assert sympy.srepr(_xreplace_shared(unwrapped, rule)) == \ + sympy.srepr(unwrapped.xreplace(rule)), name + swapped += 1 + assert swapped >= 7 # the VP (three laws, viscosity and flux) and VEP + + +def test_the_graph_guard_matches_replace(laws): """The guard every Newton node body gets (``_jit_graph.guard_half_integer_powers``), a memoised rebuild, against SymPy's own ``replace`` with the same rule.""" from underworld3.utilities._jit_graph import guard_half_integer_powers diff --git a/tests/test_0024_jit_graph_lowering.py b/tests/test_0024_jit_graph_lowering.py index 5146c4193..d632c777a 100644 --- a/tests/test_0024_jit_graph_lowering.py +++ b/tests/test_0024_jit_graph_lowering.py @@ -192,7 +192,7 @@ def test_the_guarded_lowering_is_the_guarded_tree(box): rest = {L[i, j] for i in range(2) for j in range(2)} for flux, tree_escapes_at_rest in ((viscoplastic, False), (power_law, True)): tree = flux.applyfunc(_guarded_tree) - graph = _jacobian_unwrap(flux) + graph = _jacobian_unwrap(flux, route="graph") assert graph.atoms(jg._KernelNode), "the Newton source was not lowered" escaped = _same_guarded_values(tree, graph, L, rest) assert escaped == tree_escapes_at_rest, escaped @@ -234,6 +234,7 @@ def keep(*a, **k): return modname, codeguys, diag monkeypatch.setattr(jx, "generate_c_source", keep) + solver.jit_route = "graph" solver.is_setup = False solver._setup_pointwise_functions() return seen["h"] @@ -289,8 +290,9 @@ def test_a_law_with_no_named_quantity_has_no_temporaries(monkeypatch): def test_the_manifest_is_the_constants_the_kernels_read(box): """Every constant atom a lowered kernel reads, at any depth, has a slot, ordered by - name; a constant inside a constant is folded into its holder's slot.""" - from underworld3.utilities._jitextension import _extract_constants + name; a constant inside a constant is folded into its holder's slot. The expanded + route's scan gives the same manifest.""" + from underworld3.utilities._jitextension import _extract_constants, _manifest_of mesh, T, v = box u = T.sym[0] @@ -301,9 +303,12 @@ def test_the_manifest_is_the_constants_the_kernels_read(box): a = uw.expression(r"a_{0024h}", c1 * u + c4 * u ** 2, "inner") b = uw.expression(r"b_{0024h}", a / (c2 + a ** 2), "outer") fns = (b, sympy.Matrix([[b * u, a]])) - manifest, placeholders = _extract_constants(fns, mesh) - assert [e for _, e in manifest] == [c1, c2, c4] + graph, placeholders = _manifest_of(jg.lower_callbacks(fns, mesh)) + assert [e for _, e in graph] == [c1, c2, c4] assert len(set(placeholders.values())) == 3 + # the expanded route's scan finds the same slots + expanded, _ = _extract_constants(fns, mesh) + assert [e for _, e in expanded] == [c1, c2, c4] def test_a_repeated_condition_stays_a_condition(box): diff --git a/tests/test_0026_jit_route_switch.py b/tests/test_0026_jit_route_switch.py new file mode 100644 index 000000000..228897ccc --- /dev/null +++ b/tests/test_0026_jit_route_switch.py @@ -0,0 +1,124 @@ +"""The two JIT routes, and the switch between them (#823, tier 2). + +``"graph"`` (the default) compiles each named quantity once, as one C temporary; +``"expanded"`` is the JIT before tier 2, which expands every named quantity into one +expression. The expanded route is kept as a fallback and a reference: if a model +misbehaves on one route and not the other, the JIT is at fault; if on both, the model +is. These tests hold the switch to its contract and the two routes to the same +answers, so that the fallback cannot rot unnoticed. +""" +import numpy as np +import pytest +import sympy + +import underworld3 as uw + +pytestmark = [pytest.mark.level_2, pytest.mark.tier_b] + + +@pytest.fixture(autouse=True) +def default_route(monkeypatch): + monkeypatch.delenv("UW_JIT_ROUTE", raising=False) + yield + uw.use_jit_route(None) + + +def test_the_route_setting(monkeypatch): + from underworld3.utilities._jitextension import resolve_jit_route + + assert uw.jit_route() == "graph" + monkeypatch.setenv("UW_JIT_ROUTE", "expanded") + assert uw.jit_route() == "expanded" + uw.use_jit_route("graph") # the call wins over the environment + assert uw.jit_route() == "graph" + assert resolve_jit_route("expanded") == "expanded" # an explicit route wins + uw.use_jit_route(None) + assert uw.jit_route() == "expanded" + with pytest.raises(ValueError, match="JIT route"): + uw.use_jit_route("tree") + monkeypatch.setenv("UW_JIT_ROUTE", "nonsense") + with pytest.raises(ValueError, match="JIT route"): + uw.jit_route() + + +def _viscoplastic_box(name, tangent): + mesh = uw.meshing.UnstructuredSimplexBox( + minCoords=(0.0, 0.0), maxCoords=(1.0, 1.0), cellSize=0.25) + v = uw.discretisation.MeshVariable("V" + name, mesh, 2, degree=2) + p = uw.discretisation.MeshVariable("P" + name, mesh, 1, degree=1) + stokes = uw.systems.Stokes(mesh, velocityField=v, pressureField=p) + stokes.constitutive_model = uw.constitutive_models.ViscoPlasticFlowModel + P = stokes.constitutive_model.Parameters + P.shear_viscosity_0 = 1.0 + P.yield_stress = (uw.expression(r"C_{" + name + "}", 0.3) + + uw.expression(r"\mu_{" + name + "}", 0.2) * p.sym[0]) + P.yield_stress_min = uw.expression(r"\tau_{" + name + "}", 0.01) + stokes.add_dirichlet_bc((1.0, 0.0), "Top") + stokes.add_dirichlet_bc((0.0, 0.0), "Bottom") + stokes.bodyforce = sympy.Matrix([0, -1]) + stokes.consistent_jacobian = tangent + stokes.tolerance = 1.0e-9 + stokes.petsc_options["snes_max_it"] = 200 # Picard converges linearly here + return stokes, v, p + + +@pytest.mark.parametrize("tangent", [True, False]) +def test_both_routes_solve_a_viscoplastic_box_alike(tangent): + """The same solve, Newton or Picard, on each route: the same nonlinear and linear + iteration counts and the same solution, to round-off.""" + uw.reset_default_model() + results = {} + for route in ("graph", "expanded"): + stokes, v, p = _viscoplastic_box(f"0026{route[0]}{int(tangent)}", tangent) + stokes.jit_route = route + stokes.solve(zero_init_guess=True) + report = stokes.solve_report + assert stokes.snes.getConvergedReason() > 0, (route, report.reason_str) + results[route] = (report.nl_its, report.ksp_its, + np.array(v.array), np.array(p.array)) + g, e = results["graph"], results["expanded"] + assert g[:2] == e[:2], (g[:2], e[:2]) + assert np.max(np.abs(g[2] - e[2])) <= 1.0e-9 * np.max(np.abs(e[2])) + assert np.max(np.abs(g[3] - e[3])) <= 1.0e-9 * np.max(np.abs(e[3])) + + +def test_switching_one_solver_rebuilds_it_and_keeps_its_answer(monkeypatch): + """A solver's jit_route overrides the process default; changing it rebuilds that + solver's kernels at the next solve (a new module) and nothing else.""" + import underworld3.utilities._jitextension as jx + + headers = [] + generate = jx.generate_c_source + + def keep(*args, **kwargs): + modname, codeguys, diag = generate(*args, **kwargs) + headers.append(dict(codeguys)["cy_ext.h"]) + return modname, codeguys, diag + + monkeypatch.setattr(jx, "generate_c_source", keep) + uw.reset_default_model() + stokes, v, p = _viscoplastic_box("0026s", True) + stokes.solve(zero_init_guess=True) + first_key, first_v = stokes._current_jit_cache_key, np.array(v.array) + assert "uwt_" in headers[-1] # the default is the graph + + stokes.jit_route = "expanded" + stokes.solve(zero_init_guess=True) + assert stokes._current_jit_cache_key != first_key + assert "uwt_" not in headers[-1] # one expanded expression per output + assert np.max(np.abs(np.array(v.array) - first_v)) <= 1.0e-9 * np.max(np.abs(first_v)) + + stokes.jit_route = None # back to the default + stokes.solve(zero_init_guess=True) + assert stokes._current_jit_cache_key == first_key + with pytest.raises(ValueError, match="JIT route"): + stokes.jit_route = "tree" + + +def test_the_process_default_reaches_new_solvers(): + uw.reset_default_model() + uw.use_jit_route("expanded") + stokes, v, p = _viscoplastic_box("0026d", True) + assert stokes.jit_route is None and stokes._jit_route_in_use() == "expanded" + stokes.solve(zero_init_guess=True) + assert stokes.snes.getConvergedReason() > 0 diff --git a/tests/test_0103_jit_rampable_constants.py b/tests/test_0103_jit_rampable_constants.py index aec70a572..bd9dd67a7 100644 --- a/tests/test_0103_jit_rampable_constants.py +++ b/tests/test_0103_jit_rampable_constants.py @@ -146,10 +146,11 @@ def test_same_named_placeholders_order_deterministically(): half of launches. This test is the deterministic proxy: distinct sort keys, and a sum that canonicalises the same however it is written. - Since #823 the placeholders no longer reach the generated C: the graph writes - each constant leaf as ``constants[i]`` directly, in canonical order, and - ``test_0105`` holds the C to one hash under every seed. This test keeps the - placeholders' own semantics, which code that substitutes them relies on. + On the graph route (the default since #823) the placeholders do not reach the + generated C: each constant leaf is written as ``constants[i]`` directly, in + canonical order. On the expanded route, the JIT before #823 kept as a fallback, + they do, and this test pins what that route relies on. ``test_0105`` holds the C + of both to one hash under every seed. """ import sympy from underworld3.utilities._jitextension import _JITConstant diff --git a/tests/test_0104_constant_slot_still_constant.py b/tests/test_0104_constant_slot_still_constant.py index 4d1a7f24b..8a0f1d86e 100644 --- a/tests/test_0104_constant_slot_still_constant.py +++ b/tests/test_0104_constant_slot_still_constant.py @@ -7,7 +7,8 @@ constants[] slot, and ramping ``m`` left the kernel reading a scalar where the law depends on position again. The JIT now decides constancy by structure (#823, tier 2): the expression reads a field, so it is compiled as a quantity reading ``T`` and the -``m`` slot, and ``m`` ramps without a recompile. +``m`` slot, and ``m`` ramps without a recompile. The expanded route, the JIT before +tier 2 kept as a fallback, still decides by value; there the collapsed slot raises. A slot can still stop being constant: a constant atom whose content is replaced by one that reads a field, without a rebuild. Packing a zero into it would hand the solve a @@ -21,7 +22,7 @@ pytestmark = [pytest.mark.level_1, pytest.mark.tier_a] -def _build(initial): +def _build(initial, route="graph"): import underworld3 as uw uw.reset_default_model() @@ -41,6 +42,7 @@ def _build(initial): poisson.add_dirichlet_bc(0.0, "Top") poisson.add_dirichlet_bc(0.0, "Bottom") poisson.petsc_options.delValue("ksp_monitor") + poisson.jit_route = route return uw, poisson, T, m @@ -126,3 +128,15 @@ def test_an_expression_that_stays_constant_is_unaffected(): poisson.solve(zero_init_guess=True) second = float(np.asarray(T.data)[:, 0].mean()) assert second == pytest.approx(first / 2.0, rel=1e-6) + + +def test_on_the_expanded_route_a_collapsed_slot_raises_when_it_stops_being_constant(): + """The expanded route keeps the JIT's behaviour before #823 tier 2: compiled while + the whole expression is the number 2, the diffusivity banks as one scalar slot, + and ramping m makes it depend on T again. Packing must raise, not pass a zero.""" + uw, poisson, T, m = _build(0.0, route="expanded") + poisson.solve(zero_init_guess=True) + assert len(poisson.constants_manifest) == 1 + m.sym = sympy.sympify(0.5) + with pytest.raises(RuntimeError, match="no longer.*reduces to a number"): + poisson.solve(zero_init_guess=True) diff --git a/tests/test_0105_jit_source_seed_independence.py b/tests/test_0105_jit_source_seed_independence.py index abd410da6..f9e4658a7 100644 --- a/tests/test_0105_jit_source_seed_independence.py +++ b/tests/test_0105_jit_source_seed_independence.py @@ -91,8 +91,11 @@ def _module_hash(seed, cache_dir, tmp_path): return line[0].removeprefix("MODULES:") +@pytest.mark.parametrize("route", ["graph", "expanded"]) @pytest.mark.parametrize("seeds", [(0, 1, 2)]) -def test_the_generated_module_is_the_same_under_every_hash_seed(seeds, tmp_path): +def test_the_generated_module_is_the_same_under_every_hash_seed(seeds, route, tmp_path, + monkeypatch): + monkeypatch.setenv("UW_JIT_ROUTE", route) # the children read it hashes = {} for seed in seeds: cache = tmp_path / f"cache_{seed}" # a cold cache per seed @@ -115,6 +118,8 @@ def test_the_generated_module_is_the_same_under_every_hash_seed(seeds, tmp_path) import sympy, underworld3 as uw import underworld3.utilities._jitextension as jx + uw.use_jit_route("graph") + temporaries = [] generate = jx.generate_c_source From d96aa9e459d933d5fe5c14862d819485c951a786 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Fri, 9 Oct 2026 21:00:46 -1000 Subject: [PATCH 13/14] Allowlist as on development: the restored expanded route keeps its except-pass (#823) The expanded JIT route is development's code, unedited, including its one except Exception: pass; the allowlist entry removed with that code comes back. Underworld development team with AI support from Claude Code --- scripts/deprecated_pattern_allowlist.txt | 1 + 1 file changed, 1 insertion(+) diff --git a/scripts/deprecated_pattern_allowlist.txt b/scripts/deprecated_pattern_allowlist.txt index 42f2c6235..fa55e67f7 100644 --- a/scripts/deprecated_pattern_allowlist.txt +++ b/scripts/deprecated_pattern_allowlist.txt @@ -55,6 +55,7 @@ src/underworld3/units.py:except-pass src/underworld3/utilities/_api_tools.py:except-pass src/underworld3/utilities/_interrupt.py:except-pass src/underworld3/utilities/_jit_cache.py:except-pass +src/underworld3/utilities/_jitextension.py:except-pass src/underworld3/utilities/_params.py:except-pass src/underworld3/utilities/diagnostics.py:except-pass src/underworld3/utilities/mathematical_mixin.py:except-pass From 5455449c1081634a49fdea3104a51058f06d7e44 Mon Sep 17 00:00:00 2001 From: lmoresi Date: Sat, 10 Oct 2026 08:36:21 -1000 Subject: [PATCH 14/14] Name each C temporary after the quantity it computes (#823, tier 2) Louis, 2026-10-11: the generated C carries each temporary's quantity as a comment, const double uwt_1 = ...; /* \dot\varepsilon_{II} */, so a kernel can be read by eye. The names enter the cache key; renaming a quantity recompiles. A name that holds the comment terminator is written so it cannot end the comment (test_0024::test_a_name_cannot_close_its_comment). Underworld development team with AI support from Claude Code --- .../design/jit-shared-graph-codegen.md | 18 ++++++++++-------- src/underworld3/utilities/_jit_graph.py | 14 ++++++++------ src/underworld3/utilities/_jitextension.py | 11 +++++++---- tests/test_0024_jit_graph_lowering.py | 17 +++++++++++++++++ 4 files changed, 42 insertions(+), 18 deletions(-) diff --git a/docs/developer/design/jit-shared-graph-codegen.md b/docs/developer/design/jit-shared-graph-codegen.md index c75c34b0c..dd1cdac44 100644 --- a/docs/developer/design/jit-shared-graph-codegen.md +++ b/docs/developer/design/jit-shared-graph-codegen.md @@ -271,9 +271,9 @@ Applications with equal keys compute the same C and share one temporary. The temporaries are written so that each follows the ones it uses, ties broken by key: ```c -const double uwt_0 = ...; /* one line per distinct computation, evaluated once */ -const double uwt_1 = ...; -out[0] = ...; /* outputs in terms of uwt_0, uwt_1, ... and the leaves */ +const double uwt_0 = ...; /* \dot\varepsilon_{II} */ (one line per distinct computation, +const double uwt_1 = ...; /* \eta_{\mathrm{eff}} */ named by the quantity it computes) +out[0] = ...; (outputs in terms of uwt_0, uwt_1, ... and the leaves) ``` The generated source is therefore a function of the mathematics and of the kernel's data @@ -732,10 +732,12 @@ the next begins. the expanded route keeps its own tests (`test_0022`, `test_0103`, `test_0104`, `test_0105`). Removing it is a later decision, once the graph has run in production. +- **2026-10-11: each C temporary carries its quantity's name as a comment** + (`const double uwt_3 = ...; /* \eta_{\mathrm{eff}} */`), so a kernel can be read + by eye. The names enter the cache key: renaming a quantity recompiles its kernels. + ## Questions for the maintainer -1. Whether the generated C should carry each temporary's display name as a comment. - It makes kernels readable; it also puts display names into the cache key, so a - `rename()` recompiles. -2. Whether the adjoint's own unwrap (`_peel_except`, on the adjoint branches) adopts the - nodes when those branches rebase onto this change. +1. Whether the adjoint's own unwrap (`_peel_except`, on the adjoint branches) adopts the + nodes when those branches rebase onto this change. (Louis, 2026-10-11: + probably.) diff --git a/src/underworld3/utilities/_jit_graph.py b/src/underworld3/utilities/_jit_graph.py index bb619975a..585246894 100644 --- a/src/underworld3/utilities/_jit_graph.py +++ b/src/underworld3/utilities/_jit_graph.py @@ -463,8 +463,9 @@ def emission_order(outputs, spell): their own keys. Applications with equal keys compute the same C and share one temporary; temporaries follow the ones they use, ties broken by key. - Returns ``(order, key)``: ``order`` is a list of ``(key, body)``; ``key`` maps each - node application reached to its key. + Returns ``(order, key)``: ``order`` is a list of ``(key, body, label)``, the label + being the name of the first quantity reached with that key; ``key`` maps each node + application reached to its key. """ key = {} spelt = {} @@ -498,7 +499,7 @@ def visit(app): body = body_of(app) for child in sorted(body.atoms(_KernelNode), key=key_of): visit(child) - order.append((k, body)) + order.append((k, body, app.func._label)) for out in outputs: for child in sorted(sympy.sympify(out).atoms(_KernelNode), key=key_of): @@ -510,7 +511,8 @@ def emit(fn, spell): """``fn`` (a lowered Matrix) as temporaries and outputs, ready to print. Returns ``(temporaries, outputs)``: ``temporaries`` is a list of - ``(_Temporary, body)`` in the order to write them, and ``outputs`` is ``fn`` with + ``(_Temporary, body, label)`` in the order to write them, the label being the name + of the quantity the temporary computes, and ``outputs`` is ``fn`` with each node application replaced by its temporary. In both, every leaf is replaced by a ``_CName`` holding ``spell(leaf)``. """ @@ -529,8 +531,8 @@ def rule_for(e, temp_of): return rule temp_of, temporaries = {}, [] - for i, (k, body) in enumerate(order): + for i, (k, body, label) in enumerate(order): t = _Temporary(i) - temporaries.append((t, body.xreplace(rule_for(body, temp_of)))) + temporaries.append((t, body.xreplace(rule_for(body, temp_of)), label)) temp_of[k] = t return temporaries, fn.xreplace(rule_for(fn, temp_of)) diff --git a/src/underworld3/utilities/_jitextension.py b/src/underworld3/utilities/_jitextension.py index 8b88f85ba..503eef789 100644 --- a/src/underworld3/utilities/_jitextension.py +++ b/src/underworld3/utilities/_jitextension.py @@ -1633,8 +1633,9 @@ def _unconvertible_message(symbols, index, fn_original): def _print_kernel(printer, temporaries, outputs, out): - """The C body of one kernel: a ``const double`` per temporary, in order, then - the outputs. The printer declares a number symbol it reads (``EulerGamma``, + """The C body of one kernel: a ``const double`` per temporary, in order, each + with the name of the quantity it computes as a comment (so a rename recompiles), + then the outputs. The printer declares a number symbol it reads (``EulerGamma``, ``Catalan``) before the code that reads it; each declaration is written once, at the top. A SymPy function the printer cannot write raises ``PrintMethodNotImplementedError`` (SymPy 1.14), or, in a SymPy that returns @@ -1650,11 +1651,13 @@ def without_declarations(code): rest = rest[1:] return "\n".join(rest) - for t, body in temporaries: + for t, body, label in temporaries: code = printer.doprint(body) if code.startswith("// Not supported in C:"): return code - lines.append(f"const double {t._ccodestr} = {without_declarations(code)};") + name = " ".join(str(label).replace("*/", "* /").split()) + lines.append(f"const double {t._ccodestr} = {without_declarations(code)};" + f" /* {name} */") lines.append(without_declarations(printer.doprint(outputs, out))) return "\n".join(declarations + lines) diff --git a/tests/test_0024_jit_graph_lowering.py b/tests/test_0024_jit_graph_lowering.py index d632c777a..8da9edac3 100644 --- a/tests/test_0024_jit_graph_lowering.py +++ b/tests/test_0024_jit_graph_lowering.py @@ -270,6 +270,8 @@ def test_the_emitted_source_does_not_depend_on_what_came_before(monkeypatch): second = _header(_poisson("0024f", mesh2), monkeypatch) assert first == second assert "uwt_0" in first, "the kernel has no temporaries: nothing was lowered" + # each temporary carries the name of the quantity it computes + assert "/* g_{0024f} */" in first def test_a_law_with_no_named_quantity_has_no_temporaries(monkeypatch): @@ -470,3 +472,18 @@ def test_a_number_symbol_in_a_temporary_is_declared_once(monkeypatch): assert "uwt_0" in header for body in header.split("\nvoid ")[1:]: assert body.count("const double EulerGamma =") <= 1 + + +def test_a_name_cannot_close_its_comment(box): + """Each temporary's comment is the name of its quantity; a name holding the + comment terminator is written so that it cannot end the comment early.""" + from sympy.printing.c import c_code_printers + from underworld3.utilities._jitextension import _print_kernel + + mesh, T, v = box + t = jg._Temporary(0) + printer = c_code_printers["c99"]() + code = _print_kernel(printer, [(t, sympy.Float(2.0), "a */ b\nc")], + sympy.Matrix([[t]]), sympy.MatrixSymbol("out", 1, 1)) + first = code.splitlines()[0] + assert first.endswith("/* a * / b c */") and first.count("*/") == 1, first