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 ®istry) 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