Skip to content

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

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

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

Conversation

@finsberg

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).

  • Tag a Function as identity-sensitive exactly where that fact becomes known: DirichletBC.__init__ sets a new Function._ad_bc_backing attribute.
  • FunctionAssignBlock.recompute_component now isolates a fresh snapshot (_ad_new_like()) for untagged Function targets, and keeps the original in-place mutation only for tagged (BC-backing) targets.
  • No change to assign()'s or dirichletbc()'s public signature. The PR Time-distributed control + blocked rewrite. #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 in .scratch/dirichlet-bc-recompute-identity/spec.md (knowledge repo).

Known, documented, out-of-scope residual: a Function that is both BC-backing and reassigned every tape timestep still hits the original aliasing bug under a schedule — the two needs are genuinely incompatible with no further information. Recorded in the spec's Out of Scope section and in the _ad_bc_backing docstring; a future fix would need to warn when a tagged Function is used as an assign target on a schedule-enabled tape.

Test plan

  • tests/test_assign.py::test_recompute_does_not_alias_state_across_timesteps (new) — reproduces the 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 now 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), serially and under mpirun -n 2
  • ruff check . and mypy src/dolfinx_adjoint clean

Implemented via subagent-driven development (3 tasks, each independently reviewed) plus a final whole-branch review; both this PR's description and the spec capture every finding and ruling made along the way.

🤖 Generated with Claude Code

finsberg and others added 4 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.
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