23#include "llvm/ADT/TypeSwitch.h"
26#define GEN_PASS_DEF_LINALGSPECIALIZEGENERICOPSPASS
27#include "mlir/Dialect/Linalg/Passes.h.inc"
30#define DEBUG_TYPE "linalg-specialization"
51 Block *body = genericOp.getBody();
61 "binary op uses just one block arg");
85 Block *body = genericOp.getBody();
87 for (
auto [i, v] : llvm::enumerate(op->
getOperands())) {
88 if (
auto blockArg = dyn_cast<BlockArgument>(v);
89 blockArg && blockArg.getOwner() == body)
103 .Case<arith::AddFOp, arith::AddIOp, complex::AddOp>(
104 [](
Operation *) {
return ElementwiseKind::add; })
105 .Case<arith::MulIOp, arith::MulFOp, complex::MulOp>(
106 [](
Operation *) {
return ElementwiseKind::mul; })
107 .Case([](math::ExpOp) {
return ElementwiseKind::exp; })
108 .Case([](math::AbsFOp) {
return ElementwiseKind::abs; })
109 .Case([](math::CeilOp) {
return ElementwiseKind::ceil; })
110 .Case([](math::FloorOp) {
return ElementwiseKind::floor; })
111 .Case([](arith::NegFOp) {
return ElementwiseKind::negf; })
112 .Case([](math::RoundOp) {
return ElementwiseKind::round; })
113 .Case([](math::SqrtOp) {
return ElementwiseKind::sqrt; })
114 .Case([](math::RsqrtOp) {
return ElementwiseKind::rsqrt; })
115 .Case([](math::TanhOp) {
return ElementwiseKind::tanh; })
116 .Case([](math::ErfOp) {
return ElementwiseKind::erf; })
117 .Case([](math::SinOp) {
return ElementwiseKind::sin; })
118 .Case([](math::CosOp) {
return ElementwiseKind::cos; })
119 .Case([](math::TanOp) {
return ElementwiseKind::tan; })
120 .Case([](math::AcosOp) {
return ElementwiseKind::acos; })
121 .Case([](math::AcoshOp) {
return ElementwiseKind::acosh; })
122 .Case([](math::AsinOp) {
return ElementwiseKind::asin; })
123 .Case([](math::AsinhOp) {
return ElementwiseKind::asinh; })
124 .Case([](math::AtanOp) {
return ElementwiseKind::atan; })
125 .Case([](math::AtanhOp) {
return ElementwiseKind::atanh; })
126 .Case([](math::LogOp) {
return ElementwiseKind::log; })
127 .Case([](math::Log10Op) {
return ElementwiseKind::log10; })
128 .Case([](math::Log1pOp) {
return ElementwiseKind::log1p; })
129 .Case([](math::Log2Op) {
return ElementwiseKind::log2; })
130 .Case<arith::SubIOp, arith::SubFOp, complex::SubOp>(
131 [](
Operation *) {
return ElementwiseKind::sub; })
132 .Case<arith::DivSIOp, arith::DivFOp, complex::DivOp>(
133 [](
Operation *) {
return ElementwiseKind::div; })
134 .Case([](arith::DivUIOp) {
return ElementwiseKind::div_unsigned; })
135 .Case<arith::MaxSIOp, arith::MaximumFOp>(
136 [](
Operation *) {
return ElementwiseKind::max_signed; })
137 .Case<arith::MinSIOp, arith::MinimumFOp>(
138 [](
Operation *) {
return ElementwiseKind::min_signed; })
139 .Case([](math::PowFOp) {
return ElementwiseKind::powf; })
140 .Case([](arith::MaxUIOp) {
return ElementwiseKind::max_unsigned; })
141 .Case([](arith::MinUIOp) {
return ElementwiseKind::min_unsigned; })
142 .Case([](arith::SelectOp) {
return ElementwiseKind::select; })
143 .Default([](
Operation *) {
return std::nullopt; });
173 GenericOp genericOp) {
175 unsigned arity = genericOp.getNumDpsInputs();
176 bool isUnary = arity == 1;
177 bool isBinary = arity == 2;
178 bool isTernary = arity == 3;
181 Operation *op = &genericOp.getBody()->front();
184 bool hasSwappedOperands =
186 int scalarOprIdx = -1;
193 auto replaceOp = [&](ElementwiseKind kind,
194 bool mayHoistScalarOperand =
true) -> LinalgOp {
197 if (hasSwappedOperands) {
201 std::swap(inputs[0 + isTernary], inputs[1 + isTernary]);
202 std::swap(indexingMaps[0 + isTernary], indexingMaps[1 + isTernary]);
205 if (hasScalarOperand && mayHoistScalarOperand) {
207 inputs.insert(inputs.begin() + scalarOprIdx,
209 auto scalarBroadcastMap =
212 indexingMaps.insert(indexingMaps.begin() + scalarOprIdx,
215 auto newOp = ElementwiseOp::create(
216 rewriter, genericOp.getLoc(), inputs, genericOp.getDpsInits(), kind,
224 if (
auto divOp = dyn_cast<arith::DivFOp>(op)) {
225 if (
auto constOp = dyn_cast_if_present<arith::ConstantOp>(
226 divOp.getLhs().getDefiningOp()))
227 if (cast<FloatAttr>(constOp.getValue()).getValue().isExactlyValue(1.0))
228 return replaceOp(ElementwiseKind::reciprocal,
233 if (
auto mulOp = dyn_cast<arith::MulFOp>(op))
234 if (mulOp.getLhs() == mulOp.getRhs())
235 return replaceOp(ElementwiseKind::square);
239 return v.getType().isInteger(1);
241 if (isa<arith::OrIOp>(op))
242 return replaceOp(ElementwiseKind::add);
243 if (isa<arith::AndIOp>(op))
244 return replaceOp(ElementwiseKind::mul);
251 unsigned numInputs = arity + (hasScalarOperand ? 1 : 0);
252 if (numInputs ==
static_cast<unsigned>(arityGroupAndKind.arityGroup))
253 return replaceOp(*kind);
257 genericOp,
"elementwise operation cannot be specialized to category op");
286enum class IndexMatchResult {
300static IndexMatchResult matchOperandMap(
AffineMap map,
unsigned rowDimIdx,
301 unsigned expectedPosOfRowDim,
302 unsigned expectedPosOfColDim) {
304 auto exprOfRowDim = map.
getResults()[rowDimIdx];
305 auto exprOfColDim = map.
getResults()[rowDimIdx + 1];
310 return IndexMatchResult::Mismatch;
312 auto posRowDim = cast<AffineDimExpr>(exprOfRowDim).getPosition();
313 auto posColDim = cast<AffineDimExpr>(exprOfColDim).getPosition();
315 if (expectedPosOfRowDim == posRowDim && expectedPosOfColDim == posColDim)
316 return IndexMatchResult::Match;
318 if (expectedPosOfRowDim == posColDim && expectedPosOfColDim == posRowDim)
319 return IndexMatchResult::Transposed;
321 return IndexMatchResult::Mismatch;
330template <
typename NamedOpTy>
331static LinalgOp replaceWithMatmulVariant(
RewriterBase &rewriter, GenericOp op,
332 std::optional<TypeFn> castTy,
337 if (castTy.has_value() && *castTy == TypeFn::cast_unsigned) {
339 "cast", TypeFnAttr::get(rewriter.
getContext(), *castTy));
340 attributes.push_back(castAttr);
347 return AffineMapAttr::get(map);
350 "indexing_maps", rewriter.
getArrayAttr(indexingMapsAttrVal));
351 attributes.push_back(indexingMapsAttr);
354 op,
ValueRange{op.getDpsInputs()[0], op.getDpsInputs()[1]},
363static std::optional<TypeFn> getCastTypeForMatmulLikeOp(GenericOp genericOp) {
364 bool foundCastForMatmulOutput =
false;
366 genericOp.getBody()->walk([&](CastOpInterface castOp) {
375 if (!llvm::any_of(forwardSlice, [](
Operation *op) {
378 return isa<arith::MulIOp, arith::MulFOp, complex::MulOp>(op);
380 foundCastForMatmulOutput =
true;
385 if (isa<arith::ExtUIOp, arith::UIToFPOp, arith::FPToUIOp>(castOp))
386 castTyFns.push_back(TypeFn::cast_unsigned);
387 else if (isa<arith::ExtSIOp, arith::SIToFPOp, arith::FPToSIOp>(castOp))
388 castTyFns.push_back(TypeFn::cast_signed);
393 if (foundCastForMatmulOutput)
396 if (!castTyFns.empty()) {
400 if (!llvm::all_equal(castTyFns))
402 return castTyFns.front();
406 return TypeFn::cast_signed;
409static FailureOr<LinalgOp> specializeLinalgMmt4D(
RewriterBase &rewriter,
411 std::optional<TypeFn> castTy,
414 auto indexingMaps = genericOp.getIndexingMapsArray();
415 if (llvm::any_of(indexingMaps, [](
AffineMap m) {
420 auto aOuter = matchOperandMap(indexingMaps[0], 0, dims.
m[0], dims.
k[0]);
421 auto aInner = matchOperandMap(indexingMaps[0], 2, dims.
m[1], dims.
k[1]);
423 auto bOuter = matchOperandMap(indexingMaps[1], 0, dims.
k[0], dims.
n[0]);
424 auto bInner = matchOperandMap(indexingMaps[1], 2, dims.
k[1], dims.
n[1]);
426 auto cOuter = matchOperandMap(indexingMaps[2], 0, dims.
m[0], dims.
n[0]);
427 auto cInner = matchOperandMap(indexingMaps[2], 2, dims.
m[1], dims.
n[1]);
429 if (llvm::is_contained({aOuter, bOuter, cOuter}, IndexMatchResult::Mismatch))
431 if (llvm::is_contained({aInner, bInner, cInner}, IndexMatchResult::Mismatch))
437 return replaceWithMatmulVariant<Mmt4DOp>(rewriter, genericOp, castTy,
442 if (isa<arith::MulFOp>(first) && isa<arith::AddFOp>(second))
444 if (isa<arith::MulIOp>(first) && isa<arith::AddIOp>(second))
446 if (isa<complex::MulOp>(first) && isa<complex::AddOp>(second))
448 if (isa<arith::AndIOp>(first) && isa<arith::OrIOp>(second) &&
459static FailureOr<LinalgOp>
460specializeToNamedContraction(
RewriterBase &rewriter, GenericOp genericOp,
461 std::optional<TypeFn> castTy) {
485 if (dims.
m.size() == 2 && dims.
n.size() == 2 && dims.
k.size() == 2)
486 return specializeLinalgMmt4D(rewriter, genericOp, castTy, dims);
487 if (dims.
m.size() != 1 || dims.
n.size() != 1 || dims.
k.size() != 1)
491 auto indexingMaps = genericOp.getIndexingMapsArray();
492 if (llvm::any_of(indexingMaps, [&dims](
AffineMap m) {
494 dims.
batch.size() + 2 ;
498 auto numOfBatchDims = dims.
batch.size();
499 if (indexingMaps[0].getNumDims() != numOfBatchDims + 3)
502 if (numOfBatchDims) {
506 if (llvm::any_of(indexingMaps, [numOfBatchDims](
AffineMap m) {
507 for (
unsigned i = 0; i < numOfBatchDims; ++i) {
510 cast<AffineDimExpr>(expr).getPosition() != i)
519 matchOperandMap(indexingMaps[0], numOfBatchDims, dims.
m[0], dims.
k[0]);
521 matchOperandMap(indexingMaps[1], numOfBatchDims, dims.
k[0], dims.
n[0]);
523 matchOperandMap(indexingMaps[2], numOfBatchDims, dims.
m[0], dims.
n[0]);
525 if (llvm::is_contained({a,
b, c}, IndexMatchResult::Mismatch))
529 auto *ctx = genericOp.getContext();
530 unsigned numLoopDims = numOfBatchDims + 3;
531 unsigned mIdx = numOfBatchDims;
532 unsigned nIdx = mIdx + 1;
533 unsigned kIdx = mIdx + 2;
536 auto makeMap = [&](IndexMatchResult match,
unsigned rowIdx,
unsigned colIdx) {
538 for (
unsigned i = 0; i < numOfBatchDims; ++i)
539 tensorDims.push_back(i);
540 if (match == IndexMatchResult::Transposed)
541 llvm::append_values(tensorDims, colIdx, rowIdx);
543 llvm::append_values(tensorDims, rowIdx, colIdx);
547 auto mapA = makeMap(a, mIdx, kIdx);
548 auto mapB = makeMap(
b, kIdx, nIdx);
549 auto mapC = makeMap(c, mIdx, nIdx);
554 if (numOfBatchDims) {
555 return replaceWithMatmulVariant<BatchMatmulOp>(rewriter, genericOp, castTy,
558 return replaceWithMatmulVariant<MatmulOp>(rewriter, genericOp, castTy,
564static FailureOr<LinalgOp> specializeLinalgContractions(
RewriterBase &rewriter,
566 bool emitCategoryOp) {
567 if (genericOp.getNumDpsInputs() != 2 || genericOp.getNumDpsInits() != 1)
571 auto mapRange = genericOp.getIndexingMapsArray();
572 if (llvm::any_of(mapRange,
581 isSupportedContractionPair))
586 std::optional<TypeFn> castTy = getCastTypeForMatmulLikeOp(genericOp);
589 genericOp,
"contains invalid cast ops for the named matmul op");
595 if (!emitCategoryOp) {
596 FailureOr<LinalgOp> namedOp =
597 specializeToNamedContraction(rewriter, genericOp, castTy);
598 if (succeeded(namedOp))
604 return replaceWithMatmulVariant<ContractOp>(rewriter, genericOp, castTy,
605 genericOp.getIndexingMapsArray());
610template <
typename ConvOpTy>
611static FailureOr<LinalgOp>
612specializeToConvOp(
RewriterBase &rewriter, GenericOp genericOp,
621 if constexpr (std::is_same_v<ConvOpTy, linalg::Conv1DOp> ||
622 std::is_same_v<ConvOpTy, linalg::Conv2DOp> ||
623 std::is_same_v<ConvOpTy, linalg::Conv3DOp>) {
630 genericOp, resultTypes, inputs, outputs, stridesAttr, dilationsAttr);
636static FailureOr<LinalgOp> specializeLinalgConvolutions(
RewriterBase &rewriter,
637 GenericOp genericOp) {
638#define CONV_OP_SPECIALIZER(ConvOpTy) \
639 if (std::optional<DilationsAndStrides> convParams = \
640 matchConvolutionOpOfType<ConvOpTy>(genericOp)) \
641 return specializeToConvOp<ConvOpTy>( \
642 rewriter, genericOp, convParams->dilations, convParams->strides); \
699#undef CONV_OP_SPECIALIZER
710 bool emitCategoryOps) {
720 return specializeLinalgContractions(rewriter, genericOp, emitCategoryOps);
728 genericOp,
"no matching category op specialization");
733 genericOp, genericOp.getDpsInputs()[0], genericOp.getDpsInits()[0]);
741 genericOp, *fillValue, genericOp.getDpsInits()[0]);
746 std::optional<SmallVector<int64_t>> equivalentToBroadcast =
748 if (equivalentToBroadcast) {
749 auto dims = *equivalentToBroadcast;
751 genericOp, genericOp.getDpsInputs()[0], genericOp.getDpsInits()[0],
757 std::optional<SmallVector<int64_t>> equivalentToTranspose =
759 if (equivalentToTranspose) {
760 auto permutation = *equivalentToTranspose;
762 genericOp, genericOp.getDpsInputs()[0], genericOp.getDpsInits()[0],
769 return specializeLinalgConvolutions(rewriter, genericOp);
772 "no matching named op specialization");
776struct LinalgSpecializeGenericOpsPass
777 :
public impl::LinalgSpecializeGenericOpsPassBase<
778 LinalgSpecializeGenericOpsPass> {
780 using impl::LinalgSpecializeGenericOpsPassBase<
781 LinalgSpecializeGenericOpsPass>::LinalgSpecializeGenericOpsPassBase;
782 void runOnOperation()
override;
786void LinalgSpecializeGenericOpsPass::runOnOperation() {
static bool findIndexOfScalarOperand(GenericOp genericOp, int &index)
static std::optional< ElementwiseKind > getElementwiseKind(Operation *op)
static bool areBinOpsSwapped(GenericOp genericOp, bool isTernary)
#define CONV_OP_SPECIALIZER(ConvOpTy)
static FailureOr< LinalgOp > specializeLinalgElementwise(RewriterBase &rewriter, GenericOp genericOp)
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
static AffineMap getMultiDimMapWithTargets(unsigned numDims, ArrayRef< unsigned > targets, MLIRContext *context)
Returns an affine map with numDims input dimensions and results specified by targets.
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
BlockArgument getArgument(unsigned i)
DenseIntElementsAttr getI64TensorAttr(ArrayRef< int64_t > values)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
NamedAttribute getNamedAttr(StringRef name, Attribute val)
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
IRValueT get() const
Return the current value being used by this operand.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
unsigned getNumOperands()
operand_range getOperands()
Returns an iterator on the underlying Value's.
OpOperand & getOpOperand(unsigned idx)
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,...
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.
bool isInteger() const
Return true if this is an integer type (with the specified width).
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...
Type getType() const
Return the type of this value.
static WalkResult advance()
static WalkResult interrupt()
bool isContractionBody(Block &block, function_ref< bool(Operation *, Operation *)> isaPair, llvm::raw_ostream &errs=mlir::thread_safe_nulls())
Returns true if the block contains a contraction of the following form:
std::optional< SmallVector< int64_t > > isaTransposeOpInterface(GenericOp genericOp)
Checks whether genericOp is semantically equivalent to a linalg.transpose.
bool isaElemwiseSingleUnaryOpInterface(GenericOp genericOp)
Checks whether a given genericOp is semantically equivalent to a single linalg elementwise unary op,...
bool isaCopyOpInterface(LinalgOp linalgOp)
Checks whether linalgOp is semantically equivalent to a linalg.copyOp.
void populateLinalgGenericOpsSpecializationPatterns(RewritePatternSet &patterns, bool emitCategoryOps=false)
Populates patterns with patterns to convert linalg.generic ops to named or category ops where possibl...
void populateDecomposeProjectedPermutationPatterns(RewritePatternSet &patterns)
Add patterns to make explicit broadcasts and transforms in the input operands of a genericOp.
bool isaConvolutionOpInterface(LinalgOp linalgOp, bool allowEmptyConvolvedDims=false)
Checks whether linalgOp conforms to ConvolutionOpInterface.
bool isaElemwiseSingleTernaryOpInterface(GenericOp genericOp)
Checks whether genericOp is semantically equivalent to a single linalg elementwise ternary op e....
FailureOr< LinalgOp > specializeGenericOp(RewriterBase &rewriter, GenericOp genericOp, bool emitCategoryOps=false)
Replace the given GenericOp with a namedOp or categoryOp.
std::optional< SmallVector< int64_t > > isaBroadcastOpInterface(LinalgOp linalgOp)
Checks whether linalgOp is semantically equivalent to a broadcast operation.
FailureOr< ContractionDimensions > inferContractionDims(LinalgOp linalgOp)
Find at least 2 parallel (m and n) and 1 reduction (k) dimension candidates that form a matmul subcom...
bool isaContractionOpInterface(LinalgOp linalgOp)
Checks whether linalgOp conforms to ContractionOpInterface.
ArityGroupAndKind getArityGroupAndKind(ElementwiseKind kind)
std::optional< Value > isaFillOpInterface(GenericOp genericOp)
Checks whether genericOp is semantically equivalent to a linalg.fill.
bool isaElemwiseSingleBinaryOpInterface(GenericOp genericOp)
Checks whether genericOp is semantically equivalent to a single linalg elementwise binary op e....
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...
@ DimId
Dimensional identifier.
llvm::SetVector< T, Vector, Set, N > SetVector
void getForwardSlice(Operation *op, SetVector< Operation * > *forwardSlice, const ForwardSliceOptions &options={})
Fills forwardSlice with the computed forward slice (i.e.
Positions of a Linalg op loops that correspond to different kinds of a contraction dimension.
SmallVector< unsigned, 2 > batch
SmallVector< unsigned, 2 > m
SmallVector< unsigned, 2 > n
SmallVector< unsigned, 2 > k