24#define GEN_PASS_DEF_LINALGGENERALIZENAMEDOPSPASS
25#include "mlir/Dialect/Linalg/Passes.h.inc"
28#define DEBUG_TYPE "linalg-generalization"
35 if (isa<GenericOp>(linalgOp))
38 if (linalgOp->getNumRegions() != 1) {
39 assert(linalgOp->getNumRegions() == 0 &&
"op with multiple regions");
54 if (!isa<MatmulOp, BatchMatmulOp, BatchReduceMatmulOp>(
55 namedOp.getOperation()))
63 namedOp.getIndexingMapsArray(),
69 if (
auto castAttr = namedOp->getAttrOfType<TypeFnAttr>(
"cast");
70 castAttr && castAttr.getValue() == TypeFn::cast_unsigned)
71 attributes.push_back(rewriter.
getNamedAttr(
"cast", castAttr));
75 Value lhs = namedOp.getDpsInputs()[0];
76 Value rhs = namedOp.getDpsInputs()[1];
77 Value init = namedOp.getDpsInits()[0];
78 DictionaryAttr discardableAttrs = namedOp->getDiscardableAttrDictionary();
88 contractOp->setDiscardableAttrs(discardableAttrs);
95 bool emitCategoryOps) {
100 if (emitCategoryOps) {
102 if (succeeded(contractOp))
105 "failed to categorize to named op");
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);
123 genericOp.getRegion().begin());
130 genericOp->setDiscardableAttrs(linalgOp->getDiscardableAttrDictionary());
132 rewriter.
replaceOp(linalgOp, genericOp->getResults());
133 return cast<LinalgOp>(genericOp.getOperation());
138struct LinalgGeneralizeNamedOpsPass
140 LinalgGeneralizeNamedOpsPass> {
142 LinalgGeneralizeNamedOpsPass>::LinalgGeneralizeNamedOpsPassBase;
143 void runOnOperation()
override;
148void LinalgGeneralizeNamedOpsPass::runOnOperation() {
static FailureOr< LinalgOp > generalizeToContractOp(RewriterBase &rewriter, LinalgOp namedOp)
static LogicalResult generalizeNamedOpPrecondition(LinalgOp linalgOp)
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Attributes are known-constant values of operations.
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
NamedAttribute getNamedAttr(StringRef name, Attribute val)
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 ®ion, 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.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
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 ®ion, 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.