Check each lane of a batched load before recomputing it - #3256
Open
timesselens wants to merge 1 commit into
Open
timesselens wants to merge 1 commit into
timesselens wants to merge 1 commit into
Conversation
With a vector width above 1, the reverse pass can reload a shadow pointer after the location it was loaded from has been overwritten, and use the new pointer instead of the one the forward pass loaded. The batched gradient then differs from the width-1 gradient, which gives the expected value. For example, a loop loads a pointer, sometimes replaces it and loads it again, and stores through the phi of the two loads. With width 2 and seeds 1 and 2, dx comes out as 2.431 and 4.862 instead of 5.324 and 10.648. invertPointerM builds the shadow of a load as one load per lane, combined into an aggregate, and registers only the aggregate in invertedPointers. When unwrapM rebuilds a phi of such shadows in the reverse pass, it asks legalRecompute about each lane's load. A lane's load is not an original instruction and hasUninverted does not know it, so legalRecompute reaches its final `return true` and the lane is reloaded. With width 1, the overwrite check applies through hasUninverted. Record each lane's load in unwrappedLoads, mapped to the original load, so that legalRecompute checks it like the original load. The new test ReverseModeVector/overwrittenloadphi.ll fails without this change, both its CHECK lines and its run. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
This branch has not been deployed
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
With a vector width above 1, the reverse pass can reload one lane of a shadow pointer after the location it came from has been overwritten. The batched gradient then differs from the width-1 gradient, which gives the expected value.
invertPointerMseems to register only the aggregate of the per-lane loads, solegalRecomputecannot trace a lane back to its load. This PR proposes recording each lane's load, so that it gets the same overwrite check as the original load.Before / after
A loop loads a pointer, sometimes replaces it and loads it again, and stores through the phi of the two loads:
Likely cause
GradientUtils::invertPointerMbuilds the shadow withapplyChainRule: one load per lane, combined into an aggregate. Only the aggregate is registered ininvertedPointers.unwrapMrebuilds it lane by lane and askslegalRecomputeabout each lane's load.hasUninverted, solegalRecomputereaches its final// TODO mark all the explicitly legal nodes (caches, etc)and returns true. The lane is then reloaded, althoughb->vwas overwritten after the original load. With width 1 the shadow load itself is registered, and the overwrite check applies.Proposed change
In
invertPointerM, for width above 1, record each lane's load inunwrappedLoads, mapped to the original load:legalRecomputealready followsunwrappedLoadsto the load a value derives from, as it does for the loadsunwrapMemits. The lane is then found not recomputable and is cached in the forward pass.restoreCacheonly re-unwraps entries that map to a new-function value, so it skips these (they map to an original load).Tests
New
test/Enzyme/ReverseModeVector/overwrittenloadphi.ll(the function above in IR, width 2, opaque pointers): a FileCheck that each lane is cached and the phi is rebuilt from those caches, and anllirun that printsdx: 5.324000 10.648000(mainprintsdx: 2.431000 4.862000for this IR).Open questions
No specific questions on this one, but this may miss a case where the extra caching is unwanted. The cost is caching where the value is needed: in the test the forward pass now caches each lane of the two shadow loads, where it already cached the aggregate of the first one.
Related
legalRecomputeasks about the value they stand for.invertPointerM, the shadow of a constant), Permit an unwrap's operand lookup to fail, and rebuild uncacheable values #3114 (unwrap lookups that may fail), Bound lookupM recompute recursion to avoid compile-time blowup #3005 (boundinglookupM's recompute recursion).Verification
lit (
check-enzyme,check-typeanalysis,check-activityanalysis), macOS aarch64, LLVM fromLLVM_full_jll, libEnzyme with assertions:mainOn
mainb85d5011. No existing test changes its result. The new test also fails onmainand passes with this PR on LLVM 18.1.7.The C reproducer at
-O0gives{2.431, 4.862}onmain; at-O2the loop is restructured andmaingives the expected value.This came up through Julia, where
BatchDuplicatedgives width 2 (Enzyme.jl 0.13.205, Enzyme_jll 0.0.294):With this change the batched result matches
Duplicatedfor n = 1 to 6; without it, it differs from n = 4 on. Built on v0.0.293, Enzyme.jl's CPU suite gave the same result as with the stock library (218244 passed, 61 broken, 0 failed).Only loads are recorded: the other per-lane shadow instructions (GEPs, casts) don't read memory, so recomputing them should be legal anyway.