Fix checkpointed recompute aliasing for state-transfer Function targets - #76
Open
finsberg wants to merge 5 commits into
Open
Fix checkpointed recompute aliasing for state-transfer Function targets#76finsberg wants to merge 5 commits into
finsberg wants to merge 5 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.
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>
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).This PR went through two designs before landing on the current one — recorded here since both are informative:
Functionas identity-sensitive via a newFunction._ad_bc_backingattribute, set byDirichletBC.__init__, and branchFunctionAssignBlock.recompute_componenton it. Worked, but left an implicit, easy-to-forget special case: a boolean flag that silently changes what a genericassign()call does, set from a completely different call site. It also had a documented residual gap — aFunctionthat was both BC-backing and reassigned every timestep still hit the original aliasing bug.FunctionAssignBlockisolates unconditionally for everyFunctiontarget — no tag, no branch.LinearProblemBlock/NonlinearProblemBlockinstead track each BC's backing Function (bc.g) as an explicit tape dependency, the same way every other form coefficient already is, and a newsync_bc_valueshelper (blocks/dirichletbc.py) refreshesbc.g's live array from that pinned dependency'ssaved_outputright 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 offbc.g.block_variable.saved_outputdirectly, was also tried and also empirically wrong —.block_variablealways points atbc.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 ordirichletbc()'s public signature. The PR #75RuntimeErrorguard is untouched. One unrelated cleanup: removed a now-stalexfail(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, 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 genuineRevolveschedule, 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_replaynever 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 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 pass (test_gradient_matches_uncheckpointed,test_taylor_test_under_checkpointing,test_disk_gradient_matches_uncheckpointed,test_disk_taylor_test)ruff check .andmypy src/dolfinx_adjointclean🤖 Generated with Claude Code