#include "mlir/Dialect/Bufferization/Transforms/Passes.h"
#include "mlir/Dialect/Bufferization/IR/BufferizableOpInterface.h"
#include "mlir/Dialect/Bufferization/Transforms/Bufferize.h"
#include "mlir/Dialect/Bufferization/Transforms/OneShotAnalysis.h"
#include "mlir/Dialect/Bufferization/Transforms/OneShotModuleBufferize.h"
#include "mlir/Dialect/Bufferization/Transforms/Transforms.h"
using namespace mlir;
using namespace mlir::bufferization;
LogicalResult mlir::bufferization::insertTensorCopies(
Operation *op, const OneShotBufferizationOptions &options,
const BufferizationState &bufferizationState,
BufferizationStatistics *statistics) {
OneShotAnalysisState analysisState(op, options);
if (options.bufferizeFunctionBoundaries) {
if (failed(analyzeModuleOp(op, analysisState, statistics)))
return failure();
} else {
if (failed(analyzeOp(op, analysisState, statistics)))
return failure();
}
if (options.testAnalysisOnly)
return success();
return insertTensorCopies(op, analysisState, bufferizationState);
}
LogicalResult mlir::bufferization::insertTensorCopies(
Operation *op, const AnalysisState &analysisState,
const BufferizationState &bufferizationState) {
IRRewriter rewriter(op->getContext());
WalkResult result = op->walk([&](Operation *nestedOp) {
if (op->hasTrait<OpTrait::SymbolTable>() &&
nestedOp->getParentWithTrait<OpTrait::SymbolTable>() != op)
return WalkResult::skip();
auto bufferizableOp =
analysisState.getOptions().dynCastBufferizableOp(nestedOp);
if (!bufferizableOp)
return WalkResult::skip();
rewriter.setInsertionPoint(nestedOp);
if (failed(bufferizableOp.resolveConflicts(rewriter, analysisState,
bufferizationState)))
return WalkResult::interrupt();
return WalkResult::advance();
});
return failure(result.wasInterrupted());
}