#include "mlir/Conversion/AffineToStandard/AffineToStandard.h" #include "mlir/Transforms/Passes.h" #include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" #include "src/Accelerators/PIM/Compiler/PimCompilerUtils.hpp" #include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp" #include "src/Accelerators/PIM/Pass/PIMPasses.h" #include "src/Compiler/CompilerPasses.hpp" #define DEBUG_TYPE "PimCompilerUtils" using namespace mlir; using namespace onnx_mlir; namespace onnx_mlir { void addPassesPim(OwningOpRef& module, PassManager& pm, EmissionTargetType& emissionTarget, std::string outputNameNoExt) { verifyExplicitPimCoreCount(); if (pimOnlyCodegen) { pm.addPass(createPimLocalMemoryPlanningPass()); pm.addPass(createPimVerificationPass()); pm.addPass(createEmitPimCodePass()); return; } if (emissionTarget >= EmitONNXIR) addONNXToMLIRPasses(pm, /*target CPU*/ false); if (pimEmissionTarget >= EmitSpatial) { pm.addPass(createONNXToSpatialPass()); pm.addPass(createSpatialLayoutPlanningPass()); pm.addPass(createLowerSpatialPlansPass()); pm.addPass(createTrivialGraphComputeMergePass()); pm.addPass(createMergeComputeNodesPass()); pm.addPass(createMessagePass("Onnx lowered to Spatial")); } if (pimEmissionTarget >= EmitPim) { pm.addPass(createSpatialToPimPass()); pm.addPass(createMessagePass("Spatial lowered to Pim")); } if (pimEmissionTarget >= EmitPimBufferized) { pm.addPass(createPimBufferizationPass()); pm.addPass(createMessagePass("Pim bufferized")); } if (pimEmissionTarget >= EmitPimCodegen) { pm.addPass(mlir::createLowerAffinePass()); pm.addPass(createPimHostConstantFoldingPass()); pm.addPass(createMessagePass("Pim host constants folded")); pm.addPass(createPimLocalMemoryPlanningPass()); pm.addPass(createMessagePass("Pim local memory planned")); pm.addPass(createPimVerificationPass()); pm.addPass(createMessagePass("Pim verified")); pm.addPass(createEmitPimCodePass()); pm.addPass(createMessagePass("Pim code emitted")); } } } // namespace onnx_mlir