10a391c60SMaheshRavishankar //===- TestSlicing.cpp - Testing slice functionality ----------------------===//
20a391c60SMaheshRavishankar //
30a391c60SMaheshRavishankar // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
40a391c60SMaheshRavishankar // See https://llvm.org/LICENSE.txt for license information.
50a391c60SMaheshRavishankar // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
60a391c60SMaheshRavishankar //
70a391c60SMaheshRavishankar //===----------------------------------------------------------------------===//
80a391c60SMaheshRavishankar //
90a391c60SMaheshRavishankar // This file implements a simple testing pass for slicing.
100a391c60SMaheshRavishankar //
110a391c60SMaheshRavishankar //===----------------------------------------------------------------------===//
120a391c60SMaheshRavishankar 
130a391c60SMaheshRavishankar #include "mlir/Analysis/SliceAnalysis.h"
1423aa5a74SRiver Riddle #include "mlir/Dialect/Func/IR/FuncOps.h"
15b7f2c108Sgysit #include "mlir/Dialect/Linalg/IR/Linalg.h"
160a391c60SMaheshRavishankar #include "mlir/IR/BlockAndValueMapping.h"
1765fcddffSRiver Riddle #include "mlir/IR/BuiltinOps.h"
180a391c60SMaheshRavishankar #include "mlir/IR/PatternMatch.h"
190a391c60SMaheshRavishankar #include "mlir/Pass/Pass.h"
200a391c60SMaheshRavishankar #include "mlir/Support/LLVM.h"
210a391c60SMaheshRavishankar 
220a391c60SMaheshRavishankar using namespace mlir;
230a391c60SMaheshRavishankar 
240a391c60SMaheshRavishankar /// Create a function with the same signature as the parent function of `op`
250a391c60SMaheshRavishankar /// with name being the function name and a `suffix`.
createBackwardSliceFunction(Operation * op,StringRef suffix)260a391c60SMaheshRavishankar static LogicalResult createBackwardSliceFunction(Operation *op,
270a391c60SMaheshRavishankar                                                  StringRef suffix) {
28*58ceae95SRiver Riddle   func::FuncOp parentFuncOp = op->getParentOfType<func::FuncOp>();
290a391c60SMaheshRavishankar   OpBuilder builder(parentFuncOp);
300a391c60SMaheshRavishankar   Location loc = op->getLoc();
310a391c60SMaheshRavishankar   std::string clonedFuncOpName = parentFuncOp.getName().str() + suffix.str();
32*58ceae95SRiver Riddle   func::FuncOp clonedFuncOp = builder.create<func::FuncOp>(
33*58ceae95SRiver Riddle       loc, clonedFuncOpName, parentFuncOp.getFunctionType());
340a391c60SMaheshRavishankar   BlockAndValueMapping mapper;
350a391c60SMaheshRavishankar   builder.setInsertionPointToEnd(clonedFuncOp.addEntryBlock());
3689de9cc8SMehdi Amini   for (const auto &arg : enumerate(parentFuncOp.getArguments()))
370a391c60SMaheshRavishankar     mapper.map(arg.value(), clonedFuncOp.getArgument(arg.index()));
384efb7754SRiver Riddle   SetVector<Operation *> slice;
390a391c60SMaheshRavishankar   getBackwardSlice(op, &slice);
400a391c60SMaheshRavishankar   for (Operation *slicedOp : slice)
410a391c60SMaheshRavishankar     builder.clone(*slicedOp, mapper);
4223aa5a74SRiver Riddle   builder.create<func::ReturnOp>(loc);
430a391c60SMaheshRavishankar   return success();
440a391c60SMaheshRavishankar }
450a391c60SMaheshRavishankar 
460a391c60SMaheshRavishankar namespace {
470a391c60SMaheshRavishankar /// Pass to test slice generated from slice analysis.
480a391c60SMaheshRavishankar struct SliceAnalysisTestPass
490a391c60SMaheshRavishankar     : public PassWrapper<SliceAnalysisTestPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID__anon934712fb0111::SliceAnalysisTestPass505e50dd04SRiver Riddle   MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SliceAnalysisTestPass)
515e50dd04SRiver Riddle 
52b5e22e6dSMehdi Amini   StringRef getArgument() const final { return "slice-analysis-test"; }
getDescription__anon934712fb0111::SliceAnalysisTestPass53b5e22e6dSMehdi Amini   StringRef getDescription() const final {
54b5e22e6dSMehdi Amini     return "Test Slice analysis functionality.";
55b5e22e6dSMehdi Amini   }
560a391c60SMaheshRavishankar   void runOnOperation() override;
570a391c60SMaheshRavishankar   SliceAnalysisTestPass() = default;
SliceAnalysisTestPass__anon934712fb0111::SliceAnalysisTestPass580a391c60SMaheshRavishankar   SliceAnalysisTestPass(const SliceAnalysisTestPass &) {}
590a391c60SMaheshRavishankar };
600a391c60SMaheshRavishankar } // namespace
610a391c60SMaheshRavishankar 
runOnOperation()620a391c60SMaheshRavishankar void SliceAnalysisTestPass::runOnOperation() {
630a391c60SMaheshRavishankar   ModuleOp module = getOperation();
64*58ceae95SRiver Riddle   auto funcOps = module.getOps<func::FuncOp>();
650a391c60SMaheshRavishankar   unsigned opNum = 0;
660a391c60SMaheshRavishankar   for (auto funcOp : funcOps) {
670a391c60SMaheshRavishankar     // TODO: For now this is just looking for Linalg ops. It can be generalized
680a391c60SMaheshRavishankar     // to look for other ops using flags.
690a391c60SMaheshRavishankar     funcOp.walk([&](Operation *op) {
700a391c60SMaheshRavishankar       if (!isa<linalg::LinalgOp>(op))
710a391c60SMaheshRavishankar         return WalkResult::advance();
720a391c60SMaheshRavishankar       std::string append =
730a391c60SMaheshRavishankar           std::string("__backward_slice__") + std::to_string(opNum);
74e21adfa3SRiver Riddle       (void)createBackwardSliceFunction(op, append);
750a391c60SMaheshRavishankar       opNum++;
760a391c60SMaheshRavishankar       return WalkResult::advance();
770a391c60SMaheshRavishankar     });
780a391c60SMaheshRavishankar   }
790a391c60SMaheshRavishankar }
800a391c60SMaheshRavishankar 
810a391c60SMaheshRavishankar namespace mlir {
registerSliceAnalysisTestPass()820a391c60SMaheshRavishankar void registerSliceAnalysisTestPass() {
83b5e22e6dSMehdi Amini   PassRegistration<SliceAnalysisTestPass>();
840a391c60SMaheshRavishankar }
850a391c60SMaheshRavishankar } // namespace mlir
86