MLIR 24.0.0git
Generalization.cpp
Go to the documentation of this file.
1//===- Generalization.cpp - linalg named ops to generic ops --------------===//
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 file implements the Linalg generalization pass. It converts named
10// Linalg ops to linalg.generic ops.
11//
12//===----------------------------------------------------------------------===//
13
15
18#include "mlir/IR/AffineMap.h"
19#include "mlir/IR/Builders.h"
22
23namespace mlir {
24#define GEN_PASS_DEF_LINALGGENERALIZENAMEDOPSPASS
25#include "mlir/Dialect/Linalg/Passes.h.inc"
26} // namespace mlir
27
28#define DEBUG_TYPE "linalg-generalization"
29
30using namespace mlir;
31using namespace mlir::linalg;
32
33static LogicalResult generalizeNamedOpPrecondition(LinalgOp linalgOp) {
34 // Bailout if `linalgOp` is already a generic.
35 if (isa<GenericOp>(linalgOp))
36 return failure();
37 // Check if the operation has exactly one region.
38 if (linalgOp->getNumRegions() != 1) {
39 assert(linalgOp->getNumRegions() == 0 && "op with multiple regions");
40 // TOD: Otherwise it needs to be built explicitly from the region builder.
41 return failure();
42 }
43 return success();
44}
45
46// Converts a named matmul-like op (`matmul`, `batch_matmul`, or
47// `batch_reduce_matmul`) into a `linalg.contract` category op, preserving the
48// operand indexing maps and cast semantics. Returns failure for other ops.
49static FailureOr<LinalgOp> generalizeToContractOp(RewriterBase &rewriter,
50 LinalgOp namedOp) {
51 // These are the ODS-defined matmul-like operations.
52 // For OpDSL declared contractions, please move them to ODS first,
53 // then add them to the isa<> check below + tests.
54 if (!isa<MatmulOp, BatchMatmulOp, BatchReduceMatmulOp>(
55 namedOp.getOperation()))
56 return failure();
57
59
60 // Preserve operand indexing semantics (transposition, batch/reduction dims)
61 // via the named op's indexing maps.
62 SmallVector<Attribute> indexingMaps = llvm::map_to_vector(
63 namedOp.getIndexingMapsArray(),
64 [](AffineMap map) -> Attribute { return AffineMapAttr::get(map); });
65 attributes.push_back(rewriter.getNamedAttr(
66 "indexing_maps", rewriter.getArrayAttr(indexingMaps)));
67
68 // Only the unsigned cast needs to be explicit; signed is the default.
69 if (auto castAttr = namedOp->getAttrOfType<TypeFnAttr>("cast");
70 castAttr && castAttr.getValue() == TypeFn::cast_unsigned)
71 attributes.push_back(rewriter.getNamedAttr("cast", castAttr));
72
73 // Capture the operands and discardable attributes before `namedOp` is erased
74 // by the replacement below.
75 Value lhs = namedOp.getDpsInputs()[0];
76 Value rhs = namedOp.getDpsInputs()[1];
77 Value init = namedOp.getDpsInits()[0];
78 DictionaryAttr discardableAttrs = namedOp->getDiscardableAttrDictionary();
79
80 LinalgOp contractOp = rewriter.replaceOpWithNewOp<ContractOp>(
81 namedOp, ValueRange{lhs, rhs}, ValueRange{init}, attributes);
82
83 // Discardable attributes carry user-defined metadata (e.g., annotations for
84 // downstream passes). Generalization is a semantics-preserving
85 // transformation, so dropping this metadata would be unexpected. This is safe
86 // because discardable attributes are by definition independent of op
87 // semantics.
88 contractOp->setDiscardableAttrs(discardableAttrs);
89
90 return contractOp;
91}
92
93FailureOr<LinalgOp> mlir::linalg::generalizeNamedOp(RewriterBase &rewriter,
94 LinalgOp linalgOp,
95 bool emitCategoryOps) {
96 if (failed(generalizeNamedOpPrecondition(linalgOp)))
97 return rewriter.notifyMatchFailure(linalgOp, "preconditions not met");
98
99 // Emit the `linalg.contract` category op for matmul-like named ops.
100 if (emitCategoryOps) {
101 FailureOr<LinalgOp> contractOp = generalizeToContractOp(rewriter, linalgOp);
102 if (succeeded(contractOp))
103 return contractOp;
104 return rewriter.notifyMatchFailure(linalgOp,
105 "failed to categorize to named op");
106 }
107
108 SmallVector<Value> inputs = linalgOp.getDpsInputs();
109 ValueRange outputs = linalgOp.getDpsInits();
110 SmallVector<AffineMap> indexingMaps = linalgOp.getIndexingMapsArray();
111 SmallVector<utils::IteratorType> iterators = linalgOp.getIteratorTypesArray();
112 SmallVector<Type> resultTypes = linalgOp.hasPureTensorSemantics()
113 ? TypeRange(ValueRange(outputs))
114 : TypeRange{};
115
116 // All named ops have a region attached that can be inlined.
117 assert(linalgOp->getNumRegions() == 1 &&
118 "expect named op to have one region attached");
119 GenericOp genericOp =
120 GenericOp::create(rewriter, linalgOp.getLoc(), resultTypes, inputs,
121 outputs, indexingMaps, iterators);
122 rewriter.inlineRegionBefore(linalgOp->getRegion(0), genericOp.getRegion(),
123 genericOp.getRegion().begin());
124
125 // Discardable attributes carry user-defined metadata (e.g., annotations for
126 // downstream passes). Generalization is a semantics-preserving
127 // transformation, so dropping this metadata would be unexpected. This is safe
128 // because discardable attributes are by definition independent of op
129 // semantics.
130 genericOp->setDiscardableAttrs(linalgOp->getDiscardableAttrDictionary());
131
132 rewriter.replaceOp(linalgOp, genericOp->getResults());
133 return cast<LinalgOp>(genericOp.getOperation());
134}
135
136namespace {
137
138struct LinalgGeneralizeNamedOpsPass
140 LinalgGeneralizeNamedOpsPass> {
142 LinalgGeneralizeNamedOpsPass>::LinalgGeneralizeNamedOpsPassBase;
143 void runOnOperation() override;
144};
145
146} // namespace
147
148void LinalgGeneralizeNamedOpsPass::runOnOperation() {
149 RewritePatternSet patterns(&getContext());
151 (void)applyPatternsGreedily(getOperation(), std::move(patterns));
152}
153
155 RewritePatternSet &patterns, bool emitCategoryOps) {
156 patterns.add<LinalgGeneralizationPattern>(patterns.getContext(),
157 emitCategoryOps);
158}
return success()
static FailureOr< LinalgOp > generalizeToContractOp(RewriterBase &rewriter, LinalgOp namedOp)
static LogicalResult generalizeNamedOpPrecondition(LinalgOp linalgOp)
b getContext())
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
Attributes are known-constant values of operations.
Definition Attributes.h:25
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
NamedAttribute getNamedAttr(StringRef name, Attribute val)
Definition Builders.cpp:102
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void inlineRegionBefore(Region &region, Region &parent, Region::iterator before)
Move the blocks that belong to "region" before the given position in another region "parent".
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
void populateLinalgNamedOpsGeneralizationPatterns(RewritePatternSet &patterns, bool emitCategoryOps=false)
Linalg generalization patterns.
FailureOr< LinalgOp > generalizeNamedOp(RewriterBase &rewriter, LinalgOp linalgOp, bool emitCategoryOps=false)
Create a GenericOp or CategoryOp from the given named operation linalgOp and replace the given linalg...
Include the generated interface declarations.
LogicalResult applyPatternsGreedily(Region &region, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
Linalg generalization pattern.