Skip to content

Fix checkpointed recompute aliasing for state-transfer Function targets - #76

Open
finsberg wants to merge 5 commits into
checkpointingfrom
finsberg/dirichlet-bc-recompute-identity
Open

Fix checkpointed recompute aliasing for state-transfer Function targets#76
finsberg wants to merge 5 commits into
checkpointingfrom
finsberg/dirichlet-bc-recompute-identity

Conversation

@finsberg

@finsberg finsberg commented Aug 28, 2026

Copy link
Copy Markdown
Member

Summary

FunctionAssignBlock.recompute_component always mutated block_variable.saved_output in place and returned the same object on every recompute. That's required for a Function backing a live dolfinx.fem.DirichletBC (the C++ binding reads that object's array directly, not through the tape), but wrong for an ordinary state-transfer Function (e.g. assign(uh, u_prev) in a time-stepping loop): every timestep's recompute silently overwrote the value an earlier timestep's checkpoint relied on, giving a wrong gradient once a checkpoint schedule actually forces replay (confirmed empirically: Taylor rate 0.70 instead of 2.0 under Revolve, exact 2.0 with no schedule).

This PR went through two designs before landing on the current one — recorded here since both are informative:

  1. (shipped first, since replaced) Tag a Function as identity-sensitive via a new Function._ad_bc_backing attribute, set by DirichletBC.__init__, and branch FunctionAssignBlock.recompute_component on it. Worked, but left an implicit, easy-to-forget special case: a boolean flag that silently changes what a generic assign() call does, set from a completely different call site. It also had a documented residual gap — a Function that was both BC-backing and reassigned every timestep still hit the original aliasing bug.
  2. Shipped: FunctionAssignBlock isolates unconditionally for every Function target — no tag, no branch. LinearProblemBlock/NonlinearProblemBlock instead track each BC's backing Function (bc.g) as an explicit tape dependency, the same way every other form coefficient already is, and a new sync_bc_values helper (blocks/dirichletbc.py) refreshes bc.g's live array from that pinned dependency's saved_output right before every solve, including during recompute. This closes the residual gap design 1 had to leave open, and needed no new attribute anywhere.

A first attempt at sync_bc_values, keyed off bc.g.block_variable.saved_output directly, was also tried and also empirically wrong — .block_variable always points at bc.g's most recently created BlockVariable, which after the full tape is recorded is simply the last timestep's, regardless of which point in a replay is being recomputed. Reading from the calling block's own pinned dependency (self.get_dependencies()) instead is what's actually position-aware. Recorded in the spec so nobody rediscovers this by bisection.

No change to assign()'s or dirichletbc()'s public signature. The PR #75 RuntimeError guard is untouched. One unrelated cleanup: removed a now-stale xfail(strict=True) on test_snes_time_loop_gradient_is_correct, which started passing due to an earlier, separate fix and was blocking a fully green test_checkpointing.py.

Full design rationale, including the rejected designs and the empirical evidence for each, in .scratch/dirichlet-bc-recompute-identity/spec.md (knowledge repo).

Test plan

  • tests/test_checkpointing.py::test_bc_gradient_matches_uncheckpointed (new) — a time-dependent Dirichlet BC under a genuine Revolve schedule, comparing gradients (not just forward values) against the unscheduled run. Fails on the previously-shipped tag-based design with a real ~0.6% gradient mismatch (test_time_dependent_bc_replay never enables an actual schedule, so it couldn't catch this); passes exactly on the current design.
  • tests/test_assign.py::test_recompute_does_not_alias_state_across_timesteps — reproduces the original state-transfer aliasing defect directly via assign() chains under a Revolve schedule, no PDE solve
  • tests/test_dirichlet_bc.py::test_time_dependent_bc_replay — stays green throughout (the test that would catch a regression toward "always isolate")
  • tests/test_checkpointing.py — all previously-red tests pass (test_gradient_matches_uncheckpointed, test_taylor_test_under_checkpointing, test_disk_gradient_matches_uncheckpointed, test_disk_taylor_test)
  • Full suite: 86 passed, 1 xfailed (pre-existing, unrelated residual-timestepping defect)
  • ruff check . and mypy src/dolfinx_adjoint clean

🤖 Generated with Claude Code

finsberg and others added 5 commits August 28, 2026 10:25
…argets

FunctionAssignBlock.recompute_component mutated block_variable.saved_output
in place on every recompute. This is required for _ad_bc_backing-tagged
Functions (a live DirichletBC reads that exact object's array via a C++
binding, not through the tape) but silently aliases state for ordinary
Function targets reused across a time loop (e.g. a "previous timestep
value"): once a checkpoint schedule forces genuine recompute, each
timestep's recompute overwrites the value an earlier timestep's checkpoint
was relying on.

Return an isolated snapshot (via Function._ad_new_like()) for any Function
target that is not backing a live DirichletBC, and keep the in-place update
for DirichletBC-backing Functions and non-Function outputs.

Also restores the working tape at the end of the new
test_recompute_does_not_alias_state_across_timesteps test: a tape that has
had checkpointing enabled keeps eagerly checkpointing outputs even after
clear_tape() (per the isolated_tape fixture in test_checkpointing.py), so
leaving the Revolve-enabled tape as the global working tape broke
test_time_dependent_bc_replay when the test files ran in the same session.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
…correct

The test now passes due to an unrelated SNES coefficient-replacement fix
that landed via a merge. The underlying defect is fixed, so retire the xfail marker.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
- Narrow the isinstance check in FunctionAssignBlock.recompute_component to
  the overloaded _Function type, matching the type that actually declares
  _ad_bc_backing, and simplify the accompanying comment to drop a vacuous
  "non-Function output" clause.
- Add the missing clear_tape() to test_dirichletbc_tags_its_value_function
  in tests/test_dirichlet_bc.py, matching the file's convention, after the
  final review confirmed its absence leaks a block onto the shared tape.
- Remove an unused Function/interpolate() pair in
  test_recompute_does_not_alias_state_across_timesteps
  (tests/test_assign.py); the test's actual controls come from a separate
  list.
- Document, in the _ad_bc_backing docstring, that tagging trades away
  checkpoint-aliasing safety for BC identity, so a Function needing both is
  unsupported.
FunctionAssignBlock now isolates unconditionally for every Function target,
with no special case. Instead, LinearProblemBlock and NonlinearProblemBlock
track each BC's backing Function (bc.g) as an explicit dependency, the same
way every other form coefficient already is, and sync_bc_values refreshes
bc.g's live array from that pinned dependency's own saved_output right
before each solve -- including during recompute.

An earlier version of this fix (and, before that, a version using
bc.g.block_variable.saved_output directly) both looked plausible but were
empirically wrong: bc.g.block_variable always points at bc.g's most
recently created BlockVariable, which after the tape is fully recorded is
simply the last timestep's, regardless of which point in a replay is being
recomputed. Reading from the calling block's own pinned dependency instead
is what's actually position-aware.

Adds test_bc_gradient_matches_uncheckpointed, which enables a genuine
Revolve schedule (unlike test_time_dependent_bc_replay, which only ever
does a full unscheduled replay) and would have caught this: it fails on
the previously-shipped tag-based version with a real ~0.6% gradient
mismatch, and passes exactly on this one.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant