diff --git a/src/enzyme_ad/jax/Utils.cpp b/src/enzyme_ad/jax/Utils.cpp index e07fd35944..62e846b229 100644 --- a/src/enzyme_ad/jax/Utils.cpp +++ b/src/enzyme_ad/jax/Utils.cpp @@ -350,120 +350,6 @@ bool getEffectsAfter(Operation *op, return !conservative; } -bool isCaptured(Value v, Operation *potentialUser = nullptr, - bool *seenuse = nullptr) { - SmallVector todo = {v}; - while (todo.size()) { - Value v = todo.pop_back_val(); - for (auto u : v.getUsers()) { - if (seenuse && u == potentialUser) - *seenuse = true; - if (isa(u)) - continue; - // if (isa(u)) continue - if (auto s = dyn_cast(u)) { - if (s.getValue() == v) - return true; - continue; - } - if (auto s = dyn_cast(u)) { - if (s.getValue() == v) - return true; - continue; - } - if (auto s = dyn_cast(u)) { - if (s.getValue() == v) - return true; - continue; - } - if (auto sub = dyn_cast(u)) { - todo.push_back(sub); - } - if (auto sub = dyn_cast(u)) { - todo.push_back(sub); - } - if (auto sub = dyn_cast(u)) { - todo.push_back(sub); - } - if (auto sub = dyn_cast(u)) { - continue; - } - if (auto sub = dyn_cast(u)) { - continue; - } - if (auto sub = dyn_cast(u)) { - continue; - } - if (auto sub = dyn_cast(u)) { - continue; - } - if (auto sub = dyn_cast(u)) { - todo.push_back(sub); - } - if (auto sub = dyn_cast(u)) { - continue; - } - // if (auto sub = dyn_cast(u)) { - // todo.push_back(sub); - // } - if (auto sub = dyn_cast(u)) { - todo.push_back(sub); - } - if (auto sub = dyn_cast(u)) { - todo.push_back(sub); - } - if (auto cop = dyn_cast(u)) { - if (auto callee = cop.getCallee()) { - if (getNonCapturingFunctions().count(callee->str())) - continue; - } - } - if (auto cop = dyn_cast(u)) { - if (getNonCapturingFunctions().count(cop.getCallee().str())) - continue; - } - return true; - } - } - - return false; -} - -Value getBase(Value v) { - while (true) { - // if (auto s = v.getDefiningOp()) { - // v = s.getSource(); - // continue; - // } - if (auto s = v.getDefiningOp()) { - v = s.getSource(); - continue; - } - if (auto s = v.getDefiningOp()) { - v = s.getSource(); - continue; - } - if (auto s = v.getDefiningOp()) { - v = s.getBase(); - continue; - } - if (auto s = v.getDefiningOp()) { - v = s.getArg(); - continue; - } - if (auto s = v.getDefiningOp()) { - v = s.getArg(); - continue; - } - if (auto s = v.getDefiningOp()) { - v = s.getSource(); - continue; - } - break; - } - return v; -} - bool isStackAlloca(Value v) { return v.getDefiningOp() || v.getDefiningOp() || @@ -505,9 +391,10 @@ bool mayWriteTo(Operation *op, Value val, bool ignoreBarrier) { // Calls which do not use a derived pointer of a known alloca, which is not // captured can not write to said memory. if (auto callOp = dyn_cast(op)) { - auto base = getBase(val); + auto base = enzyme::oputils::getBaseObject(val); bool seenuse = false; - if (isStackAlloca(base) && !isCaptured(base, op, &seenuse) && !seenuse) { + if (isStackAlloca(base) && + !enzyme::oputils::isCaptured(base, op, &seenuse) && !seenuse) { return false; } } @@ -1407,9 +1294,10 @@ bool mayReadFrom(Operation *op, Value val) { return false; } if (auto callOp = dyn_cast(op)) { - auto base = getBase(val); + auto base = enzyme::oputils::getBaseObject(val); bool seenuse = false; - if (isStackAlloca(base) && !isCaptured(base, op, &seenuse) && !seenuse) { + if (isStackAlloca(base) && + !enzyme::oputils::isCaptured(base, op, &seenuse) && !seenuse) { return false; } } diff --git a/workspace.bzl b/workspace.bzl index f394d269ba..5e2cf85d14 100644 --- a/workspace.bzl +++ b/workspace.bzl @@ -1,7 +1,7 @@ JAX_COMMIT = "fbf6588d55c5662ccf98558015d5b2c4c4a7bf6f" JAX_SHA256 = "" -ENZYME_COMMIT = "3cd66b79ce224b8dc35cfac0e37038ed186b2013" +ENZYME_COMMIT = "56b02e7e94265e7af38fe95e9fbf9818ef1d09c0" ENZYME_SHA256 = "" ML_TOOLCHAIN_COMMIT = "30ef4a9096f9490e8f198faa5ce5bbddd1b72fdb"