Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
80 changes: 55 additions & 25 deletions enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -193,6 +193,14 @@ struct ForOpInterfaceReverse
return arith::DivUIOp::create(builder, loc, diff, step);
}

static Value castToType(OpBuilder &builder, Location loc, Value v,
Type targetType) {
if (v.getType() == targetType)
return v;
assert(targetType.isIndex());
return arith::IndexCastOp::create(builder, loc, targetType, v);
}

static std::optional<int64_t> getCheckpointBudget(scf::ForOp forOp) {
if (auto a = forOp->getAttrOfType<IntegerAttr>("enzyme.checkpoint_period"))
return a.getInt();
Expand All @@ -204,7 +212,8 @@ struct ForOpInterfaceReverse
SmallVector<int64_t> shape;
shape.push_back(budget);
shape.append(mt.getShape().begin(), mt.getShape().end());
return MemRefType::get(shape, mt.getElementType());
return MemRefType::get(shape, mt.getElementType(),
MemRefLayoutAttrInterface{}, mt.getMemorySpace());
}
return MemRefType::get({budget}, t);
}
Expand All @@ -218,7 +227,10 @@ struct ForOpInterfaceReverse
strides.push_back(b.getIndexAttr(1));
for (int64_t i = 0, e = rowTy.getRank(); i < e; ++i) {
offsets.push_back(b.getIndexAttr(0));
sizes.push_back(b.getIndexAttr(rowTy.getDimSize(i)));
if (rowTy.isDynamicDim(i))
sizes.push_back(memref::DimOp::create(b, loc, buf, i + 1).getResult());
else
sizes.push_back(b.getIndexAttr(rowTy.getDimSize(i)));
strides.push_back(b.getIndexAttr(1));
}
auto resTy = memref::SubViewOp::inferRankReducedResultType(
Expand All @@ -244,7 +256,11 @@ struct ForOpInterfaceReverse
Type valTy) {
if (auto mt = dyn_cast<MemRefType>(valTy)) {
Value row = checkpointRow(b, loc, buf, slot, mt);
Value fresh = memref::AllocOp::create(b, loc, mt);
SmallVector<Value> dynSizes;
for (int64_t i = 0, e = mt.getRank(); i < e; ++i)
if (mt.isDynamicDim(i))
dynSizes.push_back(memref::DimOp::create(b, loc, row, i));
Value fresh = memref::AllocOp::create(b, loc, mt, dynSizes);
memref::CopyOp::create(b, loc, row, fresh);
return fresh;
}
Expand Down Expand Up @@ -280,6 +296,7 @@ struct ForOpInterfaceReverse
startV = gutils->getNewFromOriginal(forOp.getLowerBound());
stepV = gutils->getNewFromOriginal(forOp.getStep());
numItersV = getNumIterationsValue(builder, loc, forOp, gutils);
numItersV = castToType(builder, loc, numItersV, idxTy);
} else {
int64_t numIters =
ForOpEnzymeOpsRemover::getConstantNumberOfIterations(forOp).value();
Expand All @@ -302,7 +319,7 @@ struct ForOpInterfaceReverse

SetVector<Value> outsideRefs;
getUsedValuesDefinedAbove(forOp->getRegions(), outsideRefs);
SmallVector<Value> immutableRefs, mutableRefs;
SmallVector<Value> immutableRefs, mutableRefs, mutableRefsCaches;
for (auto ref : outsideRefs) {
if (isa<ClonableTypeInterface>(ref.getType()))
mutableRefs.push_back(ref);
Expand Down Expand Up @@ -349,15 +366,26 @@ struct ForOpInterfaceReverse
Value split = enzyme::BinomialProgressOp::create(builder, loc, idxTy,
numStepsRem, budgetRem);

for (auto ref : mutableRefs) {
auto iface = cast<ClonableTypeInterface>(ref.getType());
Value clone = iface.cloneValue(builder, mapping.lookupOrDefault(ref));
mutableRefsCaches.push_back(gutils->initAndPushCache(clone, builder));
}

// Inner recompute loop: advance the primal `split` steps.
auto innerFwd =
scf::ForOp::create(builder, loc, c0, split, c1,
SmallVector<Value>(state.begin(), state.end()));
preserveAttributesButCheckpointing(innerFwd, forOp);

// Remove scf.yield automatically added when there are no carried values
if (!innerFwd.getBody()->empty())
innerFwd.getBody()->front().erase();

builder.setInsertionPointToStart(innerFwd.getBody());
Value i = innerFwd.getInductionVar();
Value globalStep = arith::AddIOp::create(builder, loc, stepCtr, i);
globalStep = castToType(builder, loc, globalStep, stepV.getType());
Value iv = arith::AddIOp::create(
builder, loc, startV,
arith::MulIOp::create(builder, loc, stepV, globalStep));
Expand Down Expand Up @@ -390,11 +418,8 @@ struct ForOpInterfaceReverse
caches.push_back(gutils->initAndPushCache(buf, builder));
caches.push_back(gutils->initAndPushCache(idxBuf, builder));

for (auto ref : mutableRefs) {
auto iface = cast<ClonableTypeInterface>(ref.getType());
Value clone = iface.cloneValue(builder, mapping.lookupOrDefault(ref));
caches.push_back(gutils->initAndPushCache(clone, builder));
}
caches.append(mutableRefsCaches);

for (auto ref : immutableRefs)
caches.push_back(
gutils->initAndPushCache(mapping.lookupOrDefault(ref), builder));
Expand Down Expand Up @@ -449,13 +474,8 @@ struct ForOpInterfaceReverse
ckptBufs.push_back(gutils->popCache(caches[j], builder));
Value idxBuf = gutils->popCache(caches[numIterArgs], builder);

size_t cacheIdx = numIterArgs + 1;
SmallVector<Value> cachedMutableRefs;
for (auto ref : mutableRefs) {
Value v = gutils->popCache(caches[cacheIdx++], builder);
cachedMutableRefs.push_back(v);
mapping.map(ref, v);
}
size_t cacheIdx = numIterArgs + mutableRefs.size() + 1;

for (auto ref : immutableRefs)
mapping.map(ref, gutils->popCache(caches[cacheIdx++], builder));

Expand Down Expand Up @@ -500,6 +520,14 @@ struct ForOpInterfaceReverse
OpBuilder::InsertionGuard guard(builder);
builder.setInsertionPointToStart(revOuter.getBody());

SmallVector<Value> cachedMutableRefs;
cacheIdx = numIterArgs + 1;
for (auto ref : mutableRefs) {
Value v = gutils->popCache(caches[cacheIdx++], builder);
cachedMutableRefs.push_back(v);
mapping.map(ref, v);
}

Value ivO = revOuter.getInductionVar();
Value sp = revOuter.getBody()->getArgument(1);
auto adjArgs = revOuter.getBody()->getArguments().drop_front(2);
Expand Down Expand Up @@ -571,6 +599,9 @@ struct ForOpInterfaceReverse
scf::ForOp::create(builder, loc, pos, rematUB, c1,
SmallVector<Value>(astate.begin(), astate.end()));
preserveAttributesButCheckpointing(innerRemat, forOp);
if (!innerRemat.getBody()->empty())
innerRemat.getBody()->front().erase();

{
OpBuilder::InsertionGuard g2(builder);
builder.setInsertionPointToStart(innerRemat.getBody());
Expand Down Expand Up @@ -607,9 +638,10 @@ struct ForOpInterfaceReverse

// Adjoint of a single body step at (currentRevStep - 1).
Value stepAdj = arith::SubIOp::create(builder, loc, currentRevStep, c1);
Value stepAdjC = castToType(builder, loc, stepAdj, stepV.getType());
Value ivAdj = arith::AddIOp::create(
builder, loc, startV,
arith::MulIOp::create(builder, loc, stepV, stepAdj));
arith::MulIOp::create(builder, loc, stepV, stepAdjC));

for (auto &&[oldArg, newArg] : llvm::zip_equal(
forOp.getBody()->getArguments().drop_front(), reconState))
Expand Down Expand Up @@ -667,6 +699,10 @@ struct ForOpInterfaceReverse
}
}

for (auto ref : cachedMutableRefs)
if (auto iface = dyn_cast<ClonableTypeInterface>(ref.getType()))
iface.freeClonedValue(builder, ref);

SmallVector<Value> outerYields;
outerYields.push_back(newSp);
outerYields.append(newAdjoints.begin(), newAdjoints.end());
Expand All @@ -688,9 +724,6 @@ struct ForOpInterfaceReverse
for (auto buf : ckptBufs)
memref::DeallocOp::create(builder, loc, buf);
memref::DeallocOp::create(builder, loc, idxBuf);
for (auto ref : cachedMutableRefs)
if (auto iface = dyn_cast<ClonableTypeInterface>(ref.getType()))
iface.freeClonedValue(builder, ref);

return success(valid);
}
Expand Down Expand Up @@ -1138,11 +1171,8 @@ struct ForOpInterfaceReverse
Location loc = forOp.getInductionVar().getLoc();
auto currentIV = arith::MulIOp::create(
cacheBuilder, loc,
arith::AddIOp::create(
cacheBuilder, loc,
arith::MulIOp::create(cacheBuilder, loc,
outerFwd.getInductionVar(), nInnerCst),
innerFwd.getInductionVar()),
arith::AddIOp::create(cacheBuilder, loc, outerFwd.getInductionVar(),
innerFwd.getInductionVar()),
newForOp.getStep());

for (auto [oldArg, newArg] :
Expand Down
30 changes: 30 additions & 0 deletions enzyme/Enzyme/MLIR/Passes/EnzymeToMemRef.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -338,6 +338,34 @@ struct SetOpConversion : public OpConversionPattern<enzyme::SetOp> {
}
};

// `enzyme.load`/`enzyme.store` are `memref.load`/`memref.store` annotated
// with the (possibly dynamic) dimension sizes of the memref, for use by
// passes that need that information (e.g. mincut). The sizes are not needed
// to actually perform the load/store, so lowering simply drops them.
struct LoadOpConversion : public OpConversionPattern<enzyme::LoadOp> {
using OpConversionPattern::OpConversionPattern;

LogicalResult
matchAndRewrite(enzyme::LoadOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
rewriter.replaceOpWithNewOp<memref::LoadOp>(op, adaptor.getMemref(),
adaptor.getIndices());
return success();
}
};

struct StoreOpConversion : public OpConversionPattern<enzyme::StoreOp> {
using OpConversionPattern::OpConversionPattern;

LogicalResult
matchAndRewrite(enzyme::StoreOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
rewriter.replaceOpWithNewOp<memref::StoreOp>(
op, adaptor.getValue(), adaptor.getMemref(), adaptor.getIndices());
return success();
}
};

struct GetOpConversion : public OpConversionPattern<enzyme::GetOp> {
using OpConversionPattern<enzyme::GetOp>::OpConversionPattern;

Expand Down Expand Up @@ -403,6 +431,8 @@ struct EnzymeToMemRefPass
patterns.add<PopOpConversion>(typeConverter, context);
patterns.add<SetOpConversion>(typeConverter, context);
patterns.add<GetOpConversion>(typeConverter, context);
patterns.add<LoadOpConversion>(typeConverter, context);
patterns.add<StoreOpConversion>(typeConverter, context);

ConversionTarget target(*context);
target.addLegalDialect<memref::MemRefDialect>();
Expand Down
102 changes: 68 additions & 34 deletions enzyme/Enzyme/MLIR/Passes/RemovalUtils.h
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,35 @@ static Value ensureIndexType(Value value, OpBuilder &builder) {
builder.getIndexType(), value);
}

static Value traceToNonBlockArgSource(Value value) {
while (!isa<BlockArgument>(value)) {
if (auto svOp = value.getDefiningOp<memref::SubViewOp>())
value = svOp.getSource();
else if (auto castOp = value.getDefiningOp<memref::CastOp>())
value = castOp.getSource();
else if (auto rcOp = value.getDefiningOp<memref::ReinterpretCastOp>())
value = rcOp.getSource();
else
break;
}
return isa<BlockArgument>(value) ? nullptr : value;
}

static MultidimensionalAllocInterface
getMultiDimCacheAllocOp(Value pushedValue, LoopCacheType cacheType,
ShapedType shapeTy, Operation *forOp) {
if (cacheType != LoopCacheType::MEMREF)
return nullptr;
Value source = traceToNonBlockArgSource(pushedValue);
if (!source)
return nullptr;
auto allocOp =
dyn_cast_or_null<MultidimensionalAllocInterface>(source.getDefiningOp());
if (!allocOp || !allocOp.hoistable(forOp))
return nullptr;
return allocOp;
}

template <typename FinalClass, typename OpName>
struct ForLikeEnzymeOpsRemover
: public EnzymeOpsRemoverOpInterface::ExternalModel<FinalClass, OpName> {
Expand Down Expand Up @@ -429,18 +458,16 @@ struct ForLikeEnzymeOpsRemover
Attribute memorySpace = nullptr;
MultidimensionalAllocInterface allocOp;
if (auto ST = dyn_cast<ShapedType>(ET)) {
allocOp = dyn_cast_or_null<MultidimensionalAllocInterface>(
pushedValue.getDefiningOp());
if (cacheType == LoopCacheType::MEMREF && allocOp &&
allocOp.hoistable(forOp)) {
multiDim = true;
allocOp = getMultiDimCacheAllocOp(pushedValue, cacheType, ST, forOp);
bool fullyStatic = llvm::all_of(ST.getShape(), [](int64_t dim) {
return dim != ShapedType::kDynamic;
});
multiDim =
allocOp || (fullyStatic && cacheType == LoopCacheType::TENSOR);
if (allocOp) {
if (auto MT = dyn_cast<MemRefType>(pushedValue.getType()))
memorySpace = MT.getMemorySpace();
allocOp.appendDynamicDims(dynamicDims);
} else if (llvm::all_of(ST.getShape(), [](int64_t dim) {
return dim != ShapedType::kDynamic;
})) {
multiDim = true;
}

if (multiDim) {
Expand All @@ -449,6 +476,11 @@ struct ForLikeEnzymeOpsRemover
}
}

if (!memorySpace) {
if (auto innerMT = dyn_cast<MemRefType>(ET))
memorySpace = innerMT.getMemorySpace();
}

auto newType = cacheType == LoopCacheType::TENSOR
? cast<ShapedType>(RankedTensorType::get(newShape, ET))
: cast<ShapedType>(MemRefType::get(newShape, ET));
Expand Down Expand Up @@ -551,16 +583,21 @@ struct ForLikeEnzymeOpsRemover
MT.getShape(), cast<MemRefType>(initValue.getType()), offsets,
sizes, strides);

rewriter.setInsertionPoint(memref.getDefiningOp());
rewriter.replaceOpWithNewOp<memref::SubViewOp>(
memref.getDefiningOp(), RT, initValue,
/*offsets*/ inductionVariable,
/*sizes*/ dynSizes,
/*strides*/ ValueRange(),
/*static_offsets*/ rewriter.getDenseI64ArrayAttr(offsets),
/*static_sizes*/ rewriter.getDenseI64ArrayAttr(sizes),
/*static_strides*/ rewriter.getDenseI64ArrayAttr(strides));

rewriter.setInsertionPointAfterValue(memref);
rewriter.replaceAllUsesWith(
memref,
memref::SubViewOp::create(
rewriter, memref.getLoc(), RT, initValue,
/*offsets*/ inductionVariable,
/*sizes*/ dynSizes,
/*strides*/ ValueRange(),
/*static_offsets*/ rewriter.getDenseI64ArrayAttr(offsets),
/*static_sizes*/ rewriter.getDenseI64ArrayAttr(sizes),
/*static_strides*/ rewriter.getDenseI64ArrayAttr(strides))
.getResult());

if (auto defOp = memref.getDefiningOp())
rewriter.eraseOp(defOp);
} else {
if (dynamicDims.empty()) {
memref::StoreOp::create(rewriter, info.pushOp->getLoc(),
Expand Down Expand Up @@ -699,28 +736,25 @@ struct ForLikeEnzymeOpsRemover
if (auto ST = dyn_cast<ShapedType>(ET)) {
if (auto MT = dyn_cast<MemRefType>(ST))
memorySpace = MT.getMemorySpace();

auto svOp = info.pushedValue().getDefiningOp<memref::SubViewOp>();
if (svOp) {
allocOp = dyn_cast_or_null<MultidimensionalAllocInterface>(
svOp.getSource().getDefiningOp());
if (cacheType == LoopCacheType::MEMREF)
multiDim = true;
} else {
allocOp = dyn_cast_or_null<MultidimensionalAllocInterface>(
info.pushedValue().getDefiningOp());
if (llvm::all_of(ST.getShape(), [](int64_t dim) {
return dim != ShapedType::kDynamic;
}))
multiDim = true;
}
allocOp =
getMultiDimCacheAllocOp(info.pushedValue(), cacheType, ST, forOp);
bool fullyStatic = llvm::all_of(ST.getShape(), [](int64_t dim) {
return dim != ShapedType::kDynamic;
});
multiDim =
allocOp || (fullyStatic && cacheType == LoopCacheType::TENSOR);

if (multiDim) {
newShape.append(ST.getShape().begin(), ST.getShape().end());
ET = ST.getElementType();
}
}

if (!memorySpace) {
if (auto innerMT = dyn_cast<MemRefType>(ET))
memorySpace = innerMT.getMemorySpace();
}

auto newType = cacheType == LoopCacheType::TENSOR
? cast<ShapedType>(RankedTensorType::get(newShape, ET))
: cast<ShapedType>(MemRefType::get(newShape, ET));
Expand Down
Loading
Loading