MLIR 24.0.0git
TosaToLinalgPass.cpp
Go to the documentation of this file.
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-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This transformation pass legalizes Tosa operations to the Linalg dialect.
10//
11//===----------------------------------------------------------------------===//
12
14
30
31namespace mlir {
32#define GEN_PASS_DEF_TOSATOLINALG
33#include "mlir/Conversion/Passes.h.inc"
34} // namespace mlir
35
36using namespace mlir;
37
38namespace {
39struct TosaToLinalg : public impl::TosaToLinalgBase<TosaToLinalg> {
40public:
41 TosaToLinalg(const TosaToLinalgOptions &options)
42 : impl::TosaToLinalgBase<TosaToLinalg>(options) {}
43
44 void getDependentDialects(DialectRegistry &registry) const override {
45 registry
46 .insert<arith::ArithDialect, linalg::LinalgDialect, math::MathDialect,
47 index::IndexDialect, tensor::TensorDialect, scf::SCFDialect>();
48 }
49
50 void runOnOperation() override {
51 RewritePatternSet patterns(&getContext());
52 ConversionTarget target(getContext());
53 target.addLegalDialect<linalg::LinalgDialect, tensor::TensorDialect,
54 scf::SCFDialect>();
55 target.addIllegalDialect<tosa::TosaDialect>();
56
57 // Not every TOSA op can be legalized to linalg.
58 target.addLegalOp<tosa::ApplyScaleOp>();
59 target.addLegalOp<tosa::IfOp>();
60 target.addLegalOp<tosa::ConstOp>();
61 target.addLegalOp<tosa::ConstShapeOp>();
62 target.addLegalOp<tosa::WhileOp>();
63 target.addLegalOp<tosa::ConcatOp>();
64 target.addLegalOp<tosa::SliceOp>();
65 target.addLegalOp<tosa::ReshapeOp>();
66 target.addLegalOp<tosa::PadOp>();
67
68 target.markUnknownOpDynamicallyLegal([](Operation *) { return true; });
69
70 TypeConverter converter;
72
73 FunctionOpInterface func = getOperation();
74 TosaToLinalgOptions options;
75 options.allowNonFinites = allowNonFinites;
77 options);
78 if (failed(applyFullConversion(func, target, std::move(patterns))))
79 signalPassFailure();
80 }
81};
82} // namespace
83
84std::unique_ptr<Pass>
85mlir::tosa::createTosaToLinalg(const TosaToLinalgOptions &options) {
86 return std::make_unique<TosaToLinalg>(options);
87}
88
90 OpPassManager &pm, const TosaToLinalgOptions &options,
91 const TosaToLinalgNamedOptions &tosaToLinalgNamedOptions,
92 std::optional<tosa::TosaValidationOptions> validationOptions,
93 std::optional<TosaAttachTargetOptions> attachTargetOptions) {
94 // Optional decompositions are designed to benefit linalg.
95 if (!options.disableTosaDecompositions)
96 pm.addNestedPass<func::FuncOp>(
97 tosa::createTosaOptionalDecompositionsPass());
98 pm.addNestedPass<func::FuncOp>(createCanonicalizerPass());
99
100 pm.addNestedPass<func::FuncOp>(tosa::createTosaInferShapesPass());
101 pm.addNestedPass<func::FuncOp>(tosa::createTosaMakeBroadcastablePass());
102 pm.addNestedPass<func::FuncOp>(
103 tosa::createTosaToLinalgNamed(tosaToLinalgNamedOptions));
104 pm.addNestedPass<func::FuncOp>(createCanonicalizerPass());
105 // TODO: Remove pass that operates on const tensor and enable optionality
106 pm.addNestedPass<func::FuncOp>(tosa::createTosaLayerwiseConstantFoldPass(
107 {options.aggressiveReduceConstant}));
108 pm.addNestedPass<func::FuncOp>(tosa::createTosaMakeBroadcastablePass());
109 // tosa-attach-target writes a tosa.target_env module attribute, schedule it
110 // only when the caller actually needs one. Callers that opt out of both no
111 // longer get a tosa.target_env attribute they did not ask for.
112 if (validationOptions || attachTargetOptions) {
113 if (!attachTargetOptions) {
114 attachTargetOptions = TosaAttachTargetOptions();
115 attachTargetOptions->profiles = {"pro_int", "pro_fp"};
116 // TODO: populate with all the extensions that the tosa->linalg
117 // conversion supports
118 attachTargetOptions->extensions = {"doubleround"};
119 }
120 pm.addPass(tosa::createTosaAttachTarget(*attachTargetOptions));
121 }
122 if (validationOptions)
123 pm.addPass(tosa::createTosaValidation(*validationOptions));
125}
126
127//===----------------------------------------------------------------------===//
128// Pipeline registration.
129//===----------------------------------------------------------------------===//
130
131namespace {
132/// Options controlling the registered `tosa-to-linalg-pipeline`.
133struct TosaToLinalgPipelineOptions
134 : public PassPipelineOptions<TosaToLinalgPipelineOptions> {
135 PassOptions::Option<bool> validation{
136 *this, "validation",
137 llvm::cl::desc("Run tosa-attach-target and tosa-validate as part of the "
138 "pipeline."),
139 llvm::cl::init(true)};
140};
141} // namespace
142
145 "tosa-to-linalg-pipeline",
146 "The default pipeline for converting TOSA operators to the equivalent "
147 "operations using the tensor operations in LinAlg as well as LinAlg "
148 "named operations.",
149 [](OpPassManager &pm, const TosaToLinalgPipelineOptions &pipelineOpts) {
150 TosaToLinalgOptions tosaToLinalgOptions;
151 TosaToLinalgNamedOptions tosaToLinalgNamedOptions;
152 std::optional<TosaValidationOptions> validationOptions;
153 if (pipelineOpts.validation) {
154 validationOptions = TosaValidationOptions{
155 /*strictOpSpecAlignment=*/false,
156 /*allowInvalidOpDatatypeCombinations=*/false,
157 /*validateFunctionSignature=*/false};
158 }
159 tosa::addTosaToLinalgPasses(pm, tosaToLinalgOptions,
160 tosaToLinalgNamedOptions,
161 validationOptions);
162 });
163}
b getContext())
static llvm::ManagedStatic< PassManagerOptions > options
This class represents a pass manager that runs passes on either a specific operation type,...
Definition PassManager.h:46
void addPass(std::unique_ptr< Pass > pass)
Add the given pass to this pass manager.
Definition Pass.cpp:392
void addNestedPass(std::unique_ptr< Pass > pass)
Add the given pass to a nested pass manager for the given operation kind OpT.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
std::unique_ptr< Pass > createTosaToLinalgNamed(const TosaToLinalgNamedOptions &options=TosaToLinalgNamedOptions())
void addTosaToLinalgPasses(OpPassManager &pm, const TosaToLinalgOptions &options, const TosaToLinalgNamedOptions &tosaToLinalgNamedOptions=TosaToLinalgNamedOptions(), std::optional< tosa::TosaValidationOptions > validationOptions=tosa::TosaValidationOptions{false, false}, std::optional< TosaAttachTargetOptions > attachTargetOptions=std::nullopt)
Populates passes to convert from TOSA to Linalg.
std::unique_ptr< Pass > createTosaToLinalg(const TosaToLinalgOptions &options=TosaToLinalgOptions())
void populateTosaToLinalgConversionPatterns(const TypeConverter &converter, RewritePatternSet *patterns, const TosaToLinalgOptions &options=TosaToLinalgOptions())
Populates conversion passes from TOSA dialect to Linalg dialect.
void populateTosaTypeConversion(TypeConverter &converter)
void registerTosaToLinalgPipelines()
Populates TOSA to linalg pipelines Currently, this includes only the "tosa-to-linalg-pipeline".
Include the generated interface declarations.
std::unique_ptr< Pass > createCanonicalizerPass(const GreedyRewriteConfig &config, ArrayRef< std::string > disabledPatterns={}, ArrayRef< std::string > enabledPatterns={})
Creates an instance of the Canonicalizer pass with the specified config.
PassPipelineRegistration provides a global initializer that registers a Pass pipeline builder routine...