brintos

brintos / llvm-project-archived public Read only

0
0
Text · 4.9 KiB · e7602b4 Raw
125 lines · cpp
1//===- TosaToLinalgPass.cpp - Lowering Tosa to Linalg Dialect -------------===//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 transformation pass legalizes Tosa operations to the Linalg dialect.10//11//===----------------------------------------------------------------------===//12 13#include "mlir/Conversion/TosaToLinalg/TosaToLinalg.h"14 15#include "mlir/Dialect/Arith/IR/Arith.h"16#include "mlir/Dialect/Func/IR/FuncOps.h"17#include "mlir/Dialect/Index/IR/IndexDialect.h"18#include "mlir/Dialect/Linalg/IR/Linalg.h"19#include "mlir/Dialect/Math/IR/Math.h"20#include "mlir/Dialect/SCF/IR/SCF.h"21#include "mlir/Dialect/Tensor/IR/Tensor.h"22#include "mlir/Dialect/Tosa/IR/TosaOps.h"23#include "mlir/Dialect/Tosa/Transforms/Passes.h"24#include "mlir/IR/PatternMatch.h"25#include "mlir/Pass/PassManager.h"26#include "mlir/Transforms/DialectConversion.h"27#include "mlir/Transforms/Passes.h"28 29namespace mlir {30#define GEN_PASS_DEF_TOSATOLINALG31#include "mlir/Conversion/Passes.h.inc"32} // namespace mlir33 34using namespace mlir;35 36namespace {37struct TosaToLinalg : public impl::TosaToLinalgBase<TosaToLinalg> {38public:39  void getDependentDialects(DialectRegistry &registry) const override {40    registry41        .insert<arith::ArithDialect, linalg::LinalgDialect, math::MathDialect,42                index::IndexDialect, tensor::TensorDialect, scf::SCFDialect>();43  }44 45  void runOnOperation() override {46    RewritePatternSet patterns(&getContext());47    ConversionTarget target(getContext());48    target.addLegalDialect<linalg::LinalgDialect, tensor::TensorDialect,49                           scf::SCFDialect>();50    target.addIllegalDialect<tosa::TosaDialect>();51 52    // Not every TOSA op can be legalized to linalg.53    target.addLegalOp<tosa::ApplyScaleOp>();54    target.addLegalOp<tosa::IfOp>();55    target.addLegalOp<tosa::ConstOp>();56    target.addLegalOp<tosa::ConstShapeOp>();57    target.addLegalOp<tosa::WhileOp>();58    target.addLegalOp<tosa::ConcatOp>();59    target.addLegalOp<tosa::SliceOp>();60    target.addLegalOp<tosa::ReshapeOp>();61    target.addLegalOp<tosa::PadOp>();62 63    target.markUnknownOpDynamicallyLegal([](Operation *) { return true; });64 65    TypeConverter converter;66    tosa::populateTosaTypeConversion(converter);67 68    FunctionOpInterface func = getOperation();69    mlir::tosa::populateTosaToLinalgConversionPatterns(converter, &patterns);70    if (failed(applyFullConversion(func, target, std::move(patterns))))71      signalPassFailure();72  }73};74} // namespace75 76std::unique_ptr<Pass> mlir::tosa::createTosaToLinalg() {77  return std::make_unique<TosaToLinalg>();78}79 80void mlir::tosa::addTosaToLinalgPasses(81    OpPassManager &pm, const TosaToLinalgOptions &options,82    const TosaToLinalgNamedOptions &tosaToLinalgNamedOptions,83    std::optional<tosa::TosaValidationOptions> validationOptions) {84  // Optional decompositions are designed to benefit linalg.85  if (!options.disableTosaDecompositions)86    pm.addNestedPass<func::FuncOp>(87        tosa::createTosaOptionalDecompositionsPass());88  pm.addNestedPass<func::FuncOp>(createCanonicalizerPass());89 90  pm.addNestedPass<func::FuncOp>(tosa::createTosaInferShapesPass());91  pm.addNestedPass<func::FuncOp>(tosa::createTosaMakeBroadcastablePass());92  pm.addNestedPass<func::FuncOp>(93      tosa::createTosaToLinalgNamed(tosaToLinalgNamedOptions));94  pm.addNestedPass<func::FuncOp>(createCanonicalizerPass());95  // TODO: Remove pass that operates on const tensor and enable optionality96  pm.addNestedPass<func::FuncOp>(tosa::createTosaLayerwiseConstantFoldPass(97      {options.aggressiveReduceConstant}));98  pm.addNestedPass<func::FuncOp>(tosa::createTosaMakeBroadcastablePass());99  if (validationOptions)100    pm.addPass(tosa::createTosaValidation(*validationOptions));101  pm.addNestedPass<func::FuncOp>(tosa::createTosaToLinalg());102}103 104//===----------------------------------------------------------------------===//105// Pipeline registration.106//===----------------------------------------------------------------------===//107 108void mlir::tosa::registerTosaToLinalgPipelines() {109  PassPipelineRegistration<>(110      "tosa-to-linalg-pipeline",111      "The default pipeline for converting TOSA operators to the equivalent "112      "operations using the tensor operations in LinAlg as well as LinAlg "113      "named operations.",114      [](OpPassManager &pm) {115        TosaToLinalgOptions tosaToLinalgOptions;116        TosaToLinalgNamedOptions tosaToLinalgNamedOptions;117        TosaValidationOptions validationOptions;118        validationOptions.strictOpSpecAlignment = false;119        validationOptions.allowInvalidOpDatatypeCombinations = false;120        tosa::addTosaToLinalgPasses(pm, tosaToLinalgOptions,121                                    tosaToLinalgNamedOptions,122                                    validationOptions);123      });124}125