Fix checkpointed recompute aliasing for state-transfer Function targets - #76
Open
finsberg wants to merge 4 commits into
Open
Fix checkpointed recompute aliasing for state-transfer Function targets#76finsberg wants to merge 4 commits into
finsberg wants to merge 4 commits into
Conversation
…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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
FunctionAssignBlock.recompute_componentalways mutatedblock_variable.saved_outputin place and returned the same object on every recompute. That's required for aFunctionbacking a livedolfinx.fem.DirichletBC(the C++ binding reads that object's array directly, not through the tape), but wrong for an ordinary state-transferFunction(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 underRevolve, exact 2.0 with no schedule).Functionas identity-sensitive exactly where that fact becomes known:DirichletBC.__init__sets a newFunction._ad_bc_backingattribute.FunctionAssignBlock.recompute_componentnow isolates a fresh snapshot (_ad_new_like()) for untaggedFunctiontargets, and keeps the original in-place mutation only for tagged (BC-backing) targets.assign()'s ordirichletbc()'s public signature. The PR Time-distributed control + blocked rewrite. #75RuntimeErrorguard is untouched.xfail(strict=True)ontest_snes_time_loop_gradient_is_correct, which started passing due to an earlier, separate fix and was blocking a fully greentest_checkpointing.py.Full design rationale in
.scratch/dirichlet-bc-recompute-identity/spec.md(knowledge repo).Known, documented, out-of-scope residual: a
Functionthat 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_backingdocstring; 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 viaassign()chains under aRevolveschedule, no PDE solvetests/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)mpirun -n 2ruff check .andmypy src/dolfinx_adjointcleanImplemented 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