diff --git a/mlir/lib/Dialect/Bufferization/Transforms/StaticMemoryPlannerAnalysis.cpp b/mlir/lib/Dialect/Bufferization/Transforms/StaticMemoryPlannerAnalysis.cpp index c2ac40a8427e7..f1db0d5ee55a4 100644 --- a/mlir/lib/Dialect/Bufferization/Transforms/StaticMemoryPlannerAnalysis.cpp +++ b/mlir/lib/Dialect/Bufferization/Transforms/StaticMemoryPlannerAnalysis.cpp @@ -13,11 +13,13 @@ //===----------------------------------------------------------------------===// #include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Bufferization/IR/BufferViewFlowOpInterface.h" #include "mlir/Dialect/Bufferization/Transforms/Passes.h" #include "mlir/Dialect/Bufferization/Transforms/StaticMemoryPlanning.h" #include "mlir/Dialect/MemRef/IR/MemRef.h" #include "mlir/IR/Builders.h" #include "mlir/Interfaces/FunctionInterfaces.h" +#include "llvm/ADT/SmallPtrSet.h" #include "llvm/Support/Debug.h" #include @@ -34,10 +36,12 @@ using namespace mlir; namespace { -/// A candidate allocation with its matching deallocation and assigned offset. +/// A candidate allocation with its matching deallocation(s) and assigned +/// offset. An alloc may be freed indirectly through arith.select chains, +/// yielding multiple potential deallocs — all must be in the same block. struct AllocationCandidate { memref::AllocOp alloc; - memref::DeallocOp dealloc; + SmallVector deallocs; int64_t offset = 0; // Offset in bytes from arena start (assigned by planner) int64_t sizeInBytes = 0; // Size in bytes int64_t alignment = 1; // Required alignment in bytes @@ -47,18 +51,31 @@ struct AllocationCandidate { // Helper utilities //===----------------------------------------------------------------------===// -/// Finds the unique dealloc operation for a given alloc value. -/// Returns nullptr if there are zero or multiple deallocs. -static memref::DeallocOp findUniqueDealloc(Value allocValue) { - memref::DeallocOp deallocOp = nullptr; - for (Operation *user : allocValue.getUsers()) { +/// Collect all dealloc ops that might free the given value, following ops +/// that implement BufferViewFlowOpInterface (e.g. arith.select). For example: +/// %0 = memref.alloc() +/// %2 = arith.select %c, %0, %1 +/// memref.dealloc %2 <- this covers %0 conditionally +/// `visited` prevents cycles in the use-def graph. +/// +/// TODO: This relies on BufferViewFlowOpInterface external models being +/// registered for the ops in the IR (e.g. via +/// arith::registerBufferViewFlowOpInterfaceExternalModels). +static void findPotentialDeallocs(Value value, + SmallVectorImpl &deallocs, + SmallPtrSetImpl &visited) { + if (!visited.insert(value).second) + return; + for (Operation *user : value.getUsers()) { if (auto dealloc = dyn_cast(user)) { - if (deallocOp) - return nullptr; // Multiple deallocs found - deallocOp = dealloc; + deallocs.push_back(dealloc); + } else if (dyn_cast(user)) { + // Follow any op that propagates buffer values to its results. + for (Value result : user->getResults()) + if (isa(result.getType())) + findPotentialDeallocs(result, deallocs, visited); } } - return deallocOp; } /// Compute the size in bytes for a memref type. @@ -70,25 +87,31 @@ static int64_t computeSizeInBytes(MemRefType memrefType) { /// Build lifetime-annotated allocation descriptors from candidates. /// Returns the arena alignment (LCM of all individual alignments). +/// Uses a single block scan (O(n+m)) instead of one scan per candidate. static int64_t buildAllocInfos( MutableArrayRef candidates, SmallVectorImpl &allocInfos) { + // Build an op-index map with a single pass over the block. + DenseMap opIndex; + if (!candidates.empty()) { + Block *block = candidates.front().alloc->getBlock(); + int64_t idx = 0; + for (Operation &op : *block) + opIndex[&op] = idx++; + } + int64_t arenaAlignment = 1; for (auto &candidate : candidates) { bufferization::MemoryPlannerAlloc info; info.sizeInBytes = candidate.sizeInBytes; info.alignment = candidate.alignment; - - Block *block = candidate.alloc->getBlock(); - int64_t opIdx = 0; - for (Operation &op : *block) { - if (&op == candidate.alloc.getOperation()) - info.timeStart = opIdx; - if (&op == candidate.dealloc.getOperation()) - info.timeEnd = opIdx; - ++opIdx; - } - + info.timeStart = opIndex.lookup(candidate.alloc.getOperation()); + // Conservative: timeEnd = latest dealloc index among all potential + // deallocs. + int64_t timeEnd = 0; + for (memref::DeallocOp d : candidate.deallocs) + timeEnd = std::max(timeEnd, opIndex.lookup(d.getOperation())); + info.timeEnd = timeEnd; allocInfos.push_back(info); arenaAlignment = std::lcm(arenaAlignment, candidate.alignment); } @@ -96,46 +119,53 @@ static int64_t buildAllocInfos( } /// Collect alloc/dealloc pairs eligible for arena placement. -/// An allocation is eligible if it has a static shape and a unique dealloc -/// in the same block. -static SmallVector +/// An allocation is eligible if it has a static shape and its deallocs +/// (including those reached via BufferViewFlowOpInterface chains) +/// are in the same block. Allocations with dynamic shapes are skipped. +/// Missing deallocs or cross-block deallocs are reported as errors. +static LogicalResult collectCandidates(FunctionOpInterface funcOp, llvm::Statistic &numSkipDynamic, - llvm::Statistic &numSkipNoDealloc, - llvm::Statistic &numEligible) { - SmallVector candidates; - - funcOp->walk([&](memref::AllocOp allocOp) { + llvm::Statistic &numEligible, + SmallVector &candidates) { + bool walkFailed = false; + funcOp->walk([&](memref::AllocOp allocOp) -> WalkResult { MemRefType memrefType = allocOp.getType(); - - // Skip dynamic shapes if (!memrefType.hasStaticShape()) { ++numSkipDynamic; - return; + return WalkResult::advance(); } - // Find unique dealloc in the same block - memref::DeallocOp deallocOp = findUniqueDealloc(allocOp.getResult()); - if (!deallocOp) { - ++numSkipNoDealloc; - return; + SmallVector deallocs; + SmallPtrSet visited; + findPotentialDeallocs(allocOp.getResult(), deallocs, visited); + + if (deallocs.empty()) { + allocOp.emitError("no dealloc found; run the deallocation pipeline " + "before this pass"); + walkFailed = true; + return WalkResult::interrupt(); } - if (deallocOp->getBlock() != allocOp->getBlock()) { - ++numSkipNoDealloc; - return; + for (memref::DeallocOp d : deallocs) { + if (d->getBlock() != allocOp->getBlock()) { + allocOp.emitError("dealloc is in a different block than the alloc; " + "run the deallocation pipeline before this pass"); + walkFailed = true; + return WalkResult::interrupt(); + } } - // This allocation is eligible ++numEligible; AllocationCandidate candidate; candidate.alloc = allocOp; - candidate.dealloc = deallocOp; + candidate.deallocs = deallocs; candidate.sizeInBytes = computeSizeInBytes(memrefType); candidate.alignment = allocOp.getAlignment().value_or(1); candidates.push_back(candidate); + return WalkResult::advance(); }); - return candidates; + return failure(walkFailed); } /// Create or obtain the arena buffer based on the arena mode. @@ -182,6 +212,10 @@ static FailureOr createArena(OpBuilder &builder, /// Replace each alloc/dealloc pair with a memref.view into the arena. static void rewriteAllocations(MutableArrayRef candidates, Value arenaValue) { + SmallPtrSet deallocsToErase; + SmallVector allocsToErase; + + // Replace all alloc results with views (rewires selects too). for (auto &candidate : candidates) { OpBuilder builder(candidate.alloc); Location loc = candidate.alloc.getLoc(); @@ -191,11 +225,20 @@ static void rewriteAllocations(MutableArrayRef candidates, arith::ConstantIndexOp::create(builder, loc, candidate.offset); auto view = memref::ViewOp::create(builder, loc, originalType, arenaValue, offsetIndex, SmallVector{}); - candidate.alloc.getResult().replaceAllUsesWith(view.getResult()); - candidate.alloc.erase(); - candidate.dealloc.erase(); + allocsToErase.push_back(candidate.alloc.getOperation()); + + for (memref::DeallocOp d : candidate.deallocs) + deallocsToErase.insert(d.getOperation()); } + + // Erase deallocs first (they may reference alloc results via selects). + for (Operation *d : deallocsToErase) + d->erase(); + + // Erase allocs last (no users remain after replaceAllUsesWith). + for (Operation *allocOp : allocsToErase) + allocOp->erase(); } //===----------------------------------------------------------------------===// @@ -226,8 +269,10 @@ void StaticMemoryPlannerAnalysisPass::runOnOperation() { } // Step 1: Collect eligible allocation candidates. - SmallVector candidates = - collectCandidates(funcOp, numSkipDynamic, numSkipNoDealloc, numEligible); + SmallVector candidates; + if (failed( + collectCandidates(funcOp, numSkipDynamic, numEligible, candidates))) + return signalPassFailure(); if (candidates.empty()) return; diff --git a/mlir/test/Dialect/Bufferization/Transforms/static-memory-planner-analysis.mlir b/mlir/test/Dialect/Bufferization/Transforms/static-memory-planner-analysis.mlir index a80c0e13adc21..af900be29a22e 100644 --- a/mlir/test/Dialect/Bufferization/Transforms/static-memory-planner-analysis.mlir +++ b/mlir/test/Dialect/Bufferization/Transforms/static-memory-planner-analysis.mlir @@ -85,35 +85,7 @@ func.func @dynamic_shape_skipped(%n: index) { // ----- -// Test 5: No dealloc - should be skipped -// CHECK-LABEL: func @no_dealloc_skipped -func.func @no_dealloc_skipped() { - // CHECK: %[[ALLOC:.*]] = memref.alloc() : memref<1024xf32> - // CHECK-NOT: memref.subview - %alloc = memref.alloc() : memref<1024xf32> - return -} - -// ----- - -// Test 6: Dealloc in different block - should be skipped -// CHECK-LABEL: func @different_block_skipped -func.func @different_block_skipped(%cond: i1) { - // CHECK: %[[ALLOC:.*]] = memref.alloc() : memref<1024xf32> - // CHECK: scf.if - // CHECK: memref.dealloc %[[ALLOC]] - // CHECK-NOT: memref.subview - %alloc = memref.alloc() : memref<1024xf32> - scf.if %cond { - memref.dealloc %alloc : memref<1024xf32> - scf.yield - } - return -} - -// ----- - -// Test 7: Multiple allocations with sequential offsets +// Test 5: Multiple allocations with sequential offsets // CHECK-LABEL: func @multiple_sequential func.func @multiple_sequential() { // Arena: 1024*4 + 512*4 + 2048*4 = 14336 bytes @@ -140,7 +112,7 @@ func.func @multiple_sequential() { // ----- -// Test 8: Alignment requirements with padding +// Test 6: Alignment requirements with padding // CHECK-LABEL: func @alignment_padding func.func @alignment_padding() { // Arena: 256*4 + 128*4 + 64*4 = 1792 bytes, alignment = lcm(128,64,128) = 128 @@ -167,7 +139,7 @@ func.func @alignment_padding() { // ----- -// Test 9: LCM arena alignment (alignment=4, alignment=16 → lcm=16). +// Test 7: LCM arena alignment (alignment=4, alignment=16 → lcm=16). // For power-of-2 alignments lcm equals max, but lcm is the correct // general formula. Arena must be aligned to 16 so that all views are // correctly aligned regardless of their individual requirements. @@ -190,3 +162,64 @@ func.func @lcm_alignment() { memref.dealloc %alloc1 : memref<3xi32> return } + +// ----- + +// Test 8: Single alloc freed via arith.select-based dealloc. +// CHECK-LABEL: func @select_single_alloc +func.func @select_single_alloc() { + %c = arith.constant true + // CHECK: %[[ARENA:.*]] = memref.alloc() {alignment = 1 : i64} : memref<4096xi8> + // CHECK-NEXT: %[[C0:.*]] = arith.constant 0 : index + // CHECK-NEXT: %[[V:.*]] = memref.view %[[ARENA]][%[[C0]]][] : memref<4096xi8> to memref<1024xf32> + // CHECK-NOT: memref.alloc + // CHECK-NOT: memref.dealloc + %alloc = memref.alloc() : memref<1024xf32> + %sel = arith.select %c, %alloc, %alloc : memref<1024xf32> + memref.dealloc %sel : memref<1024xf32> + return +} + +// ----- + +// Test 9: Two allocs freed via a shared select-based dealloc. +// Group constraint: both must be eligible together or neither is. +// CHECK-LABEL: func @select_shared_dealloc +func.func @select_shared_dealloc() { + %c = arith.constant true + // CHECK: %[[ARENA:.*]] = memref.alloc() {alignment = 1 : i64} : memref<8192xi8> + // CHECK-NEXT: %[[C0:.*]] = arith.constant 0 : index + // CHECK-NEXT: %[[V0:.*]] = memref.view %[[ARENA]][%[[C0]]][] : memref<8192xi8> to memref<1024xf32> + // CHECK-NEXT: %[[C4096:.*]] = arith.constant 4096 : index + // CHECK-NEXT: %[[V1:.*]] = memref.view %[[ARENA]][%[[C4096]]][] : memref<8192xi8> to memref<1024xf32> + // CHECK-NOT: memref.alloc + // CHECK-NOT: memref.dealloc + %a = memref.alloc() : memref<1024xf32> + %b = memref.alloc() : memref<1024xf32> + %sel = arith.select %c, %a, %b : memref<1024xf32> + memref.dealloc %sel : memref<1024xf32> + return +} + +// ----- + +// Test 10: Two allocs, two select-based deallocs (mentor's canonical example). +// %a freed via dealloc(%sel1) or dealloc(%sel2), %b likewise. +// CHECK-LABEL: func @select_two_deallocs +func.func @select_two_deallocs() { + %c = arith.constant true + // CHECK: %[[ARENA:.*]] = memref.alloc() {alignment = 1 : i64} : memref<8192xi8> + // CHECK-NEXT: %[[C0:.*]] = arith.constant 0 : index + // CHECK-NEXT: %{{.*}} = memref.view %[[ARENA]][%[[C0]]][] : memref<8192xi8> to memref<1024xf32> + // CHECK-NEXT: %[[C4096:.*]] = arith.constant 4096 : index + // CHECK-NEXT: %{{.*}} = memref.view %[[ARENA]][%[[C4096]]][] : memref<8192xi8> to memref<1024xf32> + // CHECK-NOT: memref.alloc + // CHECK-NOT: memref.dealloc + %a = memref.alloc() : memref<1024xf32> + %b = memref.alloc() : memref<1024xf32> + %sel1 = arith.select %c, %a, %b : memref<1024xf32> + memref.dealloc %sel1 : memref<1024xf32> + %sel2 = arith.select %c, %b, %a : memref<1024xf32> + memref.dealloc %sel2 : memref<1024xf32> + return +} diff --git a/mlir/test/Dialect/Bufferization/Transforms/static-memory-planner-errors.mlir b/mlir/test/Dialect/Bufferization/Transforms/static-memory-planner-errors.mlir new file mode 100644 index 0000000000000..86af0649c67ba --- /dev/null +++ b/mlir/test/Dialect/Bufferization/Transforms/static-memory-planner-errors.mlir @@ -0,0 +1,24 @@ +// RUN: mlir-opt %s -pass-pipeline="builtin.module(func.func(static-memory-planner-analysis))" \ +// RUN: -split-input-file -verify-diagnostics + +// ----- + +// Test 1: Alloc with no dealloc should be an error (not silently skipped). +func.func @error_no_dealloc() { + // expected-error @+1 {{no dealloc found; run the deallocation pipeline before this pass}} + %alloc = memref.alloc() : memref<1024xf32> + return +} + +// ----- + +// Test 2: Alloc whose dealloc is in a different block should be an error. +func.func @error_cross_block_dealloc(%cond: i1) { + // expected-error @+1 {{dealloc is in a different block than the alloc; run the deallocation pipeline before this pass}} + %alloc = memref.alloc() : memref<1024xf32> + scf.if %cond { + memref.dealloc %alloc : memref<1024xf32> + scf.yield + } + return +}