MLIR 24.0.0git
Specialize.cpp
Go to the documentation of this file.
1//===- Specialize.cpp - linalg generic ops to named 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 a method to specialize generic operations to named
10// operations. Conceptually it is the opposite of generalize.cpp.
11//
12//===----------------------------------------------------------------------===//
13
23#include "llvm/ADT/TypeSwitch.h"
24
25namespace mlir {
26#define GEN_PASS_DEF_LINALGSPECIALIZEGENERICOPSPASS
27#include "mlir/Dialect/Linalg/Passes.h.inc"
28} // namespace mlir
29
30#define DEBUG_TYPE "linalg-specialization"
31
32using namespace mlir;
33using namespace mlir::linalg;
34
35//===----------------------------------------------------------------------===//
36// Specialize linalg generic to elementwise ops.
37//===----------------------------------------------------------------------===//
38
39// Given an elementwise single binary linalg generic op, checks whether the
40// binary op accesses operands as swapped. e.g.
41// this differentiates between a linalg-generic body that contains:
42// ^bb0(%a: f32, %b: f32, %c : f32):
43// %0 = arith.subf %a, %b : f32
44// linalg.yield %0: f32
45// against:
46// ^bb0(%a: f32, %b: f32, %c : f32):
47// %0 = arith.subf %b, %a : f32
48// linalg.yield %0: f32
49// Former is linalg.sub(a,b), latter is linalg.sub(b,a).
50static bool areBinOpsSwapped(GenericOp genericOp, bool isTernary) {
51 Block *body = genericOp.getBody();
52 Operation *op = &body->front();
53 bool swapped = false;
54 if (op->getOpOperand(0 + isTernary).get() !=
55 body->getArgument(0 + isTernary)) {
56 swapped = true;
57 assert(op->getOpOperand(0 + isTernary).get() ==
58 body->getArgument(1 + isTernary) &&
59 op->getOpOperand(1 + isTernary).get() ==
60 body->getArgument(0 + isTernary) &&
61 "binary op uses just one block arg");
62 }
63 return swapped;
64}
65
66// Given an elementwise single unary linalg generic op whose body operation is a
67// binary operation, check if one of its operands is a scalar value defined
68// outside the generic op, set its index, and return true. Otherwise return
69// false. The index is unique because the block argument is used at
70// least by one operand, as checked in `isaElemwiseSingleUnaryOpInterface`.
71//
72// Example:
73// %cst = arith.constant 3.14 : f32
74// %0 = linalg.generic { indexing_maps = [#mapA, #mapRes], ... }
75// ins(%A : tensor<?xf32>) outs(...) {
76// ^bb0(%a: f32, %out : f32):
77// %0 = arith.mulf %a, %cst : f32
78// linalg.yield %0: f32
79// } -> tensor<?xf32>
80// Here, the returned index is 1, and the generic op can be represented as
81// %0 = linalg.elementwise <mul>
82// indexing_maps = [#mapA, affine_map<(d0) -> ()>, #mapRes]
83// ins(%A, %cst : tensor<?xf32>, f32) outs(...) -> tensor<?xf32>
84static bool findIndexOfScalarOperand(GenericOp genericOp, int &index) {
85 Block *body = genericOp.getBody();
86 Operation *op = &body->front();
87 for (auto [i, v] : llvm::enumerate(op->getOperands())) {
88 if (auto blockArg = dyn_cast<BlockArgument>(v);
89 blockArg && blockArg.getOwner() == body)
90 continue; // not an outside value...
91 index = i;
92 return true;
93 }
94 return false;
95}
96
97// Maps a linalg.generic body operation to its corresponding elementwise kind,
98// if one exists. Handles the ops with a straightforward one-to-one mapping;
99// ops needing extra context (e.g. boolean add/mul, reciprocal, square) are
100// handled by the caller.
101static std::optional<ElementwiseKind> getElementwiseKind(Operation *op) {
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; });
144}
145
146// Attempt to specialize unary/binary/ternary linalg.generic ops
147// to linalg.elementwise.
148//
149// Example:
150// %0 = linalg.generic {
151// indexing_maps = [affine_map<(d0, d1) -> (d0, d1)>,
152// affine_map<(d0, d1) -> (d0, d1)>],
153// iterator_types = ["parallel", "parallel"]
154// } ins(%In : tensor<?x?xf32>) outs(%Out : tensor<?x?xf32>) {
155// ^bb0(%in: f32, %out: f32):
156// %1 = math.exp %in : f32
157// linalg.yield %1 : f32
158// } -> tensor<?x?xf32>
159//
160// is specialized to
161// linalg.elementwise <exp> ...
162//
163// The category op can carry non-identity indexing maps; these are
164// transferred verbatim from the `genericOp`.
165//
166// In addition to the canonical forms used by the generalization path, this
167// function can handle the following variations:
168//
169// 1) Swapped operands in binary ops (see the `areBinOpsSwapped` helper)
170// 2) Unary generic ops with a binary body op (see the
171// `findIndexOfScalarOperand` helper)
172static FailureOr<LinalgOp> specializeLinalgElementwise(RewriterBase &rewriter,
173 GenericOp genericOp) {
174 // Classify the generic op.
175 unsigned arity = genericOp.getNumDpsInputs();
176 bool isUnary = arity == 1;
177 bool isBinary = arity == 2;
178 bool isTernary = arity == 3;
179
180 // Will inspect the body operation to determine named op or elementwise kind.
181 Operation *op = &genericOp.getBody()->front();
182
183 // Detect variations from canonical forms.
184 bool hasSwappedOperands =
185 (isBinary || isTernary) && areBinOpsSwapped(genericOp, isTernary);
186 int scalarOprIdx = -1;
187 bool hasScalarOperand = isUnary && op->getNumOperands() == 2 &&
188 findIndexOfScalarOperand(genericOp, scalarOprIdx);
189
190 // Helper to dispatch between named op and `linalg.elementwise`.
191 // Lambdas with explicit template parameter list are a C++20 feature, hence
192 // the dummy op object.
193 auto replaceOp = [&](ElementwiseKind kind,
194 bool mayHoistScalarOperand = true) -> LinalgOp {
195 SmallVector<Value> inputs = genericOp.getDpsInputs();
196 SmallVector<AffineMap> indexingMaps = genericOp.getIndexingMapsArray();
197 if (hasSwappedOperands) {
198 // Ternary indices are +1, since the first is the boolean mask.
199 // If new ternary with non-booleans as first argument are created,
200 // we may need to calculate all combinations possible.
201 std::swap(inputs[0 + isTernary], inputs[1 + isTernary]);
202 std::swap(indexingMaps[0 + isTernary], indexingMaps[1 + isTernary]);
203 }
204
205 if (hasScalarOperand && mayHoistScalarOperand) {
206 // Adjust inputs and indexing maps accordingly.
207 inputs.insert(inputs.begin() + scalarOprIdx,
208 op->getOperand(scalarOprIdx));
209 auto scalarBroadcastMap =
210 AffineMap::get(genericOp.getNumParallelLoops(), /*symbolCount=*/0,
211 rewriter.getContext());
212 indexingMaps.insert(indexingMaps.begin() + scalarOprIdx,
213 scalarBroadcastMap);
214 }
215 auto newOp = ElementwiseOp::create(
216 rewriter, genericOp.getLoc(), inputs, genericOp.getDpsInits(), kind,
217 rewriter.getAffineMapArrayAttr(indexingMaps));
218
219 rewriter.replaceOp(genericOp, newOp);
220 return newOp;
221 };
222
223 // Reciprocal
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,
229 /*mayHoistScalarOperand=*/false);
230 }
231
232 // Square
233 if (auto mulOp = dyn_cast<arith::MulFOp>(op))
234 if (mulOp.getLhs() == mulOp.getRhs())
235 return replaceOp(ElementwiseKind::square);
236
237 // Boolean-typed `add` and `mul`.
238 if (isBinary && llvm::all_of(op->getOperands(), [](Value v) {
239 return v.getType().isInteger(1);
240 })) {
241 if (isa<arith::OrIOp>(op))
242 return replaceOp(ElementwiseKind::add);
243 if (isa<arith::AndIOp>(op))
244 return replaceOp(ElementwiseKind::mul);
245 }
246
247 // Table driven
248 if (std::optional<ElementwiseKind> kind = getElementwiseKind(op)) {
249 auto arityGroupAndKind = linalg::getArityGroupAndKind(*kind);
250 // A hoisted scalar operand adds one input to the elementwise op.
251 unsigned numInputs = arity + (hasScalarOperand ? 1 : 0);
252 if (numInputs == static_cast<unsigned>(arityGroupAndKind.arityGroup))
253 return replaceOp(*kind);
254 }
255
256 return rewriter.notifyMatchFailure(
257 genericOp, "elementwise operation cannot be specialized to category op");
258}
259
260//===----------------------------------------------------------------------===//
261// Specialize linalg generic to matmul variants.
262//===----------------------------------------------------------------------===//
263/// Identifies linalg.generic that is essentially named op of the form:
264// ` linalg.{batch_}?matmul{_transpose_a | _transpose_b}? `
265//
266// It is possible that a linalg.generic may be implementing a matmul but not
267// in a straight-forward way e.g. below is matrix multiply over some slice
268// ```
269// %0 = linalg.generic {
270// indexing_maps = [affine_map<(d0, d1, d2) -> (3, d1, d0)>,
271// affine_map<(d0, d1, d2) -> (d0, 5, d2)>,
272// affine_map<(d0, d1, d2) -> (d2, d1, 13)>],
273// iterator_types = ["parallel", "parallel", "parallel"]}
274// ins(%A, %B : tensor<20x20x20xf32>, tensor<20x20x20xf32>)
275// outs(%C : tensor<20x20x20xf32>) {
276// ^bb0(%a: f32, %b: f32, %c : f32):
277// %mul = arith.mulf %a, %b : f32
278// %add = arith.addf %mul, %c : f32
279// linalg.yield %add : f32
280// } -> tensor<20x20x20xf32>
281// ```
282// It is not possible to represent above as named op.
283// e.g. linalg.batch_matmul(%A, %B : tensor<20x20x20xf32>, ...) is
284// not the same as linalg.generic above.
285namespace {
286enum class IndexMatchResult {
287 Match = 0, // identity map.
288 Transposed, // transposed map.
289 Mismatch // none of the above.
290};
291
292// Checks whether the input Affine `map` contains two consecutive dims that
293// can be interpreted as accessing a 2D matrix. It is assumed that the row
294// column dimension are adjacent axis (in this order) and start at
295// `rowDimIdx` in the input map.
296//
297// e.g. consider A matrix in `C[M,N] = A[M,K] * B[K,N]`. We will check
298// whether the map of A is identity (match), transposed, or something
299// completely different (mis-match). Similar for B and C.
300static IndexMatchResult matchOperandMap(AffineMap map, unsigned rowDimIdx,
301 unsigned expectedPosOfRowDim,
302 unsigned expectedPosOfColDim) {
303 // Get the matrix multiply indices. They are past the batch indices.
304 auto exprOfRowDim = map.getResults()[rowDimIdx];
305 auto exprOfColDim = map.getResults()[rowDimIdx + 1];
306
307 // They should be pure dimension ids.
308 if (exprOfRowDim.getKind() != AffineExprKind::DimId ||
309 exprOfColDim.getKind() != AffineExprKind::DimId)
310 return IndexMatchResult::Mismatch;
311
312 auto posRowDim = cast<AffineDimExpr>(exprOfRowDim).getPosition();
313 auto posColDim = cast<AffineDimExpr>(exprOfColDim).getPosition();
314
315 if (expectedPosOfRowDim == posRowDim && expectedPosOfColDim == posColDim)
316 return IndexMatchResult::Match;
317
318 if (expectedPosOfRowDim == posColDim && expectedPosOfColDim == posRowDim)
319 return IndexMatchResult::Transposed;
320
321 return IndexMatchResult::Mismatch;
322}
323
324// Replaces genericOp with `NamedOpTy` op, supplied as a template arg.
325// All the variants expressed as pseudo regular expression:
326// `linalg.{batch_}?matmul` have same number of ins/out, so it's easy to
327// stamp different versions.
328// `castTy` is an optional type function that indicates whether (and which) cast
329// attribute is needed for the named matmul op variant.
330template <typename NamedOpTy>
331static LinalgOp replaceWithMatmulVariant(RewriterBase &rewriter, GenericOp op,
332 std::optional<TypeFn> castTy,
333 ArrayRef<AffineMap> indexingMaps) {
335 // Only explicitly specify the cast attribute for unsigned cast; signed is
336 // the default for linalg.matmul/linalg.batch_matmul.
337 if (castTy.has_value() && *castTy == TypeFn::cast_unsigned) {
338 auto castAttr = rewriter.getNamedAttr(
339 "cast", TypeFnAttr::get(rewriter.getContext(), *castTy));
340 attributes.push_back(castAttr);
341 }
342
343 // Set the original generic's maps to preserve operand indexing semantics like
344 // transposition.
345 SmallVector<Attribute, 3> indexingMapsAttrVal =
346 llvm::map_to_vector(indexingMaps, [](AffineMap map) -> Attribute {
347 return AffineMapAttr::get(map);
348 });
349 auto indexingMapsAttr = rewriter.getNamedAttr(
350 "indexing_maps", rewriter.getArrayAttr(indexingMapsAttrVal));
351 attributes.push_back(indexingMapsAttr);
352
353 LinalgOp namedOp = rewriter.replaceOpWithNewOp<NamedOpTy>(
354 op, ValueRange{op.getDpsInputs()[0], op.getDpsInputs()[1]},
355 ValueRange{op.getDpsInits()[0]}, attributes);
356
357 return namedOp;
358}
359
360// Returns the cast type to use for a matmul-like named op. If the generic
361// contains casts that cannot be represented (e.g. output casts or mixed
362// signedness), return std::nullopt.
363static std::optional<TypeFn> getCastTypeForMatmulLikeOp(GenericOp genericOp) {
364 bool foundCastForMatmulOutput = false;
365 SmallVector<TypeFn> castTyFns;
366 genericOp.getBody()->walk([&](CastOpInterface castOp) {
367 // Collect forward slice of the cast op to check if it is for the matmul
368 // output.
369 SetVector<Operation *> forwardSlice;
370 getForwardSlice(castOp, &forwardSlice);
371
372 // If there is no multiplication op in the forward slice, then this cast
373 // op is for the matmul output. Cast ops on matmul output cannot be
374 // expressed by the matmul op variant.
375 if (!llvm::any_of(forwardSlice, [](Operation *op) {
376 // We check explicitly for these multiplication ops in
377 // `specializeLinalgContractions()` to infer matmul-like ops.
378 return isa<arith::MulIOp, arith::MulFOp, complex::MulOp>(op);
379 })) {
380 foundCastForMatmulOutput = true;
381 return WalkResult::interrupt();
382 }
383
384 // Determine the cast type.
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);
389
390 return WalkResult::advance();
391 });
392
393 if (foundCastForMatmulOutput)
394 return std::nullopt;
395
396 if (!castTyFns.empty()) {
397 // If there were multiple different cast types found, then we can't express
398 // them using matmul-like ops. They only allow a single cast type for all
399 // inputs.
400 if (!llvm::all_equal(castTyFns))
401 return std::nullopt;
402 return castTyFns.front();
403 }
404
405 // Default to signed cast for matmul-like ops.
406 return TypeFn::cast_signed;
407}
408
409static FailureOr<LinalgOp> specializeLinalgMmt4D(RewriterBase &rewriter,
410 GenericOp genericOp,
411 std::optional<TypeFn> castTy,
412 ContractionDimensions &dims) {
413 // Should all be rank 4 and dim 6
414 auto indexingMaps = genericOp.getIndexingMapsArray();
415 if (llvm::any_of(indexingMaps, [](AffineMap m) {
416 return m.getResults().size() != 4 || m.getNumDims() != 6;
417 }))
418 return failure();
419
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]);
422
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]);
425
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]);
428
429 if (llvm::is_contained({aOuter, bOuter, cOuter}, IndexMatchResult::Mismatch))
430 return failure();
431 if (llvm::is_contained({aInner, bInner, cInner}, IndexMatchResult::Mismatch))
432 return failure();
433
434 SmallVector<AffineMap> namedOpMaps = {indexingMaps[0], indexingMaps[1],
435 indexingMaps[2]};
436
437 return replaceWithMatmulVariant<Mmt4DOp>(rewriter, genericOp, castTy,
438 namedOpMaps);
439}
440
441static bool isSupportedContractionPair(Operation *first, Operation *second) {
442 if (isa<arith::MulFOp>(first) && isa<arith::AddFOp>(second))
443 return true;
444 if (isa<arith::MulIOp>(first) && isa<arith::AddIOp>(second))
445 return true;
446 if (isa<complex::MulOp>(first) && isa<complex::AddOp>(second))
447 return true;
448 if (isa<arith::AndIOp>(first) && isa<arith::OrIOp>(second) &&
449 first->getResult(0).getType().isInteger(1))
450 return true;
451
452 return false;
453}
454
455// Attempts to specialize `genericOp` to a specific named matmul variant
456// (`matmul`, `batch_matmul`, or `mmt4d`). Returns failure without modifying the
457// IR if no named variant matches. `castTy` is the cast type inferred for the
458// contraction body.
459static FailureOr<LinalgOp>
460specializeToNamedContraction(RewriterBase &rewriter, GenericOp genericOp,
461 std::optional<TypeFn> castTy) {
462 // Linalg generic contraction can be across multiple axis e.g.
463 // ```
464 // linalg.generic
465 // {indexing_maps = [affine_map<(m, n, k1, k2) -> (m, k1, k2)>,
466 // affine_map<(m, n, k1, k2) -> (k2, k1, n)>,
467 // affine_map<(m, n, k1, k2) -> (m, n)>],
468 // iterator_types = ["parallel", "parallel",
469 // "reduction", "reduction"]}
470 // ins(%A, %B : tensor<10x20x30xf32>, tensor<30x20x40xf32>)
471 // outs(%C : tensor<10x40xf32>) {
472 // ^bb0(%a: f32, %b: f32, %c: f32):
473 // %1 = arith.mulf %a, %b : f32
474 // %2 = arith.addf %c, %1 : f32
475 // linalg.yield %2 : f32
476 // } -> tensor<10x40xf32>
477 // ```
478 // In above contraction, there are two reduction dimensions {k1, k2}
479 // and although a valid linalg contraction, it is not a named-op
480 // matrix multiply kind. Therefore, reject multi-dim reduction.
481 auto res = inferContractionDims(genericOp);
482 if (!succeeded(res))
483 return failure();
484 auto dims = *res;
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)
488 return failure();
489
490 // Check rank of operands
491 auto indexingMaps = genericOp.getIndexingMapsArray();
492 if (llvm::any_of(indexingMaps, [&dims](AffineMap m) {
493 return m.getResults().size() !=
494 dims.batch.size() + 2 /* any two of {m,n,k} */;
495 }))
496 return failure();
497
498 auto numOfBatchDims = dims.batch.size();
499 if (indexingMaps[0].getNumDims() != numOfBatchDims + 3)
500 return failure();
501
502 if (numOfBatchDims) {
503 // Each operand in a linalg generic contraction could express different
504 // permutations for its batch dimension. But for named op it must be
505 // identity since separate maps are not specified.
506 if (llvm::any_of(indexingMaps, [numOfBatchDims](AffineMap m) {
507 for (unsigned i = 0; i < numOfBatchDims; ++i) {
508 auto expr = m.getResults()[i];
509 if (expr.getKind() != AffineExprKind::DimId ||
510 cast<AffineDimExpr>(expr).getPosition() != i)
511 return true;
512 }
513 return false;
514 }))
515 return failure();
516 }
517
518 auto a =
519 matchOperandMap(indexingMaps[0], numOfBatchDims, dims.m[0], dims.k[0]);
520 auto b =
521 matchOperandMap(indexingMaps[1], numOfBatchDims, dims.k[0], dims.n[0]);
522 auto c =
523 matchOperandMap(indexingMaps[2], numOfBatchDims, dims.m[0], dims.n[0]);
524
525 if (llvm::is_contained({a, b, c}, IndexMatchResult::Mismatch))
526 return failure();
527
528 // Build indexing maps for the named op in its canonical dimension ordering
529 auto *ctx = genericOp.getContext();
530 unsigned numLoopDims = numOfBatchDims + 3;
531 unsigned mIdx = numOfBatchDims;
532 unsigned nIdx = mIdx + 1;
533 unsigned kIdx = mIdx + 2;
534
535 // TODO: add support for indexing_maps with broadcasts.
536 auto makeMap = [&](IndexMatchResult match, unsigned rowIdx, unsigned colIdx) {
537 SmallVector<unsigned> tensorDims;
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);
542 else
543 llvm::append_values(tensorDims, rowIdx, colIdx);
544 return AffineMap::getMultiDimMapWithTargets(numLoopDims, tensorDims, ctx);
545 };
546
547 auto mapA = makeMap(a, mIdx, kIdx);
548 auto mapB = makeMap(b, kIdx, nIdx);
549 auto mapC = makeMap(c, mIdx, nIdx);
550
551 SmallVector<AffineMap> namedOpMaps = {mapA, mapB, mapC};
552
553 // Codegen the different matmul variants.
554 if (numOfBatchDims) {
555 return replaceWithMatmulVariant<BatchMatmulOp>(rewriter, genericOp, castTy,
556 namedOpMaps);
557 }
558 return replaceWithMatmulVariant<MatmulOp>(rewriter, genericOp, castTy,
559 namedOpMaps);
560}
561
562// Converts linalg.generic to named linalg.*matmul* where possible, falling
563// back to the generic `linalg.contract` op otherwise.
564static FailureOr<LinalgOp> specializeLinalgContractions(RewriterBase &rewriter,
565 GenericOp genericOp,
566 bool emitCategoryOp) {
567 if (genericOp.getNumDpsInputs() != 2 || genericOp.getNumDpsInits() != 1)
568 return failure();
569
570 // Early exit if not projected permutations.
571 auto mapRange = genericOp.getIndexingMapsArray();
572 if (llvm::any_of(mapRange,
573 [](AffineMap m) { return !m.isProjectedPermutation(); }))
574 return failure();
575
576 // Only contractions that can be represented by named linalg ops are
577 // eligible for specialization:
578 // - mul + add (floating-point, integer, complex)
579 // - and + or (bool)
580 if (!mlir::linalg::detail::isContractionBody(*genericOp.getBlock(),
581 isSupportedContractionPair))
582 return failure();
583
584 // Determine the cast type for the named matmul op, or bail out if casts
585 // cannot be represented by the named op.
586 std::optional<TypeFn> castTy = getCastTypeForMatmulLikeOp(genericOp);
587 if (!castTy)
588 return rewriter.notifyMatchFailure(
589 genericOp, "contains invalid cast ops for the named matmul op");
590
591 // TODO: When `emitCategoryOp` is set, skip the named-variant matching below
592 // and go straight to the `linalg.contract` category op.
593
594 // Try to specialize to a specific named matmul variant first.
595 if (!emitCategoryOp) {
596 FailureOr<LinalgOp> namedOp =
597 specializeToNamedContraction(rewriter, genericOp, castTy);
598 if (succeeded(namedOp))
599 return namedOp;
600 }
601
602 // No named variant matched; fall back to the generic `linalg.contract` op,
603 // which supports a wider range of variants.
604 return replaceWithMatmulVariant<ContractOp>(rewriter, genericOp, castTy,
605 genericOp.getIndexingMapsArray());
606}
607
608/// Utility to specialize a `genericOp` with a convolution op of type `ConvOpTy`
609/// with `dilations` and `strides`.
610template <typename ConvOpTy>
611static FailureOr<LinalgOp>
612specializeToConvOp(RewriterBase &rewriter, GenericOp genericOp,
613 ArrayRef<int64_t> dilations, ArrayRef<int64_t> strides) {
614 SmallVector<Value> inputs = genericOp.getDpsInputs();
615 ValueRange outputs = genericOp.getDpsInits();
616 SmallVector<Type> resultTypes = genericOp.hasPureTensorSemantics()
617 ? TypeRange(ValueRange(outputs))
618 : TypeRange{};
619 LinalgOp namedOp;
620 // Ops with no dilations and no strides.
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>) {
624 namedOp = rewriter.replaceOpWithNewOp<ConvOpTy>(genericOp, resultTypes,
625 inputs, outputs);
626 } else {
627 Attribute stridesAttr = rewriter.getI64TensorAttr(strides);
628 Attribute dilationsAttr = rewriter.getI64TensorAttr(dilations);
629 namedOp = rewriter.replaceOpWithNewOp<ConvOpTy>(
630 genericOp, resultTypes, inputs, outputs, stridesAttr, dilationsAttr);
631 }
632 return namedOp;
633}
634
635/// Converts linalg.generic to named linalg.*conv/pooling* where possible.
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); \
643 // -----------------------------
644 // Convolution ops.
645 // -----------------------------
646 CONV_OP_SPECIALIZER(linalg::Conv1DOp);
647 CONV_OP_SPECIALIZER(linalg::Conv1DNwcWcfOp);
648 CONV_OP_SPECIALIZER(linalg::Conv1DNcwFcwOp);
649 CONV_OP_SPECIALIZER(linalg::Conv2DOp);
650 CONV_OP_SPECIALIZER(linalg::Conv2DNhwcHwcfOp);
651 CONV_OP_SPECIALIZER(linalg::Conv2DNhwcHwcfQOp);
652 CONV_OP_SPECIALIZER(linalg::Conv2DNhwcFhwcOp);
653 CONV_OP_SPECIALIZER(linalg::Conv2DNhwcFhwcQOp);
654 CONV_OP_SPECIALIZER(linalg::Conv2DNchwFchwOp);
655 CONV_OP_SPECIALIZER(linalg::Conv2DNchwFchwQOp);
656 CONV_OP_SPECIALIZER(linalg::Conv2DNgchwFgchwOp);
657 CONV_OP_SPECIALIZER(linalg::Conv2DNgchwGfchwOp);
658 CONV_OP_SPECIALIZER(linalg::Conv2DNgchwGfchwQOp);
659 CONV_OP_SPECIALIZER(linalg::Conv2DNhwgcGfhwcOp);
660 CONV_OP_SPECIALIZER(linalg::Conv2DNhwgcGfhwcQOp);
661 CONV_OP_SPECIALIZER(linalg::Conv3DOp);
662 CONV_OP_SPECIALIZER(linalg::Conv3DNdhwcDhwcfOp);
663 CONV_OP_SPECIALIZER(linalg::Conv3DNdhwcDhwcfQOp);
664 CONV_OP_SPECIALIZER(linalg::Conv3DNcdhwFcdhwOp);
665 // -----------------------------
666 // Depthwise Convolution ops.
667 // -----------------------------
668 CONV_OP_SPECIALIZER(linalg::DepthwiseConv1DNcwCwOp);
669 CONV_OP_SPECIALIZER(linalg::DepthwiseConv1DNwcWcOp);
670 CONV_OP_SPECIALIZER(linalg::DepthwiseConv1DNwcWcmOp);
671 CONV_OP_SPECIALIZER(linalg::DepthwiseConv2DNchwChwOp);
672 CONV_OP_SPECIALIZER(linalg::DepthwiseConv2DNhwcHwcOp);
673 CONV_OP_SPECIALIZER(linalg::DepthwiseConv2DNhwcHwcQOp);
674 CONV_OP_SPECIALIZER(linalg::DepthwiseConv2DNhwcHwcmOp);
675 CONV_OP_SPECIALIZER(linalg::DepthwiseConv2DNhwcHwcmQOp);
676 CONV_OP_SPECIALIZER(linalg::DepthwiseConv3DNdhwcDhwcOp);
677 CONV_OP_SPECIALIZER(linalg::DepthwiseConv3DNcdhwCdhwOp);
678 CONV_OP_SPECIALIZER(linalg::DepthwiseConv3DNdhwcDhwcmOp);
679 // -----------------------------
680 // Pooling ops.
681 // -----------------------------
682 CONV_OP_SPECIALIZER(linalg::PoolingNhwcMaxOp);
683 CONV_OP_SPECIALIZER(linalg::PoolingNhwcMinOp);
684 CONV_OP_SPECIALIZER(linalg::PoolingNhwcSumOp);
685 CONV_OP_SPECIALIZER(linalg::PoolingNhwcMaxUnsignedOp);
686 CONV_OP_SPECIALIZER(linalg::PoolingNhwcMinUnsignedOp);
687 CONV_OP_SPECIALIZER(linalg::PoolingNchwSumOp);
688 CONV_OP_SPECIALIZER(linalg::PoolingNchwMaxOp);
689 CONV_OP_SPECIALIZER(linalg::PoolingNwcSumOp);
690 CONV_OP_SPECIALIZER(linalg::PoolingNcwSumOp);
691 CONV_OP_SPECIALIZER(linalg::PoolingNwcMaxOp);
692 CONV_OP_SPECIALIZER(linalg::PoolingNwcMaxUnsignedOp);
693 CONV_OP_SPECIALIZER(linalg::PoolingNcwMaxOp);
694 CONV_OP_SPECIALIZER(linalg::PoolingNwcMinOp);
695 CONV_OP_SPECIALIZER(linalg::PoolingNwcMinUnsignedOp);
696 CONV_OP_SPECIALIZER(linalg::PoolingNdhwcSumOp);
697 CONV_OP_SPECIALIZER(linalg::PoolingNdhwcMaxOp);
698 CONV_OP_SPECIALIZER(linalg::PoolingNdhwcMinOp);
699#undef CONV_OP_SPECIALIZER
700 return failure();
701}
702
703} // namespace
704
705//===----------------------------------------------------------------------===//
706// Categorize linalg generic to named op where possible.
707//===----------------------------------------------------------------------===//
709 GenericOp genericOp,
710 bool emitCategoryOps) {
711 // Elementwise - e.g. exp, add are always category ops
712 if (isaElemwiseSingleUnaryOpInterface(genericOp) ||
715 return specializeLinalgElementwise(rewriter, genericOp);
716 }
717
718 // Contraction - e.g. matmul
719 if (isaContractionOpInterface(genericOp)) {
720 return specializeLinalgContractions(rewriter, genericOp, emitCategoryOps);
721 }
722
723 // Early exit in case of category specialization.
724 // TODO: Remove when matches for other ops account for both named and
725 // category.
726 if (emitCategoryOps)
727 return rewriter.notifyMatchFailure(
728 genericOp, "no matching category op specialization");
729
730 // Copy
731 if (isaCopyOpInterface(genericOp)) {
732 LinalgOp namedOp = rewriter.replaceOpWithNewOp<CopyOp>(
733 genericOp, genericOp.getDpsInputs()[0], genericOp.getDpsInits()[0]);
734 return namedOp;
735 }
736
737 // Fill
738 if (std::optional<Value> fillValue = isaFillOpInterface(genericOp)) {
739 // Always use the detected fill value, regardless of pattern
740 LinalgOp namedOp = rewriter.replaceOpWithNewOp<FillOp>(
741 genericOp, *fillValue, genericOp.getDpsInits()[0]);
742 return namedOp;
743 }
744
745 // Broadcast
746 std::optional<SmallVector<int64_t>> equivalentToBroadcast =
747 isaBroadcastOpInterface(genericOp);
748 if (equivalentToBroadcast) {
749 auto dims = *equivalentToBroadcast;
750 LinalgOp namedOp = rewriter.replaceOpWithNewOp<BroadcastOp>(
751 genericOp, genericOp.getDpsInputs()[0], genericOp.getDpsInits()[0],
752 dims);
753 return namedOp;
754 }
755
756 // Transpose
757 std::optional<SmallVector<int64_t>> equivalentToTranspose =
758 isaTransposeOpInterface(genericOp);
759 if (equivalentToTranspose) {
760 auto permutation = *equivalentToTranspose;
761 LinalgOp namedOp = rewriter.replaceOpWithNewOp<TransposeOp>(
762 genericOp, genericOp.getDpsInputs()[0], genericOp.getDpsInits()[0],
763 permutation);
764 return namedOp;
765 }
766
767 // Convolution - e.g. *conv/pooling*
768 if (isaConvolutionOpInterface(genericOp))
769 return specializeLinalgConvolutions(rewriter, genericOp);
770
771 return rewriter.notifyMatchFailure(genericOp,
772 "no matching named op specialization");
773}
774
775namespace {
776struct LinalgSpecializeGenericOpsPass
777 : public impl::LinalgSpecializeGenericOpsPassBase<
778 LinalgSpecializeGenericOpsPass> {
779
780 using impl::LinalgSpecializeGenericOpsPassBase<
781 LinalgSpecializeGenericOpsPass>::LinalgSpecializeGenericOpsPassBase;
782 void runOnOperation() override;
783};
784} // namespace
785
786void LinalgSpecializeGenericOpsPass::runOnOperation() {
787 RewritePatternSet patterns(&getContext());
790
791 if (failed(applyPatternsGreedily(getOperation(), std::move(patterns))))
792 signalPassFailure();
793}
794
796 RewritePatternSet &patterns, bool emitCategoryOps) {
797 patterns.add<LinalgSpecializationPattern>(patterns.getContext(),
798 emitCategoryOps);
799}
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
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.
Definition AffineMap.h:46
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.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
BlockArgument getArgument(unsigned i)
Definition Block.h:154
Operation & front()
Definition Block.h:178
DenseIntElementsAttr getI64TensorAttr(ArrayRef< int64_t > values)
Definition Builders.cpp:194
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
MLIRContext * getContext() const
Definition Builders.h:56
NamedAttribute getNamedAttr(StringRef name, Attribute val)
Definition Builders.cpp:102
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
Definition Builders.cpp:327
IRValueT get() const
Return the current value being used by this operand.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
unsigned getNumOperands()
Definition Operation.h:371
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
OpOperand & getOpOperand(unsigned idx)
Definition Operation.h:413
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.
Definition TypeRange.h:40
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
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
Type getType() const
Return the type of this value.
Definition Value.h:105
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
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....
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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...
@ DimId
Dimensional identifier.
Definition AffineExpr.h:59
llvm::SetVector< T, Vector, Set, N > SetVector
Definition LLVM.h:125
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