From 3cb9ccf2a76fdfca3908e359c938faaeab1afdfb Mon Sep 17 00:00:00 2001 From: Paul Berg <9824244+Pangoraw@users.noreply.github.com> Date: Mon, 27 Jul 2026 10:36:43 -0500 Subject: [PATCH 1/5] Fixes for mutable memory in scf.for reverse --- .../SCFAutoDiffOpInterfaceImpl.cpp | 73 ++++++++---- enzyme/Enzyme/MLIR/Passes/RemovalUtils.h | 102 +++++++++++------ ...checkpointing_binomial_mutable_memory.mlir | 106 ++++++++++++++++++ .../scf_for_checkpointing_mutable_memory.mlir | 95 ++++++++++++++++ ...scf_for_memref_binomial_checkpointing.mlir | 14 ++- .../ReverseMode/scf_for_mutable_memory.mlir | 71 ++++++++++++ 6 files changed, 401 insertions(+), 60 deletions(-) create mode 100644 enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_binomial_mutable_memory.mlir create mode 100644 enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_mutable_memory.mlir create mode 100644 enzyme/test/MLIR/ReverseMode/scf_for_mutable_memory.mlir diff --git a/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp b/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp index 88e0df4aa6b1..8dad493f9fdb 100644 --- a/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp +++ b/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp @@ -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 getCheckpointBudget(scf::ForOp forOp) { if (auto a = forOp->getAttrOfType("enzyme.checkpoint_period")) return a.getInt(); @@ -204,7 +212,8 @@ struct ForOpInterfaceReverse SmallVector 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); } @@ -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( @@ -244,7 +256,11 @@ struct ForOpInterfaceReverse Type valTy) { if (auto mt = dyn_cast(valTy)) { Value row = checkpointRow(b, loc, buf, slot, mt); - Value fresh = memref::AllocOp::create(b, loc, mt); + SmallVector 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; } @@ -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(); @@ -302,7 +319,7 @@ struct ForOpInterfaceReverse SetVector outsideRefs; getUsedValuesDefinedAbove(forOp->getRegions(), outsideRefs); - SmallVector immutableRefs, mutableRefs; + SmallVector immutableRefs, mutableRefs, mutableRefsCaches; for (auto ref : outsideRefs) { if (isa(ref.getType())) mutableRefs.push_back(ref); @@ -349,15 +366,26 @@ struct ForOpInterfaceReverse Value split = enzyme::BinomialProgressOp::create(builder, loc, idxTy, numStepsRem, budgetRem); + for (auto ref : mutableRefs) { + auto iface = cast(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(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)); @@ -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(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)); @@ -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 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)); @@ -500,6 +520,14 @@ struct ForOpInterfaceReverse OpBuilder::InsertionGuard guard(builder); builder.setInsertionPointToStart(revOuter.getBody()); + SmallVector 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); @@ -571,6 +599,9 @@ struct ForOpInterfaceReverse scf::ForOp::create(builder, loc, pos, rematUB, c1, SmallVector(astate.begin(), astate.end())); preserveAttributesButCheckpointing(innerRemat, forOp); + if (!innerRemat.getBody()->empty()) + innerRemat.getBody()->front().erase(); + { OpBuilder::InsertionGuard g2(builder); builder.setInsertionPointToStart(innerRemat.getBody()); @@ -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)) @@ -667,6 +699,10 @@ struct ForOpInterfaceReverse } } + for (auto ref : cachedMutableRefs) + if (auto iface = dyn_cast(ref.getType())) + iface.freeClonedValue(builder, ref); + SmallVector outerYields; outerYields.push_back(newSp); outerYields.append(newAdjoints.begin(), newAdjoints.end()); @@ -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(ref.getType())) - iface.freeClonedValue(builder, ref); return success(valid); } diff --git a/enzyme/Enzyme/MLIR/Passes/RemovalUtils.h b/enzyme/Enzyme/MLIR/Passes/RemovalUtils.h index 13f91c2396ba..d64afa9f8673 100644 --- a/enzyme/Enzyme/MLIR/Passes/RemovalUtils.h +++ b/enzyme/Enzyme/MLIR/Passes/RemovalUtils.h @@ -98,6 +98,35 @@ static Value ensureIndexType(Value value, OpBuilder &builder) { builder.getIndexType(), value); } +static Value traceToNonBlockArgSource(Value value) { + while (!isa(value)) { + if (auto svOp = value.getDefiningOp()) + value = svOp.getSource(); + else if (auto castOp = value.getDefiningOp()) + value = castOp.getSource(); + else if (auto rcOp = value.getDefiningOp()) + value = rcOp.getSource(); + else + break; + } + return isa(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(source.getDefiningOp()); + if (!allocOp || !allocOp.hoistable(forOp)) + return nullptr; + return allocOp; +} + template struct ForLikeEnzymeOpsRemover : public EnzymeOpsRemoverOpInterface::ExternalModel { @@ -429,18 +458,16 @@ struct ForLikeEnzymeOpsRemover Attribute memorySpace = nullptr; MultidimensionalAllocInterface allocOp; if (auto ST = dyn_cast(ET)) { - allocOp = dyn_cast_or_null( - 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(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) { @@ -449,6 +476,11 @@ struct ForLikeEnzymeOpsRemover } } + if (!memorySpace) { + if (auto innerMT = dyn_cast(ET)) + memorySpace = innerMT.getMemorySpace(); + } + auto newType = cacheType == LoopCacheType::TENSOR ? cast(RankedTensorType::get(newShape, ET)) : cast(MemRefType::get(newShape, ET)); @@ -551,16 +583,21 @@ struct ForLikeEnzymeOpsRemover MT.getShape(), cast(initValue.getType()), offsets, sizes, strides); - rewriter.setInsertionPoint(memref.getDefiningOp()); - rewriter.replaceOpWithNewOp( - 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(), @@ -699,21 +736,13 @@ struct ForLikeEnzymeOpsRemover if (auto ST = dyn_cast(ET)) { if (auto MT = dyn_cast(ST)) memorySpace = MT.getMemorySpace(); - - auto svOp = info.pushedValue().getDefiningOp(); - if (svOp) { - allocOp = dyn_cast_or_null( - svOp.getSource().getDefiningOp()); - if (cacheType == LoopCacheType::MEMREF) - multiDim = true; - } else { - allocOp = dyn_cast_or_null( - 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()); @@ -721,6 +750,11 @@ struct ForLikeEnzymeOpsRemover } } + if (!memorySpace) { + if (auto innerMT = dyn_cast(ET)) + memorySpace = innerMT.getMemorySpace(); + } + auto newType = cacheType == LoopCacheType::TENSOR ? cast(RankedTensorType::get(newShape, ET)) : cast(MemRefType::get(newShape, ET)); diff --git a/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_binomial_mutable_memory.mlir b/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_binomial_mutable_memory.mlir new file mode 100644 index 000000000000..89c0881bbd9e --- /dev/null +++ b/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_binomial_mutable_memory.mlir @@ -0,0 +1,106 @@ +// RUN: %eopt %s --enzyme-wrap="infn=reduce_sum outfn= argTys=enzyme_dup retTys=enzyme_active mode=ReverseModeCombined" --canonicalize --enzyme-simplify-math --remove-unnecessary-enzyme-ops --canonicalize | FileCheck %s + +module { + func.func @reduce_sum(%buf: memref<10xf64>) -> f64 { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c10 = arith.constant 10 : index + %init = arith.constant 0.0 : f64 + + %sum = scf.for %i = %c0 to %c10 step %c1 iter_args(%acc = %init) -> (f64) { + %val = memref.load %buf[%i] : memref<10xf64> + %new_acc = arith.addf %acc, %val : f64 + memref.store %new_acc, %buf[%c0] : memref<10xf64> + scf.yield %new_acc : f64 + } {enzyme.enable_checkpointing = true, + enzyme.binomial_checkpointing, + enzyme.checkpoint_period=4, + enzyme.disable_mincut=true} + + return %sum : f64 + } +} + + +// CHECK: func.func @reduce_sum(%arg0: memref<10xf64>, %arg1: memref<10xf64>, %arg2: f64) { +// CHECK-NEXT: %c9 = arith.constant 9 : index +// CHECK-NEXT: %c4 = arith.constant 4 : index +// CHECK-NEXT: %c10 = arith.constant 10 : index +// CHECK-NEXT: %c1 = arith.constant 1 : index +// CHECK-NEXT: %c0 = arith.constant 0 : index +// CHECK-NEXT: %cst = arith.constant 0.000000e+00 : f64 +// CHECK-NEXT: %alloc = memref.alloc() : memref<4xf64> +// CHECK-NEXT: %alloc_0 = memref.alloc() : memref<4xindex> +// CHECK-NEXT: %alloc_1 = memref.alloc() : memref<4x10xf64> +// CHECK-NEXT: %0:2 = scf.for %arg3 = %c0 to %c4 step %c1 iter_args(%arg4 = %c0, %arg5 = %cst) -> (index, f64) { +// CHECK-NEXT: memref.store %arg5, %alloc[%arg3] : memref<4xf64> +// CHECK-NEXT: memref.store %arg4, %alloc_0[%arg3] : memref<4xindex> +// CHECK-NEXT: %3 = arith.subi %c10, %arg4 : index +// CHECK-NEXT: %4 = arith.subi %c4, %arg3 : index +// CHECK-NEXT: %5 = arith.minui %4, %3 : index +// CHECK-NEXT: %6 = enzyme.binomial_progress %3, %5 : index +// CHECK-NEXT: %subview = memref.subview %alloc_1[%arg3, 0] [1, 10] [1, 1] : memref<4x10xf64> to memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: memref.copy %arg0, %subview : memref<10xf64> to memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: %7 = scf.for %arg6 = %c0 to %6 step %c1 iter_args(%arg7 = %arg5) -> (f64) { +// CHECK-NEXT: %9 = arith.addi %arg4, %arg6 : index +// CHECK-NEXT: %10 = memref.load %arg0[%9] : memref<10xf64> +// CHECK-NEXT: %11 = arith.addf %arg7, %10 : f64 +// CHECK-NEXT: memref.store %11, %arg0[%c0] : memref<10xf64> +// CHECK-NEXT: scf.yield %11 : f64 +// CHECK-NEXT: } {enzyme.disable_mincut = true} +// CHECK-NEXT: %8 = arith.addi %arg4, %6 : index +// CHECK-NEXT: scf.yield %8, %7 : index, f64 +// CHECK-NEXT: } +// CHECK-NEXT: %1 = arith.addf %arg2, %cst : f64 +// CHECK-NEXT: %2:2 = scf.for %arg3 = %c0 to %c10 step %c1 iter_args(%arg4 = %c4, %arg5 = %1) -> (index, f64) { +// CHECK-NEXT: %3 = arith.subi %c9, %arg3 : index +// CHECK-NEXT: %subview = memref.subview %alloc_1[%3, 0] [1, 10] [1, 1] : memref<4x10xf64> to memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: %4 = arith.subi %arg4, %c1 : index +// CHECK-NEXT: %5 = arith.subi %c10, %arg3 : index +// CHECK-NEXT: %6 = memref.load %alloc[%4] : memref<4xf64> +// CHECK-NEXT: %7 = memref.load %alloc_0[%4] : memref<4xindex> +// CHECK-NEXT: %8:3 = scf.while (%arg6 = %7, %arg7 = %4, %arg8 = %6) : (index, index, f64) -> (index, index, f64) { +// CHECK-NEXT: %19 = arith.addi %arg6, %c1 : index +// CHECK-NEXT: %20 = arith.cmpi slt, %19, %5 : index +// CHECK-NEXT: scf.condition(%20) %arg6, %arg7, %arg8 : index, index, f64 +// CHECK-NEXT: } do { +// CHECK-NEXT: ^bb0(%arg6: index, %arg7: index, %arg8: f64): +// CHECK-NEXT: %19 = arith.subi %5, %arg6 : index +// CHECK-NEXT: %20 = arith.subi %c4, %arg7 : index +// CHECK-NEXT: %21 = arith.minui %20, %19 : index +// CHECK-NEXT: %22 = enzyme.binomial_progress %19, %21 : index +// CHECK-NEXT: memref.store %arg8, %alloc[%arg7] : memref<4xf64> +// CHECK-NEXT: memref.store %arg6, %alloc_0[%arg7] : memref<4xindex> +// CHECK-NEXT: %23 = arith.addi %arg6, %22 : index +// CHECK-NEXT: %24 = arith.cmpi eq, %23, %5 : index +// CHECK-NEXT: %25 = arith.subi %23, %c1 : index +// CHECK-NEXT: %26 = arith.select %24, %25, %23 : index +// CHECK-NEXT: %27 = scf.for %arg9 = %arg6 to %26 step %c1 iter_args(%arg10 = %arg8) -> (f64) { +// CHECK-NEXT: %29 = memref.load %subview[%arg9] : memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: %30 = arith.addf %arg10, %29 : f64 +// CHECK-NEXT: memref.store %30, %subview[%c0] : memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: scf.yield %30 : f64 +// CHECK-NEXT: } {enzyme.disable_mincut = true} +// CHECK-NEXT: %28 = arith.addi %arg7, %c1 : index +// CHECK-NEXT: scf.yield %23, %28, %27 : index, index, f64 +// CHECK-NEXT: } +// CHECK-NEXT: %9 = arith.subi %c9, %arg3 : index +// CHECK-NEXT: %10 = memref.load %subview[%9] : memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: %11 = arith.addf %8#2, %10 : f64 +// CHECK-NEXT: memref.store %11, %subview[%c0] : memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: %12 = arith.addf %arg5, %cst : f64 +// CHECK-NEXT: %13 = memref.load %arg1[%c0] : memref<10xf64> +// CHECK-NEXT: %14 = arith.addf %12, %13 : f64 +// CHECK-NEXT: memref.store %cst, %arg1[%c0] : memref<10xf64> +// CHECK-NEXT: %15 = arith.addf %14, %cst : f64 +// CHECK-NEXT: %16 = arith.addf %14, %cst : f64 +// CHECK-NEXT: %17 = memref.load %arg1[%9] : memref<10xf64> +// CHECK-NEXT: %18 = arith.addf %17, %16 : f64 +// CHECK-NEXT: memref.store %18, %arg1[%9] : memref<10xf64> +// CHECK-NEXT: scf.yield %8#1, %15 : index, f64 +// CHECK-NEXT: } {enzyme.disable_mincut = true} +// CHECK-NEXT: memref.dealloc %alloc_1 : memref<4x10xf64> +// CHECK-NEXT: memref.dealloc %alloc : memref<4xf64> +// CHECK-NEXT: memref.dealloc %alloc_0 : memref<4xindex> +// CHECK-NEXT: return +// CHECK-NEXT: } diff --git a/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_mutable_memory.mlir b/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_mutable_memory.mlir new file mode 100644 index 000000000000..51e863402fca --- /dev/null +++ b/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_mutable_memory.mlir @@ -0,0 +1,95 @@ +// RUN: %eopt %s --enzyme-wrap="infn=reduce_sum outfn= argTys=enzyme_dup retTys=enzyme_active mode=ReverseModeCombined" --canonicalize --enzyme-simplify-math --remove-unnecessary-enzyme-ops | FileCheck %s + +func.func @reduce_sum(%buf: memref<10xf64>) -> f64 { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c10 = arith.constant 10 : index + %init = arith.constant 0.0 : f64 + + %sum = scf.for %i = %c0 to %c10 step %c1 iter_args(%acc = %init) -> (f64) { + %val = memref.load %buf[%i] : memref<10xf64> + %new_acc = arith.addf %acc, %val : f64 + memref.store %new_acc, %buf[%c0] : memref<10xf64> + scf.yield %new_acc : f64 + } {enzyme.enable_checkpointing = true, + enzyme.checkpoint_period=4, + enzyme.disable_mincut=true} + + return %sum : f64 +} + +// CHECK: func.func @reduce_sum(%arg0: memref<10xf64>, %arg1: memref<10xf64>, %arg2: f64) { +// CHECK-NEXT: %c4 = arith.constant 4 : index +// CHECK-NEXT: %c9 = arith.constant 9 : index +// CHECK-NEXT: %c12 = arith.constant 12 : index +// CHECK-NEXT: %c3 = arith.constant 3 : index +// CHECK-NEXT: %c1 = arith.constant 1 : index +// CHECK-NEXT: %c0 = arith.constant 0 : index +// CHECK-NEXT: %cst = arith.constant 0.000000e+00 : f64 +// CHECK-NEXT: %alloc = memref.alloc() : memref<4x10xf64> +// CHECK-NEXT: %alloc_0 = memref.alloc() : memref<4xf64> +// CHECK-NEXT: %0 = scf.for %arg3 = %c0 to %c12 step %c3 iter_args(%arg4 = %cst) -> (f64) { +// CHECK-NEXT: %3 = arith.divui %arg3, %c3 : index +// CHECK-NEXT: %4 = arith.cmpi eq, %arg3, %c9 : index +// CHECK-NEXT: %5 = arith.select %4, %c1, %c3 : index +// CHECK-NEXT: %subview = memref.subview %alloc[%3, 0] [1, 10] [1, 1] : memref<4x10xf64> to memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: memref.copy %arg0, %subview : memref<10xf64> to memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: %6 = arith.muli %arg3, %c3 : index +// CHECK-NEXT: %7 = scf.for %arg5 = %c0 to %5 step %c1 iter_args(%arg6 = %arg4) -> (f64) { +// CHECK-NEXT: %8 = arith.addi %6, %arg5 : index +// CHECK-NEXT: %9 = memref.load %arg0[%8] : memref<10xf64> +// CHECK-NEXT: %10 = arith.addf %arg6, %9 : f64 +// CHECK-NEXT: memref.store %10, %arg0[%c0] : memref<10xf64> +// CHECK-NEXT: scf.yield %10 : f64 +// CHECK-NEXT: } {enzyme.disable_mincut = true} +// CHECK-NEXT: memref.store %arg4, %alloc_0[%3] : memref<4xf64> +// CHECK-NEXT: scf.yield %7 : f64 +// CHECK-NEXT: } +// CHECK-NEXT: %1 = arith.addf %arg2, %cst : f64 +// CHECK-NEXT: %2 = scf.for %arg3 = %c0 to %c4 step %c1 iter_args(%arg4 = %1) -> (f64) { +// CHECK-NEXT: %3 = arith.subi %c3, %arg3 : index +// CHECK-NEXT: %subview = memref.subview %alloc[%3, 0] [1, 10] [1, 1] : memref<4x10xf64> to memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: %4 = arith.subi %c3, %arg3 : index +// CHECK-NEXT: %5 = memref.load %alloc_0[%3] : memref<4xf64> +// CHECK-NEXT: %6 = arith.cmpi eq, %arg3, %c0 : index +// CHECK-NEXT: %7 = arith.select %6, %c1, %c3 : index +// CHECK-NEXT: %8 = arith.muli %4, %c3 : index +// CHECK-NEXT: %alloc_1 = memref.alloc(%7) : memref> +// CHECK-NEXT: %alloc_2 = memref.alloc(%7) : memref +// CHECK-NEXT: %alloc_3 = memref.alloc(%7) : memref +// CHECK-NEXT: %9 = scf.for %arg5 = %c0 to %7 step %c1 iter_args(%arg6 = %5) -> (f64) { +// CHECK-NEXT: %11 = arith.addi %8, %arg5 : index +// CHECK-NEXT: enzyme.store %11, %alloc_2[%arg5] ([%7]) : memref +// CHECK-NEXT: %12 = memref.load %subview[%11] : memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: %13 = arith.addf %arg6, %12 : f64 +// CHECK-NEXT: enzyme.store %arg1, %alloc_1[%arg5] ([%7]) : memref> +// CHECK-NEXT: enzyme.store %c0, %alloc_3[%arg5] ([%7]) : memref +// CHECK-NEXT: memref.store %13, %subview[%c0] : memref<10xf64, strided<[1], offset: ?>> +// CHECK-NEXT: scf.yield %13 : f64 +// CHECK-NEXT: } +// CHECK-NEXT: %10 = scf.for %arg5 = %c0 to %7 step %c1 iter_args(%arg6 = %arg4) -> (f64) { +// CHECK-NEXT: %11 = arith.subi %7, %c1 : index +// CHECK-NEXT: %12 = arith.subi %11, %arg5 : index +// CHECK-NEXT: %13 = arith.addf %arg6, %cst : f64 +// CHECK-NEXT: %14 = enzyme.load %alloc_1[%12] ([%7]) : memref> +// CHECK-NEXT: %15 = enzyme.load %alloc_3[%12] ([%7]) : memref +// CHECK-NEXT: %16 = memref.load %14[%15] : memref<10xf64> +// CHECK-NEXT: %17 = arith.addf %13, %16 : f64 +// CHECK-NEXT: memref.store %cst, %14[%15] : memref<10xf64> +// CHECK-NEXT: %18 = arith.addf %17, %cst : f64 +// CHECK-NEXT: %19 = arith.addf %17, %cst : f64 +// CHECK-NEXT: %20 = enzyme.load %alloc_2[%12] ([%7]) : memref +// CHECK-NEXT: %21 = memref.load %14[%20] : memref<10xf64> +// CHECK-NEXT: %22 = arith.addf %21, %19 : f64 +// CHECK-NEXT: memref.store %22, %14[%20] : memref<10xf64> +// CHECK-NEXT: scf.yield %18 : f64 +// CHECK-NEXT: } {enzyme.disable_mincut = true} +// CHECK-NEXT: memref.dealloc %alloc_3 : memref +// CHECK-NEXT: memref.dealloc %alloc_2 : memref +// CHECK-NEXT: memref.dealloc %alloc_1 : memref> +// CHECK-NEXT: scf.yield %10 : f64 +// CHECK-NEXT: } {enzyme.disable_mincut = true} +// CHECK-NEXT: memref.dealloc %alloc_0 : memref<4xf64> +// CHECK-NEXT: memref.dealloc %alloc : memref<4x10xf64> +// CHECK-NEXT: return +// CHECK-NEXT: } diff --git a/enzyme/test/MLIR/ReverseMode/scf_for_memref_binomial_checkpointing.mlir b/enzyme/test/MLIR/ReverseMode/scf_for_memref_binomial_checkpointing.mlir index 37290d6b3fdf..5f7d015ef3a9 100644 --- a/enzyme/test/MLIR/ReverseMode/scf_for_memref_binomial_checkpointing.mlir +++ b/enzyme/test/MLIR/ReverseMode/scf_for_memref_binomial_checkpointing.mlir @@ -26,6 +26,7 @@ module { // CHECK-LABEL: func.func @main( // CHECK-DAG: %[[STATE:.+]] = memref.alloc() : memref<3xf32> // CHECK-DAG: %[[IDX:.+]] = memref.alloc() : memref<3xindex> +// CHECK-NEXT: %[[CLONE_ARG0:.+]] = memref.alloc() : memref<3xf32> // Forward checkpoint-placement loop (budget = 3). // CHECK: scf.for {{.*}} = %c0 to %c3 step %c1 @@ -33,15 +34,16 @@ module { // CHECK: memref.store {{.*}}, %[[IDX]] // The mutable outside reference is snapshotted (cloned) for the reverse pass. -// CHECK: %[[CLONE:.+]] = memref.alloc() : memref -// CHECK: memref.copy %arg0, %[[CLONE]] +// CHECK: %[[SUBVIEW:.+]] = memref.subview %[[CLONE_ARG0]][%{{.+}}] [1] [1] : memref<3xf32> to memref> +// CHECK-NEXT: memref.copy %arg0, %[[SUBVIEW]] : memref to memref> // Reverse loop over all 9 steps, with the remat scf.while. // CHECK: scf.for {{.*}} = %c0 to %c9 step %c1 +// CHECK: %[[SUBVIEWREV:.+]] = memref.subview %[[CLONE_ARG0]][%2] [1] [1] : memref<3xf32> to memref> // CHECK: scf.while - +// CHECK: %[[VAL:.+]] = memref.load %[[SUBVIEWREV]][] : memref> // All allocations are freed. -// CHECK-DAG: memref.dealloc %[[STATE]] -// CHECK-DAG: memref.dealloc %[[IDX]] -// CHECK-DAG: memref.dealloc %[[CLONE]] +// CHECK-DAG: memref.dealloc %[[STATE]] : memref<3xf32> +// CHECK-DAG: memref.dealloc %[[IDX]] : memref<3xindex> +// CHECK-DAG: memref.dealloc %[[CLONE_ARG0]] : memref<3xf32> // CHECK: return diff --git a/enzyme/test/MLIR/ReverseMode/scf_for_mutable_memory.mlir b/enzyme/test/MLIR/ReverseMode/scf_for_mutable_memory.mlir new file mode 100644 index 000000000000..28294c85bb18 --- /dev/null +++ b/enzyme/test/MLIR/ReverseMode/scf_for_mutable_memory.mlir @@ -0,0 +1,71 @@ +// RUN: %eopt %s --enzyme-wrap="infn=reduce_sum outfn= argTys=enzyme_dup retTys=enzyme_active mode=ReverseModeCombined" --canonicalize --enzyme-simplify-math --remove-unnecessary-enzyme-ops | FileCheck %s + +// Check that the accumulated gradients of a memref that is both read and +// written to inside the loop (%buf[%i] is read, %buf[%c0] is written every +// iteration) end up accumulated into the shadow memref (%arg1), the dup of +// the original argument, rather than being lost or accumulated somewhere +// else. + +func.func @reduce_sum(%buf: memref<10xf64>) -> f64 { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c10 = arith.constant 10 : index + %init = arith.constant 0.0 : f64 + + %sum = scf.for %i = %c0 to %c10 step %c1 iter_args(%acc = %init) -> (f64) { + %val = memref.load %buf[%i] : memref<10xf64> + %new_acc = arith.addf %acc, %val : f64 + memref.store %new_acc, %buf[%c0] : memref<10xf64> + scf.yield %new_acc : f64 + } {enzyme.enable_checkpointing = false, + enzyme.checkpoint_period=4, enzyme.disable_mincut=true} + + return %sum : f64 +} + +// CHECK-LABEL: func.func @reduce_sum( +// CHECK-SAME: %[[ARG0:.*]]: memref<10xf64>, %[[ARG1:.*]]: memref<10xf64>, %[[ARG2:.*]]: f64) { +// CHECK: %[[C9:.*]] = arith.constant 9 : index +// CHECK: %[[C10:.*]] = arith.constant 10 : index +// CHECK: %[[C1:.*]] = arith.constant 1 : index +// CHECK: %[[C0:.*]] = arith.constant 0 : index +// CHECK: %[[CST:.*]] = arith.constant 0.000000e+00 : f64 +// CHECK: %[[ALLOC:.*]] = memref.alloc() : memref<10xmemref<10xf64>> +// CHECK: %[[ALLOC_0:.*]] = memref.alloc() : memref<10xindex> +// CHECK: %[[ALLOC_1:.*]] = memref.alloc() : memref<10xindex> +// CHECK: %[[FOR_0:.*]] = scf.for %[[IV:.*]] = %[[C0]] to %[[C10]] step %[[C1]] iter_args(%[[ACC:.*]] = %[[CST]]) -> (f64) { +// CHECK: memref.store %[[IV]], %[[ALLOC_0]]{{\[}}%[[IV]]] : memref<10xindex> +// CHECK: %[[LOAD_0:.*]] = memref.load %[[ARG0]]{{\[}}%[[IV]]] : memref<10xf64> +// CHECK: %[[ADDF_0:.*]] = arith.addf %[[ACC]], %[[LOAD_0]] : f64 +// Every iteration caches the shadow memref itself (%[[ARG1]]), not a fresh +// per-iteration copy -- it is a mutable handle whose identity must be +// preserved across iterations. +// CHECK: memref.store %[[ARG1]], %[[ALLOC]]{{\[}}%[[IV]]] : memref<10xmemref<10xf64>> +// CHECK: memref.store %[[C0]], %[[ALLOC_1]]{{\[}}%[[IV]]] : memref<10xindex> +// CHECK: memref.store %[[ADDF_0]], %[[ARG0]]{{\[}}%[[C0]]] : memref<10xf64> +// CHECK: scf.yield %[[ADDF_0]] : f64 +// CHECK: } +// CHECK: %[[ADDF_1:.*]] = arith.addf %[[ARG2]], %[[CST]] : f64 +// CHECK: %[[FOR_1:.*]] = scf.for %[[IV_REV:.*]] = %[[C0]] to %[[C10]] step %[[C1]] iter_args(%[[DACC:.*]] = %[[ADDF_1]]) -> (f64) { +// CHECK: %[[IDX:.*]] = arith.subi %[[C9]], %[[IV_REV]] : index +// CHECK: %[[ADDF_2:.*]] = arith.addf %[[DACC]], %[[CST]] : f64 +// CHECK: %[[SHADOW:.*]] = memref.load %[[ALLOC]]{{\[}}%[[IDX]]] : memref<10xmemref<10xf64>> +// CHECK: %[[STOREIDX:.*]] = memref.load %[[ALLOC_1]]{{\[}}%[[IDX]]] : memref<10xindex> +// CHECK: %[[DVAL_0:.*]] = memref.load %[[SHADOW]]{{\[}}%[[STOREIDX]]] : memref<10xf64> +// CHECK: %[[ADDF_3:.*]] = arith.addf %[[ADDF_2]], %[[DVAL_0]] : f64 +// CHECK: memref.store %[[CST]], %[[SHADOW]]{{\[}}%[[STOREIDX]]] : memref<10xf64> +// CHECK: %[[ADDF_4:.*]] = arith.addf %[[ADDF_3]], %[[CST]] : f64 +// CHECK: %[[ADDF_5:.*]] = arith.addf %[[ADDF_3]], %[[CST]] : f64 +// CHECK: %[[LOADIDX:.*]] = memref.load %[[ALLOC_0]]{{\[}}%[[IDX]]] : memref<10xindex> +// CHECK: %[[DVAL_1:.*]] = memref.load %[[SHADOW]]{{\[}}%[[LOADIDX]]] : memref<10xf64> +// CHECK: %[[ADDF_6:.*]] = arith.addf %[[DVAL_1]], %[[ADDF_5]] : f64 +// The gradient contribution from the load is accumulated back into the +// shadow memref (%[[SHADOW]], i.e. %[[ARG1]]) in place. +// CHECK: memref.store %[[ADDF_6]], %[[SHADOW]]{{\[}}%[[LOADIDX]]] : memref<10xf64> +// CHECK: scf.yield %[[ADDF_4]] : f64 +// CHECK: } {enzyme.disable_mincut = true} +// CHECK: memref.dealloc %[[ALLOC_1]] : memref<10xindex> +// CHECK: memref.dealloc %[[ALLOC_0]] : memref<10xindex> +// CHECK: memref.dealloc %[[ALLOC]] : memref<10xmemref<10xf64>> +// CHECK: return +// CHECK: } From e0f08d5eaa66d6ed28a69849766fb0cc52393f0d Mon Sep 17 00:00:00 2001 From: Paul Berg <9824244+Pangoraw@users.noreply.github.com> Date: Mon, 27 Jul 2026 11:03:17 -0500 Subject: [PATCH 2/5] remove dyn dims stuff --- .../Implementations/SCFAutoDiffOpInterfaceImpl.cpp | 10 ++-------- 1 file changed, 2 insertions(+), 8 deletions(-) diff --git a/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp b/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp index 8dad493f9fdb..957927c0ca2e 100644 --- a/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp +++ b/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp @@ -227,10 +227,7 @@ 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)); - 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))); + sizes.push_back(b.getIndexAttr(rowTy.getDimSize(i))); strides.push_back(b.getIndexAttr(1)); } auto resTy = memref::SubViewOp::inferRankReducedResultType( @@ -257,10 +254,7 @@ struct ForOpInterfaceReverse if (auto mt = dyn_cast(valTy)) { Value row = checkpointRow(b, loc, buf, slot, mt); SmallVector 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); + Value fresh = memref::AllocOp::create(b, loc, mt); memref::CopyOp::create(b, loc, row, fresh); return fresh; } From 56efdc303c5c9a5e2809c1af69fe8aa1edde9ca8 Mon Sep 17 00:00:00 2001 From: Paul Berg <9824244+Pangoraw@users.noreply.github.com> Date: Mon, 27 Jul 2026 11:15:51 -0500 Subject: [PATCH 3/5] Revert "remove dyn dims stuff" This reverts commit e0f08d5eaa66d6ed28a69849766fb0cc52393f0d. --- .../Implementations/SCFAutoDiffOpInterfaceImpl.cpp | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp b/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp index 957927c0ca2e..8dad493f9fdb 100644 --- a/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp +++ b/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp @@ -227,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( @@ -254,7 +257,10 @@ struct ForOpInterfaceReverse if (auto mt = dyn_cast(valTy)) { Value row = checkpointRow(b, loc, buf, slot, mt); SmallVector dynSizes; - Value fresh = memref::AllocOp::create(b, loc, mt); + 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; } From 759e808698831363f3b6e01c3912122ae6a87968 Mon Sep 17 00:00:00 2001 From: Paul Berg <9824244+Pangoraw@users.noreply.github.com> Date: Mon, 27 Jul 2026 11:38:26 -0500 Subject: [PATCH 4/5] also fix inner iv rematerialization --- .../SCFAutoDiffOpInterfaceImpl.cpp | 7 ++----- .../scf_for_checkpointing_mutable_memory.mlir | 15 +++++++-------- 2 files changed, 9 insertions(+), 13 deletions(-) diff --git a/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp b/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp index 8dad493f9fdb..36c26ad6729e 100644 --- a/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp +++ b/enzyme/Enzyme/MLIR/Implementations/SCFAutoDiffOpInterfaceImpl.cpp @@ -1171,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] : diff --git a/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_mutable_memory.mlir b/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_mutable_memory.mlir index 51e863402fca..8996503e53ff 100644 --- a/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_mutable_memory.mlir +++ b/enzyme/test/MLIR/ReverseMode/scf_for_checkpointing_mutable_memory.mlir @@ -34,16 +34,15 @@ func.func @reduce_sum(%buf: memref<10xf64>) -> f64 { // CHECK-NEXT: %5 = arith.select %4, %c1, %c3 : index // CHECK-NEXT: %subview = memref.subview %alloc[%3, 0] [1, 10] [1, 1] : memref<4x10xf64> to memref<10xf64, strided<[1], offset: ?>> // CHECK-NEXT: memref.copy %arg0, %subview : memref<10xf64> to memref<10xf64, strided<[1], offset: ?>> -// CHECK-NEXT: %6 = arith.muli %arg3, %c3 : index -// CHECK-NEXT: %7 = scf.for %arg5 = %c0 to %5 step %c1 iter_args(%arg6 = %arg4) -> (f64) { -// CHECK-NEXT: %8 = arith.addi %6, %arg5 : index -// CHECK-NEXT: %9 = memref.load %arg0[%8] : memref<10xf64> -// CHECK-NEXT: %10 = arith.addf %arg6, %9 : f64 -// CHECK-NEXT: memref.store %10, %arg0[%c0] : memref<10xf64> -// CHECK-NEXT: scf.yield %10 : f64 +// CHECK-NEXT: %6 = scf.for %arg5 = %c0 to %5 step %c1 iter_args(%arg6 = %arg4) -> (f64) { +// CHECK-NEXT: %7 = arith.addi %arg3, %arg5 : index +// CHECK-NEXT: %8 = memref.load %arg0[%7] : memref<10xf64> +// CHECK-NEXT: %9 = arith.addf %arg6, %8 : f64 +// CHECK-NEXT: memref.store %9, %arg0[%c0] : memref<10xf64> +// CHECK-NEXT: scf.yield %9 : f64 // CHECK-NEXT: } {enzyme.disable_mincut = true} // CHECK-NEXT: memref.store %arg4, %alloc_0[%3] : memref<4xf64> -// CHECK-NEXT: scf.yield %7 : f64 +// CHECK-NEXT: scf.yield %6 : f64 // CHECK-NEXT: } // CHECK-NEXT: %1 = arith.addf %arg2, %cst : f64 // CHECK-NEXT: %2 = scf.for %arg3 = %c0 to %c4 step %c1 iter_args(%arg4 = %1) -> (f64) { From a9759bcdb32668063b4593a85f91c8097f85c84e Mon Sep 17 00:00:00 2001 From: Paul Berg <9824244+Pangoraw@users.noreply.github.com> Date: Mon, 27 Jul 2026 11:57:04 -0500 Subject: [PATCH 5/5] Add execution-based numerical tests for scf.for checkpointing with mutable memory The existing scf_for*mutable_memory* tests only FileCheck the generated IR shape; they can't catch numerical bugs like the checkpointing forward-loop induction variable double-scaling fixed in the previous commit. Add mlir-runner-executed variants (checkpointing off, uniform, and binomial) that JIT-run the differentiated function and check the printed primal/gradient values match across all three checkpointing strategies. This requires two small additions: - %mlir-opt/%mlir-runner/%mlir_runner_utils/%mlir_c_runner_utils lit substitutions, gated behind a `mlir-runner` feature so the new tests are skipped (not failed) wherever these tools aren't built. - Lowering enzyme.load/enzyme.store (dynamic-size-annotated memref load/store) in --convert-enzyme-to-memref, which previously only handled init/push/pop/get/set; the uniform checkpointing path emits load/store for its per-iteration caches, so this is needed to fully lower to executable IR. --- enzyme/Enzyme/MLIR/Passes/EnzymeToMemRef.cpp | 30 ++++++++++++++++ .../test/MLIR/Integration/ReverseMode/BUILD | 30 ++++++++++++++++ .../Inputs/exec_main_10xf64.mlir.inc | 16 +++++++++ ...pointing_binomial_mutable_memory_exec.mlir | 30 ++++++++++++++++ ...for_checkpointing_mutable_memory_exec.mlir | 35 +++++++++++++++++++ .../scf_for_mutable_memory_exec.mlir | 33 +++++++++++++++++ enzyme/test/lit.site.cfg.py.in | 27 ++++++++++++++ 7 files changed, 201 insertions(+) create mode 100644 enzyme/test/MLIR/Integration/ReverseMode/BUILD create mode 100644 enzyme/test/MLIR/Integration/ReverseMode/Inputs/exec_main_10xf64.mlir.inc create mode 100644 enzyme/test/MLIR/Integration/ReverseMode/scf_for_checkpointing_binomial_mutable_memory_exec.mlir create mode 100644 enzyme/test/MLIR/Integration/ReverseMode/scf_for_checkpointing_mutable_memory_exec.mlir create mode 100644 enzyme/test/MLIR/Integration/ReverseMode/scf_for_mutable_memory_exec.mlir diff --git a/enzyme/Enzyme/MLIR/Passes/EnzymeToMemRef.cpp b/enzyme/Enzyme/MLIR/Passes/EnzymeToMemRef.cpp index 471a033505e1..66290ebcf91e 100644 --- a/enzyme/Enzyme/MLIR/Passes/EnzymeToMemRef.cpp +++ b/enzyme/Enzyme/MLIR/Passes/EnzymeToMemRef.cpp @@ -338,6 +338,34 @@ struct SetOpConversion : public OpConversionPattern { } }; +// `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 { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(enzyme::LoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, adaptor.getMemref(), + adaptor.getIndices()); + return success(); + } +}; + +struct StoreOpConversion : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult + matchAndRewrite(enzyme::StoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp( + op, adaptor.getValue(), adaptor.getMemref(), adaptor.getIndices()); + return success(); + } +}; + struct GetOpConversion : public OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -403,6 +431,8 @@ struct EnzymeToMemRefPass patterns.add(typeConverter, context); patterns.add(typeConverter, context); patterns.add(typeConverter, context); + patterns.add(typeConverter, context); + patterns.add(typeConverter, context); ConversionTarget target(*context); target.addLegalDialect(); diff --git a/enzyme/test/MLIR/Integration/ReverseMode/BUILD b/enzyme/test/MLIR/Integration/ReverseMode/BUILD new file mode 100644 index 000000000000..e602063d241e --- /dev/null +++ b/enzyme/test/MLIR/Integration/ReverseMode/BUILD @@ -0,0 +1,30 @@ +package( + default_applicable_licenses = [], + default_visibility = ["//visibility:public"], +) + +# Execution-based numerical tests for scf.for reverse-mode checkpointing with +# mutable memory: unlike the FileCheck-only IR-shape tests in +# //test/MLIR/ReverseMode, these actually JIT-run the differentiated function +# via mlir-runner and check the printed numerical result. + +load("@llvm-project//llvm:lit_test.bzl", "lit_test") + +[ + lit_test( + name = "%s.test" % src, + srcs = [src], + data = [ + "Inputs/exec_main_10xf64.mlir.inc", + "//:enzymemlir-opt", + "//test:lit.cfg.py", + "//test:lit.site.cfg.py", + "@llvm-project//llvm:FileCheck", + "@llvm-project//mlir:libmlir_c_runner_utils.so", + "@llvm-project//mlir:libmlir_runner_utils.so", + "@llvm-project//mlir:mlir-opt", + "@llvm-project//mlir:mlir-runner", + ], + ) + for src in glob(["*.mlir"]) +] diff --git a/enzyme/test/MLIR/Integration/ReverseMode/Inputs/exec_main_10xf64.mlir.inc b/enzyme/test/MLIR/Integration/ReverseMode/Inputs/exec_main_10xf64.mlir.inc new file mode 100644 index 000000000000..890a769ecc11 --- /dev/null +++ b/enzyme/test/MLIR/Integration/ReverseMode/Inputs/exec_main_10xf64.mlir.inc @@ -0,0 +1,16 @@ +memref.global "private" @__buf : memref<10xf64> = dense<[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0]> +memref.global "private" @__dbuf : memref<10xf64> = dense<0.0> + +func.func @main() { + %buf = memref.get_global @__buf : memref<10xf64> + %dbuf = memref.get_global @__dbuf : memref<10xf64> + %seed = arith.constant 1.0 : f64 + call @reduce_sum(%buf, %dbuf, %seed) : (memref<10xf64>, memref<10xf64>, f64) -> () + %ubuf = memref.cast %buf : memref<10xf64> to memref<*xf64> + %udbuf = memref.cast %dbuf : memref<10xf64> to memref<*xf64> + call @printMemrefF64(%ubuf) : (memref<*xf64>) -> () + call @printMemrefF64(%udbuf) : (memref<*xf64>) -> () + return +} + +func.func private @printMemrefF64(%ptr : memref<*xf64>) diff --git a/enzyme/test/MLIR/Integration/ReverseMode/scf_for_checkpointing_binomial_mutable_memory_exec.mlir b/enzyme/test/MLIR/Integration/ReverseMode/scf_for_checkpointing_binomial_mutable_memory_exec.mlir new file mode 100644 index 000000000000..f9e1162e804b --- /dev/null +++ b/enzyme/test/MLIR/Integration/ReverseMode/scf_for_checkpointing_binomial_mutable_memory_exec.mlir @@ -0,0 +1,30 @@ +// REQUIRES: mlir-runner +// +// Same numeric check as scf_for_checkpointing_mutable_memory_exec.mlir, but +// for binomial (Revolve) checkpointing. +// +// RUN: (%eopt %s --enzyme-wrap="infn=reduce_sum outfn= argTys=enzyme_dup retTys=enzyme_active mode=ReverseModeCombined" --canonicalize --enzyme-simplify-math --remove-unnecessary-enzyme-ops --canonicalize --lower-enzyme-binomial-progress --convert-enzyme-to-memref | tail -n +2 | head -n -2; cat %S/Inputs/exec_main_10xf64.mlir.inc) | %mlir-opt --convert-scf-to-cf --expand-strided-metadata --lower-affine --convert-arith-to-llvm --finalize-memref-to-llvm --convert-cf-to-llvm --convert-func-to-llvm --reconcile-unrealized-casts | %mlir-runner -e main -entry-point-result=void -shared-libs=%mlir_runner_utils,%mlir_c_runner_utils | FileCheck %s + +func.func @reduce_sum(%buf: memref<10xf64>) -> f64 { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c10 = arith.constant 10 : index + %init = arith.constant 0.0 : f64 + + %sum = scf.for %i = %c0 to %c10 step %c1 iter_args(%acc = %init) -> (f64) { + %val = memref.load %buf[%i] : memref<10xf64> + %new_acc = arith.addf %acc, %val : f64 + memref.store %new_acc, %buf[%c0] : memref<10xf64> + scf.yield %new_acc : f64 + } {enzyme.enable_checkpointing = true, + enzyme.binomial_checkpointing, + enzyme.checkpoint_period=4, + enzyme.disable_mincut=true} + + return %sum : f64 +} + +// CHECK: Unranked Memref base@ = {{0x[0-9a-f]*}} rank = 1 offset = 0 sizes = [10] strides = [1] data = +// CHECK-NEXT: [55, 2, 3, 4, 5, 6, 7, 8, 9, 10] +// CHECK: Unranked Memref base@ = {{0x[0-9a-f]*}} rank = 1 offset = 0 sizes = [10] strides = [1] data = +// CHECK-NEXT: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1] diff --git a/enzyme/test/MLIR/Integration/ReverseMode/scf_for_checkpointing_mutable_memory_exec.mlir b/enzyme/test/MLIR/Integration/ReverseMode/scf_for_checkpointing_mutable_memory_exec.mlir new file mode 100644 index 000000000000..d996f69ff6bc --- /dev/null +++ b/enzyme/test/MLIR/Integration/ReverseMode/scf_for_checkpointing_mutable_memory_exec.mlir @@ -0,0 +1,35 @@ +// REQUIRES: mlir-runner +// +// Numeric regression test for the uniform (sqrt-chunked) checkpointing +// forward-loop index bug: the forward chunk-caching loop's induction +// variable was double-scaled by the chunk size (nInner), so for a +// checkpoint_period that didn't evenly divide the 10-element buffer, the +// recomputed primal accessed the buffer out of bounds. Compare against +// scf_for_mutable_memory_exec.mlir (enable_checkpointing=false): both must +// produce the exact same primal/gradient result, since checkpointing is +// purely a memory/recompute strategy and must not change the numerics. +// +// RUN: (%eopt %s --enzyme-wrap="infn=reduce_sum outfn= argTys=enzyme_dup retTys=enzyme_active mode=ReverseModeCombined" --canonicalize --enzyme-simplify-math --remove-unnecessary-enzyme-ops --convert-enzyme-to-memref | tail -n +2 | head -n -2; cat %S/Inputs/exec_main_10xf64.mlir.inc) | %mlir-opt --convert-scf-to-cf --expand-strided-metadata --lower-affine --convert-arith-to-llvm --finalize-memref-to-llvm --convert-cf-to-llvm --convert-func-to-llvm --reconcile-unrealized-casts | %mlir-runner -e main -entry-point-result=void -shared-libs=%mlir_runner_utils,%mlir_c_runner_utils | FileCheck %s + +func.func @reduce_sum(%buf: memref<10xf64>) -> f64 { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c10 = arith.constant 10 : index + %init = arith.constant 0.0 : f64 + + %sum = scf.for %i = %c0 to %c10 step %c1 iter_args(%acc = %init) -> (f64) { + %val = memref.load %buf[%i] : memref<10xf64> + %new_acc = arith.addf %acc, %val : f64 + memref.store %new_acc, %buf[%c0] : memref<10xf64> + scf.yield %new_acc : f64 + } {enzyme.enable_checkpointing = true, + enzyme.checkpoint_period=4, + enzyme.disable_mincut=true} + + return %sum : f64 +} + +// CHECK: Unranked Memref base@ = {{0x[0-9a-f]*}} rank = 1 offset = 0 sizes = [10] strides = [1] data = +// CHECK-NEXT: [55, 2, 3, 4, 5, 6, 7, 8, 9, 10] +// CHECK: Unranked Memref base@ = {{0x[0-9a-f]*}} rank = 1 offset = 0 sizes = [10] strides = [1] data = +// CHECK-NEXT: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1] diff --git a/enzyme/test/MLIR/Integration/ReverseMode/scf_for_mutable_memory_exec.mlir b/enzyme/test/MLIR/Integration/ReverseMode/scf_for_mutable_memory_exec.mlir new file mode 100644 index 000000000000..4cbc70cd6847 --- /dev/null +++ b/enzyme/test/MLIR/Integration/ReverseMode/scf_for_mutable_memory_exec.mlir @@ -0,0 +1,33 @@ +// REQUIRES: mlir-runner +// +// Executes the reverse-mode derivative of a scf.for loop whose body both +// reads and writes the same memref (see ../../ReverseMode/scf_for_mutable_memory.mlir +// for the FileCheck-only IR-shape test), then checks the actual numerical +// result: the primal `buf` should end with the same final mutation trace as +// running the loop directly, and the gradient of `sum` w.r.t. `buf` should be +// all-ones (since the mutation of buf[0] is never re-read, `sum` is just a +// sum of the original buf entries). +// +// RUN: (%eopt %s --enzyme-wrap="infn=reduce_sum outfn= argTys=enzyme_dup retTys=enzyme_active mode=ReverseModeCombined" --canonicalize --enzyme-simplify-math --remove-unnecessary-enzyme-ops | tail -n +2 | head -n -2; cat %S/Inputs/exec_main_10xf64.mlir.inc) | %mlir-opt --convert-scf-to-cf --expand-strided-metadata --lower-affine --convert-arith-to-llvm --finalize-memref-to-llvm --convert-cf-to-llvm --convert-func-to-llvm --reconcile-unrealized-casts | %mlir-runner -e main -entry-point-result=void -shared-libs=%mlir_runner_utils,%mlir_c_runner_utils | FileCheck %s + +func.func @reduce_sum(%buf: memref<10xf64>) -> f64 { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c10 = arith.constant 10 : index + %init = arith.constant 0.0 : f64 + + %sum = scf.for %i = %c0 to %c10 step %c1 iter_args(%acc = %init) -> (f64) { + %val = memref.load %buf[%i] : memref<10xf64> + %new_acc = arith.addf %acc, %val : f64 + memref.store %new_acc, %buf[%c0] : memref<10xf64> + scf.yield %new_acc : f64 + } {enzyme.enable_checkpointing = false, + enzyme.checkpoint_period=4, enzyme.disable_mincut=true} + + return %sum : f64 +} + +// CHECK: Unranked Memref base@ = {{0x[0-9a-f]*}} rank = 1 offset = 0 sizes = [10] strides = [1] data = +// CHECK-NEXT: [55, 2, 3, 4, 5, 6, 7, 8, 9, 10] +// CHECK: Unranked Memref base@ = {{0x[0-9a-f]*}} rank = 1 offset = 0 sizes = [10] strides = [1] data = +// CHECK-NEXT: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1] diff --git a/enzyme/test/lit.site.cfg.py.in b/enzyme/test/lit.site.cfg.py.in index bb7d494764f3..2ec8c9780887 100644 --- a/enzyme/test/lit.site.cfg.py.in +++ b/enzyme/test/lit.site.cfg.py.in @@ -66,6 +66,33 @@ if len("@ENZYME_BINARY_DIR@") == 0: eclang += "-I " + os.path.dirname(os.path.abspath(__file__)) + "/Integration" config.substitutions.append(('%eopt', emopt)) + +# mlir-opt/mlir-runner and the runner support libs live alongside the LLVM +# tools/libs when LLVM is built with MLIR enabled (the CMake case here, since +# this project builds against a preexisting LLVM+MLIR install/build). Under +# Bazel, @llvm-project//mlir data deps land as a sibling of this file's own +# runfiles directory (runfiles_root/llvm-project/mlir/...), same as how +# %eopt's fallback above locates //:enzymemlir-opt relative to __file__. +mlir_tools_dir = config.llvm_tools_dir +mlir_libs_dir = config.llvm_libs_dir +if len("@ENZYME_BINARY_DIR@") == 0: + mlir_tools_dir = (os.path.dirname(os.path.abspath(__file__)) + + "/../../llvm-project/mlir") + mlir_libs_dir = mlir_tools_dir + +mlir_opt = mlir_tools_dir + "/mlir-opt" +mlir_runner = mlir_tools_dir + "/mlir-runner" +mlir_runner_utils = mlir_libs_dir + "/libmlir_runner_utils" + config.llvm_shlib_ext +mlir_c_runner_utils = mlir_libs_dir + "/libmlir_c_runner_utils" + config.llvm_shlib_ext + +if os.path.exists(mlir_runner): + config.available_features.add('mlir-runner') + +config.substitutions.append(('%mlir-opt', mlir_opt)) +config.substitutions.append(('%mlir-runner', mlir_runner)) +config.substitutions.append(('%mlir_runner_utils', mlir_runner_utils)) +config.substitutions.append(('%mlir_c_runner_utils', mlir_c_runner_utils)) + config.substitutions.append(('%llvmver', config.llvm_ver)) config.substitutions.append(('%FileCheck', config.llvm_tools_dir + "/FileCheck")) config.substitutions.append(('%clang', eclang))