Skip to content

Check each lane of a batched load before recomputing it - #3256

Open
timesselens wants to merge 1 commit into
EnzymeAD:mainfrom
timesselens:pr/batch-lane-loads
Open

timesselens wants to merge 1 commit into
EnzymeAD:mainfrom
timesselens:pr/batch-lane-loads

Conversation

@timesselens

Copy link
Copy Markdown

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. invertPointerM seems to register only the aggregate of the per-lane loads, so legalRecompute cannot 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:

struct Box { double *v, *spare; };

// Every other iteration moves the data into the spare buffer and swaps the two,
// so b->v changes and v is read again. out = x^n v[0].
__attribute__((noinline)) void f(double *out, struct Box *b, double *x, long n) {
  for (long i = 1; i <= n; i++) {
    double *v = b->v;
    double t = v[0];
    if (i % 2 == 0) {
      double *s = b->spare;
      s[0] = t;
      b->spare = v;
      b->v = s;
      v = b->v;
    }
    v[0] = *x * t;
  }
  *out = b->v[0];
}
// __enzyme_autodiff(f, enzyme_width, 2, &out, &dout1, &dout2, &b, &db1, &db2, &x, &dx1, &dx2, enzyme_const, n)
// with n = 4, x = 1.1, seeds 1 and 2
$ clang -O1 -S -emit-llvm lanes.c -o - | opt -load-pass-plugin=LLVMEnzyme-16.dylib -passes=enzyme -S | lli -
main:    dx = {3.63, 7.26}
this PR: dx = {5.324, 10.648}      (d out/dx = n x^(n-1) = 5.324; width 1 gives 5.324 on main too)

Likely cause

  • For a load, GradientUtils::invertPointerM builds the shadow with applyChainRule: one load per lane, combined into an aggregate. Only the aggregate is registered in invertedPointers.
  • The store through the phi needs the phi's shadow in the reverse pass. unwrapM rebuilds it lane by lane and asks legalRecompute about each lane's load.
  • A lane's load is neither an original instruction nor known to hasUninverted, so legalRecompute reaches its final // TODO mark all the explicitly legal nodes (caches, etc) and returns true. The lane is then reloaded, although b->v was 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 in unwrappedLoads, mapped to the original load:

       li->setSyncScopeID(arg->getSyncScopeID());
+      // With a vector width > 1 the shadow registered for `arg` is the
+      // aggregate of these per-lane loads, so hasUninverted cannot map a lane
+      // back to `arg`. Record it here so that legalRecompute applies the same
+      // overwrite check to each lane as to `arg`.
+      if (getWidth() > 1)
+        unwrappedLoads[li] = arg;
       idx++;
       return li;

legalRecompute already follows unwrappedLoads to the load a value derives from, as it does for the loads unwrapM emits. The lane is then found not recomputable and is cached in the forward pass. restoreCache only 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 an lli run that prints dx: 5.324000 10.648000 (main prints dx: 2.431000 4.862000 for this IR).

llvm-lit -v build/test/Enzyme/ReverseModeVector/overwrittenloadphi.ll

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

Verification
  • lit (check-enzyme, check-typeanalysis, check-activityanalysis), macOS aarch64, LLVM from LLVM_full_jll, libEnzyme with assertions:

    LLVM main this PR
    15.0.7 1184 passed, 11 expectedly failed 1185 passed, 11 expectedly failed
    16.0.6 1184 passed, 11 expectedly failed 1185 passed, 11 expectedly failed

    On main b85d5011. No existing test changes its result. The new test also fails on main and passes with this PR on LLVM 18.1.7.

  • The C reproducer at -O0 gives {2.431, 4.862} on main; at -O2 the loop is restructured and main gives the expected value.

  • This came up through Julia, where BatchDuplicated gives width 2 (Enzyme.jl 0.13.205, Enzyme_jll 0.0.294):

    using Enzyme
    mutable struct Box; v::Vector{Float64}; end
    @noinline replace!(b) = (b.v = copy(b.v); nothing)
    function f!(out, b, x, n)
        for i in 1:n
            v = b.v; t = v[1]
            if i % 2 == 0; replace!(b); v = b.v; end   # sometimes replace b.v, then read it again
            v[1] = x[] * t                              # through the phi of the two loads of b.v
        end
        out[] = b.v[1]; return nothing
    end
    b = Box([1.0]); db1 = Box([0.0]); db2 = Box([0.0]); dx1 = Ref(0.0); dx2 = Ref(0.0)
    autodiff(Reverse, f!, Const, BatchDuplicated(Ref(0.0), (Ref(1.0), Ref(2.0))),
             BatchDuplicated(b, (db1, db2)), BatchDuplicated(Ref(1.1), (dx1, dx2)), Const(4))
    (dx1[], dx2[] / 2)   # main: (3.8720000000000008, 3.8720000000000008); this PR: (5.324000000000002, 5.324000000000002), as Duplicated

    With this change the batched result matches Duplicated for 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.

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

No deployments
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