96 lines · cpp
1//===- MemRefToEmitC.cpp - MemRef to EmitC conversion ---------------------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8//9// This file implements a pass to convert memref ops into emitc ops.10//11//===----------------------------------------------------------------------===//12 13#include "mlir/Conversion/MemRefToEmitC/MemRefToEmitCPass.h"14 15#include "mlir/Conversion/MemRefToEmitC/MemRefToEmitC.h"16#include "mlir/Dialect/EmitC/IR/EmitC.h"17#include "mlir/Dialect/MemRef/IR/MemRef.h"18#include "mlir/IR/Attributes.h"19#include "mlir/Pass/Pass.h"20#include "mlir/Transforms/DialectConversion.h"21#include "llvm/ADT/SmallSet.h"22#include "llvm/ADT/StringRef.h"23 24namespace mlir {25#define GEN_PASS_DEF_CONVERTMEMREFTOEMITC26#include "mlir/Conversion/Passes.h.inc"27} // namespace mlir28 29using namespace mlir;30 31namespace {32 33emitc::IncludeOp addStandardHeader(OpBuilder &builder, ModuleOp module,34 StringRef headerName) {35 StringAttr includeAttr = builder.getStringAttr(headerName);36 return emitc::IncludeOp::create(37 builder, module.getLoc(), includeAttr,38 /*is_standard_include=*/builder.getUnitAttr());39}40 41struct ConvertMemRefToEmitCPass42 : public impl::ConvertMemRefToEmitCBase<ConvertMemRefToEmitCPass> {43 using Base::Base;44 void runOnOperation() override {45 TypeConverter converter;46 ConvertMemRefToEmitCOptions options;47 options.lowerToCpp = this->lowerToCpp;48 // Fallback for other types.49 converter.addConversion([](Type type) -> std::optional<Type> {50 if (!emitc::isSupportedEmitCType(type))51 return {};52 return type;53 });54 55 populateMemRefToEmitCTypeConversion(converter);56 57 RewritePatternSet patterns(&getContext());58 populateMemRefToEmitCConversionPatterns(patterns, converter);59 60 ConversionTarget target(getContext());61 target.addIllegalDialect<memref::MemRefDialect>();62 target.addLegalDialect<emitc::EmitCDialect>();63 64 if (failed(applyPartialConversion(getOperation(), target,65 std::move(patterns))))66 return signalPassFailure();67 68 mlir::ModuleOp module = getOperation();69 llvm::SmallSet<StringRef, 4> existingHeaders;70 mlir::OpBuilder builder(module.getBody(), module.getBody()->begin());71 module.walk([&](mlir::emitc::IncludeOp includeOp) {72 if (includeOp.getIsStandardInclude())73 existingHeaders.insert(includeOp.getInclude());74 });75 76 module.walk([&](mlir::emitc::CallOpaqueOp callOp) {77 StringRef expectedHeader;78 if (callOp.getCallee() == alignedAllocFunctionName ||79 callOp.getCallee() == mallocFunctionName)80 expectedHeader = options.lowerToCpp ? cppStandardLibraryHeader81 : cStandardLibraryHeader;82 else if (callOp.getCallee() == memcpyFunctionName)83 expectedHeader =84 options.lowerToCpp ? cppStringLibraryHeader : cStringLibraryHeader;85 else86 return mlir::WalkResult::advance();87 if (!existingHeaders.contains(expectedHeader)) {88 addStandardHeader(builder, module, expectedHeader);89 existingHeaders.insert(expectedHeader);90 }91 return mlir::WalkResult::advance();92 });93 }94};95} // namespace96