brintos

brintos / llvm-project-archived public Read only

0
0
Text · 3.4 KiB · a073a9a Raw
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