#include "flang/Optimizer/Dialect/FIRDialect.h"
#include "flang/Optimizer/Dialect/FIROps.h"
#include "flang/Optimizer/Dialect/FIRType.h"
#include "flang/Optimizer/Transforms/Passes.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/IR/Diagnostics.h"
#include "mlir/Pass/Pass.h"
#include "mlir/Transforms/DialectConversion.h"
#include "mlir/Transforms/Passes.h"
#include "llvm/ADT/TypeSwitch.h"
namespace fir {
#define GEN_PASS_DEF_MEMORYALLOCATIONOPT
#include "flang/Optimizer/Transforms/Passes.h.inc"
}
#define DEBUG_TYPE "flang-memory-allocation-opt"
static constexpr std::size_t unlimitedArraySize = ~static_cast<std::size_t>(0);
namespace {
struct MemoryAllocationOptions {
bool dynamicArrayOnHeap = false;
std::size_t maxStackArraySize = unlimitedArraySize;
};
class ReturnAnalysis {
public:
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(ReturnAnalysis)
ReturnAnalysis(mlir::Operation *op) {
if (auto func = mlir::dyn_cast<mlir::func::FuncOp>(op))
for (mlir::Block &block : func)
for (mlir::Operation &i : block)
if (mlir::isa<mlir::func::ReturnOp>(i)) {
returnMap[op].push_back(&i);
break;
}
}
llvm::SmallVector<mlir::Operation *> getReturns(mlir::Operation *func) const {
auto iter = returnMap.find(func);
if (iter != returnMap.end())
return iter->second;
return {};
}
private:
llvm::DenseMap<mlir::Operation *, llvm::SmallVector<mlir::Operation *>>
returnMap;
};
}
static inline bool keepStackAllocation(fir::AllocaOp alloca, mlir::Block *entry,
const MemoryAllocationOptions &options) {
if (alloca->getBlock() != entry)
return true;
if (auto seqTy = alloca.getInType().dyn_cast<fir::SequenceType>()) {
if (fir::hasDynamicSize(seqTy)) {
if (options.dynamicArrayOnHeap)
return false;
} else {
std::int64_t numberOfElements = 1;
for (std::int64_t i : seqTy.getShape()) {
numberOfElements *= i;
if (numberOfElements <= 0)
return true;
}
if (static_cast<std::size_t>(numberOfElements) >
options.maxStackArraySize) {
LLVM_DEBUG(llvm::dbgs()
<< "memory allocation opt: found " << alloca << '\n');
return false;
}
}
}
return true;
}
namespace {
class AllocaOpConversion : public mlir::OpRewritePattern<fir::AllocaOp> {
public:
using OpRewritePattern::OpRewritePattern;
AllocaOpConversion(mlir::MLIRContext *ctx,
llvm::ArrayRef<mlir::Operation *> rets)
: OpRewritePattern(ctx), returnOps(rets) {}
mlir::LogicalResult
matchAndRewrite(fir::AllocaOp alloca,
mlir::PatternRewriter &rewriter) const override {
auto loc = alloca.getLoc();
mlir::Type varTy = alloca.getInType();
auto unpackName =
[](std::optional<llvm::StringRef> opt) -> llvm::StringRef {
if (opt)
return *opt;
return {};
};
auto uniqName = unpackName(alloca.getUniqName());
auto bindcName = unpackName(alloca.getBindcName());
auto heap = rewriter.create<fir::AllocMemOp>(
loc, varTy, uniqName, bindcName, alloca.getTypeparams(),
alloca.getShape());
auto insPt = rewriter.saveInsertionPoint();
for (mlir::Operation *retOp : returnOps) {
rewriter.setInsertionPoint(retOp);
[[maybe_unused]] auto free = rewriter.create<fir::FreeMemOp>(loc, heap);
LLVM_DEBUG(llvm::dbgs() << "memory allocation opt: add free " << free
<< " for " << heap << '\n');
}
rewriter.restoreInsertionPoint(insPt);
rewriter.replaceOpWithNewOp<fir::ConvertOp>(
alloca, fir::ReferenceType::get(varTy), heap);
LLVM_DEBUG(llvm::dbgs() << "memory allocation opt: replaced " << alloca
<< " with " << heap << '\n');
return mlir::success();
}
private:
llvm::ArrayRef<mlir::Operation *> returnOps;
};
class MemoryAllocationOpt
: public fir::impl::MemoryAllocationOptBase<MemoryAllocationOpt> {
public:
MemoryAllocationOpt() {
options = {dynamicArrayOnHeap, maxStackArraySize};
}
MemoryAllocationOpt(bool dynOnHeap, std::size_t maxStackSize) {
options = {dynOnHeap, maxStackSize};
}
inline void useCommandLineOptions() {
if (dynamicArrayOnHeap)
options.dynamicArrayOnHeap = dynamicArrayOnHeap;
if (maxStackArraySize != unlimitedArraySize)
options.maxStackArraySize = maxStackArraySize;
}
void runOnOperation() override {
auto *context = &getContext();
auto func = getOperation();
mlir::RewritePatternSet patterns(context);
mlir::ConversionTarget target(*context);
useCommandLineOptions();
LLVM_DEBUG(llvm::dbgs()
<< "dynamic arrays on heap: " << options.dynamicArrayOnHeap
<< "\nmaximum number of elements of array on stack: "
<< options.maxStackArraySize << '\n');
if (func.empty())
return;
const auto &analysis = getAnalysis<ReturnAnalysis>();
target.addLegalDialect<fir::FIROpsDialect, mlir::arith::ArithDialect,
mlir::func::FuncDialect>();
target.addDynamicallyLegalOp<fir::AllocaOp>([&](fir::AllocaOp alloca) {
return keepStackAllocation(alloca, &func.front(), options);
});
patterns.insert<AllocaOpConversion>(context, analysis.getReturns(func));
if (mlir::failed(
mlir::applyPartialConversion(func, target, std::move(patterns)))) {
mlir::emitError(func.getLoc(),
"error in memory allocation optimization\n");
signalPassFailure();
}
}
private:
MemoryAllocationOptions options;
};
}
std::unique_ptr<mlir::Pass> fir::createMemoryAllocationPass() {
return std::make_unique<MemoryAllocationOpt>();
}
std::unique_ptr<mlir::Pass>
fir::createMemoryAllocationPass(bool dynOnHeap, std::size_t maxStackSize) {
return std::make_unique<MemoryAllocationOpt>(dynOnHeap, maxStackSize);
}