MLIR 24.0.0git
LinalgInterfaces.cpp
Go to the documentation of this file.
1//===- LinalgInterfaces.cpp - Linalg interfaces implementation ------------===//
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
10
16#include "mlir/IR/AffineExpr.h"
18#include "mlir/IR/AffineMap.h"
20#include "mlir/IR/MLIRContext.h"
22#include "llvm/ADT/STLExtras.h"
23#include "llvm/ADT/SetOperations.h"
24#include "llvm/ADT/SmallBitVector.h"
25#include "llvm/ADT/SmallVector.h"
26#include "llvm/Support/Casting.h"
27#include "llvm/Support/raw_ostream.h"
28#include <optional>
29
30using namespace mlir;
31using namespace mlir::linalg;
32
33/// Include the definitions of the copy operation interface.
34#include "mlir/Dialect/Linalg/IR/LinalgInterfaces.cpp.inc"
35
36//===----------------------------------------------------------------------===//
37// Interface utility functions
38//===----------------------------------------------------------------------===//
39
41 linalg::LinalgOp linalgOp, ArrayRef<OpOperand *> droppedOperands) {
42 SmallVector<AffineMap> indexingMaps;
43 for (auto &opOperand : linalgOp->getOpOperands()) {
44 if (llvm::is_contained(droppedOperands, &opOperand))
45 continue;
46 indexingMaps.push_back(linalgOp.getMatchingIndexingMap(&opOperand));
47 }
48 if (indexingMaps.empty()) {
49 // If there are no indexing maps, the operand can only be dropped
50 // if the op has no loops.
51 return linalgOp.getNumLoops() == 0;
52 }
54 indexingMaps, linalgOp.getContext())) != AffineMap();
55}
56
57//===----------------------------------------------------------------------===//
58// CopyOpInterface implementation
59//===----------------------------------------------------------------------===//
60
61bool linalg::isaCopyOpInterface(LinalgOp op) {
62 // Check all loops are parallel and linalgOp is single input and output.
63 if (!op.isAllParallelLoops() || !op.isSingleInputOutput())
64 return false;
65
66 auto mapRange = op.getIndexingMapsArray();
67 if (mapRange.size() != 2 || !mapRange.front().isIdentity() ||
68 !mapRange.back().isIdentity()) {
69 return false;
70 }
71 // Check yield first block argument.
72 Block *body = op.getBlock();
73 if (body->getOperations().size() != 1)
74 return false;
75 auto yieldOp = dyn_cast<linalg::YieldOp>(body->back());
76 if (!yieldOp || yieldOp.getNumOperands() != 1)
77 return false;
78 return yieldOp->getOperand(0) == body->getArgument(0);
79}
80
81//===----------------------------------------------------------------------===//
82// FillOpInterface implementation
83//===----------------------------------------------------------------------===//
84/// Detects if a linalg.generic operation represents a fill with an inlined
85/// constant. If so, returns the constant value. Otherwise, returns
86/// std::nullopt.
87static std::optional<Value> isaInlinedFillOp(GenericOp op) {
88 if (!op.isAllParallelLoops() || op.getNumDpsInits() != 1 ||
89 op.getNumDpsInputs() != 0)
90 return std::nullopt;
91
92 // Init should not be referenced.
93 if (op.payloadUsesValueFromOperand(op.getDpsInitOperand(0)))
94 return std::nullopt;
95
96 Block *body = op.getBody();
97 if (body->getOperations().size() != 1)
98 return std::nullopt;
99
100 auto yieldOp = dyn_cast<linalg::YieldOp>(body->back());
101 if (!yieldOp || yieldOp.getNumOperands() != 1)
102 return std::nullopt;
103
104 Value yieldOperand = yieldOp->getOperand(0);
105 if (!yieldOperand.getDefiningOp<arith::ConstantOp>() &&
106 !yieldOperand.getDefiningOp<complex::ConstantOp>())
107 return std::nullopt;
108
109 return yieldOperand;
110}
111
112/// Detects if a linalg.generic operation represents an external scalar input.
113/// If so, returns the constant value. Otherwise, returns std::nullopt.
114static std::optional<Value> isaExternalFillOp(GenericOp op) {
115 // Structural.
116 if (!op.isAllParallelLoops() || !op.isSingleInputOutput() ||
117 !op.isSingleYieldOp())
118 return std::nullopt;
119
120 // Input should be referenced and init should not.
121 if (!op.payloadUsesValueFromOperand(op.getDpsInputOperand(0)) ||
122 op.payloadUsesValueFromOperand(op.getDpsInitOperand(0)))
123 return std::nullopt;
124
125 OpOperand *value = op.getDpsInputOperand(0);
126 if (!op.isScalar(value))
127 return std::nullopt;
128 return value->get();
129}
130
131std::optional<Value> linalg::isaFillOpInterface(GenericOp op) {
132 if (auto fillVal = isaInlinedFillOp(op))
133 return fillVal;
134 return isaExternalFillOp(op);
135}
136
137//===----------------------------------------------------------------------===//
138// BroadcastOpInterface implementation
139//===----------------------------------------------------------------------===//
140std::optional<SmallVector<int64_t>>
142 if (auto broadcastOp = dyn_cast<BroadcastOp>(linalgOp.getOperation()))
143 return SmallVector<int64_t>(broadcastOp.getDimensions().begin(),
144 broadcastOp.getDimensions().end());
145
146 auto op = dyn_cast<GenericOp>(linalgOp.getOperation());
147 if (!op)
148 return std::nullopt;
149
150 // Structural.
151 if (!op.isAllParallelLoops() || !op.isSingleInputOutput() ||
152 !op.isSingleYieldOp())
153 return std::nullopt;
154
155 auto srcTy = op.getDpsInputOperand(0)->get().getType();
156 auto dstTy = op.getDpsInitOperand(0)->get().getType();
157 if (!isa<MemRefType, RankedTensorType>(srcTy) ||
158 !isa<MemRefType, RankedTensorType>(dstTy))
159 return std::nullopt;
160
161 // Check output is identity map. Broadcast could additionally be
162 // employing permutation of indices and that would be expressible
163 // in linalg.generic but is not expressible for named broadcast op.
164 auto dstMap = op.getIndexingMapsArray()[1];
165 if (!dstMap.isIdentity())
166 return std::nullopt;
167
168 SmallVector<int64_t> position;
169 auto srcMap = op.getIndexingMapsArray()[0];
170
171 if (srcMap.getResults().size() >= dstMap.getResults().size())
172 return std::nullopt;
173
174 // Check input map is monotonically increasing DimIds.
175 for (unsigned i = 0; i < srcMap.getNumResults(); ++i) {
176 auto expr = llvm::dyn_cast<AffineDimExpr>(srcMap.getResults()[i]);
177 if (!expr)
178 return std::nullopt;
179 int64_t pos = expr.getPosition();
180 if (i > 0 && pos <= position[i - 1])
181 return std::nullopt;
182 position.push_back(expr.getPosition());
183 }
184
185 SmallVector<int64_t> broadcastedDims;
186 auto numDims = srcMap.getNumDims();
187 // This is quadratic but number of items is generally small.
188 for (auto dim : llvm::seq<int64_t>(0, numDims)) {
189 if (!llvm::is_contained(position, dim))
190 broadcastedDims.push_back(dim);
191 }
192 return broadcastedDims;
193}
194
195//===----------------------------------------------------------------------===//
196// TransposeOpInterface implementation
197//===----------------------------------------------------------------------===//
198std::optional<SmallVector<int64_t>>
200 // To specialize as a transpose op, the genericOp must be
201 // all parallel loops, single input, single output, and its body
202 // should be just a yield op, yielding input as output as is (no compute).
203 if (!op.isAllParallelLoops() || !op.isSingleInputOutput() ||
204 !op.isSingleYieldOp())
205 return std::nullopt;
206
207 auto mapRange = op.getIndexingMapsArray();
208 if (mapRange.size() != 2)
209 return std::nullopt;
210
211 auto mapOfInput = mapRange.front();
212 auto mapOfResult = mapRange.back();
213
214 // linalg.transpose permutes the dimensions of input using this
215 // rule: dim(result, i) = dim(input, permutation[i])
216 if (!mapOfResult.isIdentity() || !mapOfInput.isPermutation())
217 return std::nullopt;
218
219 SmallVector<int64_t> permutation(mapOfInput.getNumDims());
220 for (unsigned i = 0; i < mapOfInput.getNumDims(); ++i) {
221 auto expr = llvm::cast<AffineDimExpr>(mapOfInput.getResults()[i]);
222 permutation[expr.getPosition()] = i;
223 }
224 return permutation;
225}
226
227//===----------------------------------------------------------------------===//
228// Elementwise Single Unary/Binary-OpInterface implementation
229//===----------------------------------------------------------------------===//
230
231static bool isaElemwiseSingleOpInterface(linalg::GenericOp op, unsigned arity) {
232 // Check all loops are parallel.
233 if (!op.isAllParallelLoops() || op.getNumLoops() < 1)
234 return false;
235
236 // Check there are arity-inputs and 1-output. Non-identity indexing maps are
237 // allowed as they can be represented by the category op.
238 if (op.getNumDpsInputs() != arity || op.getNumDpsInits() != 1)
239 return false;
240
241 // Init should not be referenced for elementwise operations.
242 if (op.payloadUsesValueFromOperand(op.getDpsInitOperand(0)))
243 return false;
244
245 // A linalg.generic could be series of elementwise ops e.g. exp(neg(x)) such
246 // as resulting from producer-consumer fusion. Here, we restrict to two ops in
247 // the body, where the first is the elementwise single op and the second a
248 // yield.
249 Block *body = op.getBody();
250 if (body->getOperations().size() != 2)
251 return false;
252
253 // The payload op must have one result and at least arity-many operands
254 // (otherwise not all inputs can be used). It can have additional operands
255 // from outside of the generic op (e.g. div(1, x) for elementwise reciprocal)
256 // or use an input more than once (e.g. mul(x, x) for elementwise square).
257 Operation *oper = &body->front();
258 if (oper->getNumOperands() < arity || oper->getNumResults() != 1)
259 return false;
260
261 auto yieldOp = dyn_cast<linalg::YieldOp>(body->back());
262 return !(!yieldOp || yieldOp.getNumOperands() != 1 ||
263 yieldOp->getOperand(0).getDefiningOp() != oper);
264}
265
266bool linalg::isaElemwiseSingleUnaryOpInterface(linalg::GenericOp op) {
267 // All basic elemwise checks.
269 return false;
270
271 // Check input is actually used.
272 if (!op.payloadUsesValueFromOperand(op.getDpsInputOperand(0)))
273 return false;
274 return true;
275}
276
277bool linalg::isaElemwiseSingleBinaryOpInterface(linalg::GenericOp op) {
278 // All basic elemwise checks.
280 return false;
281
282 // Check both inputs are used (elementwise).
283 OpOperand *inputOpOperand0 = op.getDpsInputOperand(0);
284 OpOperand *inputOpOperand1 = op.getDpsInputOperand(1);
285 return !(!op.payloadUsesValueFromOperand(inputOpOperand0) ||
286 !op.payloadUsesValueFromOperand(inputOpOperand1));
287}
288
289bool linalg::isaElemwiseSingleTernaryOpInterface(linalg::GenericOp op) {
290 // All basic elemwise checks.
292 return false;
293
294 // The only ternary (select) has a boolean argument as its first operand.
295 // If we add more ternaries later, we need to change this check.
296 // But for now, it simplifies other checks, like checking for swapped
297 // operands.
298 if (!getElementTypeOrSelf(op.getDpsInputOperand(0)->get().getType())
299 .isInteger(1))
300 return false;
301
302 // Check all three inputs are used (elementwise).
303 OpOperand *inputOpOperand0 = op.getDpsInputOperand(0);
304 OpOperand *inputOpOperand1 = op.getDpsInputOperand(1);
305 OpOperand *inputOpOperand2 = op.getDpsInputOperand(2);
306 return !(!op.payloadUsesValueFromOperand(inputOpOperand0) ||
307 !op.payloadUsesValueFromOperand(inputOpOperand1) ||
308 !op.payloadUsesValueFromOperand(inputOpOperand2));
309}
310
311//===----------------------------------------------------------------------===//
312// ContractionOpInterface implementation
313//===----------------------------------------------------------------------===//
314
315/// If the value is defined by a chain of unary side effect-free, go up the
316/// use-def chain until the first value that isn't defined by such an op.
317// TODO: relax to multi-operands with constants, which are technically unary ops
318// as needed (e.g. add5).
320 Operation *op = value.getDefiningOp();
321 while (op && op->getNumOperands() == 1) {
322 auto iface = dyn_cast<MemoryEffectOpInterface>(op);
323 if (!iface || !iface.hasNoEffect())
324 break;
325 value = op->getOperand(0);
326 op = value.getDefiningOp();
327 }
328 return value;
329}
330
332 Block &block, function_ref<bool(Operation *, Operation *)> isaPair,
333 llvm::raw_ostream &errs) {
334 if (block.empty() || !block.back().mightHaveTrait<OpTrait::IsTerminator>()) {
335 errs << "no terminator in the block";
336 return false;
337 }
338
339 if (block.getNumArguments() != 3) {
340 errs << "expected block with 3 arguments";
341 return false;
342 }
343
344 Operation *terminator = block.getTerminator();
345 if (terminator->getNumOperands() != 1) {
346 errs << "expected terminator with 1 operand";
347 return false;
348 }
349
350 Value yielded = getSourceSkipUnary(terminator->getOperand(0));
351 Operation *reductionOp = yielded.getDefiningOp();
352 if (!reductionOp || reductionOp->getNumResults() != 1 ||
353 reductionOp->getNumOperands() != 2) {
354 errs << "expected reduction op to be binary";
355 return false;
356 }
357
358 Value reductionLHS = getSourceSkipUnary(reductionOp->getOperand(0));
359 Value reductionRHS = getSourceSkipUnary(reductionOp->getOperand(1));
360
361 if (reductionLHS != block.getArgument(2) &&
362 reductionRHS != block.getArgument(2)) {
363 errs << "expected reduction to take block argument #2 as one of the "
364 "operands (modulo unary casts)";
365 return false;
366 }
367
368 Value contributed = getSourceSkipUnary(
369 isa<BlockArgument>(reductionLHS) ? reductionRHS : reductionLHS);
370 Operation *elementwiseOp = contributed.getDefiningOp();
371 if (!elementwiseOp || elementwiseOp->getNumResults() != 1 ||
372 elementwiseOp->getNumOperands() != 2) {
373 errs << "expected elementwise op to be binary";
374 return false;
375 }
376
377 if (!isaPair(elementwiseOp, reductionOp)) {
378 errs << "expected reduction/elementwise op kind not satisfied";
379 return false;
380 }
381
382 Value elementwiseLHS = getSourceSkipUnary(elementwiseOp->getOperand(0));
383 Value elementwiseRHS = getSourceSkipUnary(elementwiseOp->getOperand(1));
384 if ((elementwiseLHS == block.getArgument(0) &&
385 elementwiseRHS == block.getArgument(1)) ||
386 (elementwiseLHS == block.getArgument(1) &&
387 elementwiseRHS == block.getArgument(0))) {
388 return true;
389 }
390
391 errs << "expected elementwise op to apply to block arguments (modulo unary "
392 "casts)";
393 return false;
394}
395
396/// Returns true if the two operations are of the kinds specified by a pair of
397/// consecutive template arguments.
398template <typename AddOpTy, typename MulOpTy, typename... Args>
400 static_assert(sizeof...(Args) % 2 == 0,
401 "expected an even number of template arguments");
402 if (isa<AddOpTy>(add) && isa<MulOpTy>(mul))
403 return true;
404
405 if constexpr (sizeof...(Args) > 0)
407 else
408 return false;
409}
410
411/// Returns true if the block is a body of a contraction with the kinds of
412/// operations given pairwise by template arguments.
413template <typename... Args>
417
418/// Given an `indexingMap` and its corresponding `iterators`, returns
419/// the positions of the iterators of type `iter` that are indexed by
420/// the `indexingMap` as a permutation. This is useful to infer various
421/// subcomputations on a `LinalgOp`. This is performed by looking up
422/// each result in the `indexingMap` and determining whether:
423/// - It is a single AffineDimExpr.
424/// - It is the only result involving this AffineDimExpr.
425static llvm::SmallDenseSet<int64_t>
428 utils::IteratorType iter) {
429 assert(iterators.size() == indexingMap.getNumDims());
430 llvm::SmallDenseSet<int64_t> res;
431 for (AffineExpr e : indexingMap.getResults()) {
432 if (auto d = dyn_cast<AffineDimExpr>(e)) {
433 if (iterators[d.getPosition()] == iter &&
434 llvm::count_if(indexingMap.getResults(), [d](AffineExpr e) {
435 return e.isFunctionOfDim(d.getPosition());
436 }) == 1)
437 res.insert(d.getPosition());
438 }
439 }
440 return res;
441}
442
443namespace {
444auto par = utils::IteratorType::parallel;
445auto red = utils::IteratorType::reduction;
446} // namespace
447
448/// Infer the iterator types from the init affine map. This looks at which dims
449/// are present in the map results, and returns an iterator types array with
450/// parallel types for dims that are present, and reduction types for dims that
451/// are not present.
452static FailureOr<SmallVector<utils::IteratorType>>
454 if (!map.isProjectedPermutation())
455 return failure();
456 SmallVector<utils::IteratorType> iterators(map.getNumDims(), red);
457 for (auto expr : map.getResults())
458 if (auto dim = dyn_cast<AffineDimExpr>(expr))
459 iterators[dim.getPosition()] = par;
460 return iterators;
461}
462
463/// Find 2 parallel (m and n) and 1 reduction (k) dimension candidates that form
464/// a matmul subcomputation within `linalgOp`. These dimensions are such that:
465/// 1. The m dimension is involved in an outer-product along LHS
466/// (i.e. it is a permutation on RES and LHS and does not appear in RHS).
467/// 2. The n dimension is involved in an outer-product along RHS
468/// (i.e. it is a permutation on RES and RHS and does not appear in LHS).
469/// 3. The k dimension appears as a permutation on LHS and RHS.
470/// 4. m, n and k appear only once in any given indexing.
471/// 5. Optional batch dimensions that appear in all operands are captured.
472/// This allows e.g. detecting that some contraction is embedded within
473/// `linalgOp` with some orthogonal heuristic.
474static FailureOr<ContractionDimensions>
477 llvm::SmallDenseSet<int64_t> a =
478 findPermutationsIndexingOperand(indexingMaps[0], iterators, par);
479 llvm::SmallDenseSet<int64_t> b =
480 findPermutationsIndexingOperand(indexingMaps[1], iterators, par);
481 llvm::SmallDenseSet<int64_t> c =
482 findPermutationsIndexingOperand(indexingMaps[2], iterators, par);
483
484 // A & C - B are the iterators involved in an outer-product along A (the LHS).
485 llvm::SmallDenseSet<int64_t> ac = a;
486 llvm::set_intersect(ac, c);
487 llvm::set_subtract(ac, b);
488 // B & C - A are the iterators involved in an outer-product along B (the RHS).
489 llvm::SmallDenseSet<int64_t> bc = b;
490 llvm::set_intersect(bc, c);
491 llvm::set_subtract(bc, a);
492 // A & B & C are the "batch" dimensions.
493 llvm::SmallDenseSet<int64_t> batches = a;
494 llvm::set_intersect(batches, b);
495 llvm::set_intersect(batches, c);
496
497 // A & B red are the reduction dimensions.
498 llvm::SmallDenseSet<int64_t> ra =
499 findPermutationsIndexingOperand(indexingMaps[0], iterators, red);
500 llvm::SmallDenseSet<int64_t> rb =
501 findPermutationsIndexingOperand(indexingMaps[1], iterators, red);
502 llvm::set_intersect(ra, rb);
503
504 // Return each set in sorted order.
505 ContractionDimensions dimensions{
506 SmallVector<unsigned, 2>(batches.begin(), batches.end()),
507 SmallVector<unsigned, 2>(ac.begin(), ac.end()),
508 SmallVector<unsigned, 2>(bc.begin(), bc.end()),
509 SmallVector<unsigned, 2>(ra.begin(), ra.end())};
510 llvm::sort(dimensions.batch);
511 llvm::sort(dimensions.m);
512 llvm::sort(dimensions.n);
513 llvm::sort(dimensions.k);
514 return dimensions;
515}
516
517FailureOr<ContractionDimensions>
519 if (linalgOp.getNumDpsInits() != 1 || linalgOp.getNumDpsInputs() != 2)
520 return failure();
521 return inferContractionDimsImpl(linalgOp.getIndexingMapsArray(),
522 linalgOp.getIteratorTypesArray());
523}
524
525FailureOr<ContractionDimensions>
527 if (indexingMaps.size() != 3)
528 return failure();
529 auto iterators = inferIteratorsFromOutMap(indexingMaps[2]);
530 if (failed(iterators))
531 return failure();
532 return inferContractionDimsImpl(indexingMaps, iterators.value());
533}
534
535namespace mlir::linalg::detail {
544} // namespace mlir::linalg::detail
545
549 auto linalgOp = dyn_cast<linalg::LinalgOp>(op);
550 if (!linalgOp)
552 if (linalgOp.getNumDpsInputs() != 2 || linalgOp.getNumDpsInits() != 1)
554 auto mapRange = linalgOp.getIndexingMapsArray();
555 if (linalgOp.getNumReductionLoops() == 0)
557 if (llvm::any_of(mapRange,
558 [](AffineMap m) { return !m.isProjectedPermutation(); }))
560 // TODO: more fields than add/mul.
561 // clang-format off
563 arith::MulFOp, arith::AddFOp,
564 arith::MulIOp, arith::AddIOp,
565 complex::MulOp, complex::AddOp,
566 arith::AndIOp, arith::OrIOp>(
567 *linalgOp.getBlock())) {
569 }
570 // clang-format on
571
572 if (dimensions) {
573 FailureOr<ContractionDimensions> res = inferContractionDims(linalgOp);
574 assert(succeeded(res) && "unexpected failure to infer contraction dims");
575 *dimensions = *res;
576 }
578}
579
580StringRef
582 switch (res) {
584 return "expected a LinalgOp";
586 return "expected op with 2 inputs and 1 output";
588 return "expected at least 1 reduction";
590 return "expected indexing maps to be projected permutations";
592 return "expected add/mul op in the body";
594 return "";
595 }
596 llvm_unreachable("unhandled MatchContractionResult case");
597}
598
600 if (!linalgOp)
601 return false;
602 Operation *op = linalgOp.getOperation();
603 return isa<ContractionOpInterface>(op) ||
606}
607
608/// Verify that a LinalgOp `op` is a contraction.
609/// A Linalg contraction is defined in general terms:
610/// 1. Has 2 input and 1 output shapes.
611/// 2. Has at least one reduction dimension.
612/// 3. Has only projected permutation indexing maps.
613/// 4. its body computes `u5(u1(c) + u2(u3(a) * u4(b)))` on some field
614/// (AddOpType, MulOpType), where u1, u2, u3, u4 and u5 represent scalar unary
615/// operations that may change the type (e.g. for mixed-precision).
616/// As a consequence, when vectorization of such an op occurs, the only special
617/// behavior is that the (unique) MulOpType is vectorized into a
618/// `vector.contract`. All other ops are handled in a generic fashion.
619/// In the future, we may wish to allow more input arguments and elementwise and
620/// constant operations that do not involve the reduction dimension(s).
627
628//===----------------------------------------------------------------------===//
629// ConvolutionOpInterface implementation
630//===----------------------------------------------------------------------===//
631
632/// Of the given two expressions returns one that is of type T (`lhs` gets
633/// preference over `rhs`)
634template <typename T>
636 return isa<T>(lhs) ? cast<T>(lhs) : (isa<T>(rhs) ? cast<T>(rhs) : nullptr);
637}
638
639namespace {
640/// Walk the indexing expressions for input of a convolution operation to verify
641/// its of the right form, either
642/// - AffineDimExpr
643/// - AffineDimExpr (`*` (AffineSymbolExpr | AffineConstantExpr))?
644/// (`+` AffineDimExpr (`*` (AffineSymbolExpr | AffineConstantExpr))?)*
645///
646/// classifies the AffineDimExpr as convolved dimensions or unconvolved
647/// dimensions and verifies each dimension occurs only once.
648struct ConvAccessExprWalker
649 : public AffineExprVisitor<ConvAccessExprWalker, LogicalResult> {
650 // Stores dimensions used in expressions of the above form.
651 llvm::SmallDenseSet<int64_t> convolvedDims;
652 // Stores the dual mapping between LHS and RHS of convolution exprs.
653 llvm::SmallDenseMap<int64_t, int64_t> convolvedDimMapping;
654 // Stores single use dimensions used by an AffineDimExpr.
655 llvm::SmallDenseSet<int64_t> unConvolvedDims;
656 // Stores a mapping from convolved dims to their coefficient.
657 llvm::SmallDenseMap<int64_t, AffineExpr> strideAndDilationMapping;
658
659 // Removes dims with multiple uses in the source input map from dimension
660 // sets tracked by this walker.
661 void clearMultiUseDims(AffineMap map) {
662 for (int dimPos = 0, e = map.getNumDims(); dimPos < e; ++dimPos) {
663 if (llvm::count_if(map.getResults(), [dimPos](AffineExpr e) {
664 return e.isFunctionOfDim(dimPos);
665 }) > 1) {
666 convolvedDims.erase(dimPos);
667 unConvolvedDims.erase(dimPos);
668 // If a duplicate dim is marked as convolved, the pair of the duplicate
669 // dim must be removed from the map as well.
670 auto it = convolvedDimMapping.find(dimPos);
671 if (it != convolvedDimMapping.end()) {
672 int64_t pairedDim = it->second;
673 convolvedDims.erase(pairedDim);
674 unConvolvedDims.erase(pairedDim);
675 strideAndDilationMapping.erase(pairedDim);
676 convolvedDimMapping.erase(dimPos);
677 convolvedDimMapping.erase(pairedDim);
678 }
679 }
680 }
681 }
682
683 LogicalResult visitDimExpr(AffineDimExpr dimExpr) {
684 unsigned position = dimExpr.getPosition();
685 if (unConvolvedDims.count(position) || convolvedDims.count(position)) {
686 return failure();
687 }
688 unConvolvedDims.insert(position);
689 return success();
690 }
691
692 LogicalResult visitSymbolExpr(AffineSymbolExpr expr) { return failure(); }
693
694 LogicalResult visitConstantExpr(AffineConstantExpr expr) { return failure(); }
695
696 LogicalResult visitAffineBinaryOpExpr(AffineBinaryOpExpr binaryExpr) {
697 // In pre-order visit, top level op has to be an add op.
698 if (binaryExpr.getKind() != AffineExprKind::Add)
699 return failure();
700 auto lhsDimPos = getDimExprOrMulExprDimPos(binaryExpr.getLHS());
701 auto rhsDimPos = getDimExprOrMulExprDimPos(binaryExpr.getRHS());
702 if (failed(lhsDimPos) || failed(rhsDimPos))
703 return failure();
704 convolvedDimMapping[*lhsDimPos] = *rhsDimPos;
705 convolvedDimMapping[*rhsDimPos] = *lhsDimPos;
706 return success();
707 }
708
709 FailureOr<int64_t> getDimExprOrMulExprDimPos(AffineExpr expr) {
710 if (auto dimExpr = dyn_cast<AffineDimExpr>(expr)) {
711 int64_t dim = dimExpr.getPosition();
712 if (convolvedDims.count(dim) || unConvolvedDims.count(dim))
713 return failure();
714 // Stride/dilation for this dim is implicitly 1.
715 strideAndDilationMapping[dim] =
717 convolvedDims.insert(dim);
718 return dim;
719 }
720 if (auto symbolMulExpr = dyn_cast<AffineBinaryOpExpr>(expr)) {
721 if (symbolMulExpr.getKind() != AffineExprKind::Mul)
722 return failure();
723 auto lhsExpr = symbolMulExpr.getLHS();
724 auto rhsExpr = symbolMulExpr.getRHS();
725 // Check for symbol expression.
726 AffineExpr mulExpr =
728 // If there was no symbol expr, check for constant expression.
729 if (!mulExpr) {
730 mulExpr = getAffineExprOfType<AffineConstantExpr>(lhsExpr, rhsExpr);
731 }
732 auto dimExpr = getAffineExprOfType<AffineDimExpr>(lhsExpr, rhsExpr);
733 if (!mulExpr || !dimExpr)
734 return failure();
735 int64_t dim = dimExpr.getPosition();
736 if (convolvedDims.count(dim) || unConvolvedDims.count(dim))
737 return failure();
738 strideAndDilationMapping[dim] = mulExpr;
739 convolvedDims.insert(dim);
740 return dim;
741 }
742 return failure();
743 }
744};
745} // namespace
746
747static llvm::SmallDenseSet<int64_t> getPreservedDims(AffineMap map) {
748 assert(map.isProjectedPermutation() &&
749 "expected map to have projected permutations");
750 llvm::SmallDenseSet<int64_t> preservedDims;
751 for (auto expr : map.getResults())
752 preservedDims.insert(cast<AffineDimExpr>(expr).getPosition());
753 return preservedDims;
754}
755
759 for (auto e : exprs) {
760 auto constantExpr = dyn_cast<AffineConstantExpr>(e);
761 assert(constantExpr && "Found non-constant stride/dilation");
762 vals.push_back(constantExpr.getValue());
763 }
764 return vals;
765}
766
767/// Classifies dimensions in the `indexingMaps` used by a convolution
768/// subcomputation, as captured by `inputExprWalker`. If
769/// `allowEmptyConvolvedDims` is not set this will fail if there is not
770/// at least one convolved dimension pair (output image + filter loop).
771///
772/// The returned dimensions are ordered as follows:
773/// - `outputImage` is sorted by dimension index.
774/// - `filterLoop` is ordered to match the pairing with `outputImage`, i.e.,
775/// `outputImage[i]` and `filterLoop[i]` are paired dimensions from the
776/// convolution access pattern (e.g., `oh + kh` pairs `oh` with `kh`).
777/// - `strides[i]` corresponds to `outputImage[i]`.
778/// - `dilations[i]` corresponds to `filterLoop[i]`.
779/// - Other dimension sets (batch, outputChannel, etc.) are sorted by index.
780///
781/// `nativeStrides` and `nativeDilations`, when non-null, are the op-carried
782/// `strides`/`dilations` attributes and take precedence over the values derived
783/// from the convolution access pattern. They are null for the maps-based
784/// overload.
785static FailureOr<ConvolutionDimensions> inferConvolutionDimsImpl(
787 ConvAccessExprWalker &inputExprWalker, bool allowEmptyConvolvedDims,
788 DenseIntElementsAttr nativeStrides, DenseIntElementsAttr nativeDilations) {
789 AffineMap filterMap = indexingMaps[1];
790 AffineMap outputMap = indexingMaps.back();
791 llvm::SmallDenseSet<int64_t> filterDims =
792 findPermutationsIndexingOperand(filterMap, iterators, par);
793 llvm::SmallDenseSet<int64_t> outputDims =
794 findPermutationsIndexingOperand(outputMap, iterators, par);
795
796 // unConvolvedDims & outputDims - filterDims are the batch iterators.
797 llvm::SmallDenseSet<int64_t> batch = inputExprWalker.unConvolvedDims;
798 llvm::set_intersect(batch, outputDims);
799 llvm::set_subtract(batch, filterDims);
800
801 // convolvedDims & outputDims are the output image iterators.
802 llvm::SmallDenseSet<int64_t> oi = inputExprWalker.convolvedDims;
803 llvm::set_intersect(oi, outputDims);
804
805 // filterDims & outputDims - unConvolvedDims are the output channel iterators.
806 llvm::SmallDenseSet<int64_t> oc = filterDims;
807 llvm::set_intersect(oc, outputDims);
808 llvm::set_subtract(oc, inputExprWalker.unConvolvedDims);
809
810 // filterDims & outputDims & unConvolvedDims are the depth iterators.
811 llvm::SmallDenseSet<int64_t> depth = filterDims;
812 llvm::set_intersect(depth, outputDims);
813 llvm::set_intersect(depth, inputExprWalker.unConvolvedDims);
814
815 llvm::SmallDenseSet<int64_t> filterReducedDims =
816 findPermutationsIndexingOperand(filterMap, iterators, red);
817
818 // convolvedDims & filterReducedDims are the filter loop iterators.
819 llvm::SmallDenseSet<int64_t> fl = inputExprWalker.convolvedDims;
820 llvm::set_intersect(fl, filterReducedDims);
821
822 // unConvolvedDims & filterReducedDims are the input channel iterators.
823 llvm::SmallDenseSet<int64_t> ic = inputExprWalker.unConvolvedDims;
824 llvm::set_intersect(ic, filterReducedDims);
825
826 if (oi.empty() && !allowEmptyConvolvedDims)
827 return failure();
828
829 // Return each set in sorted order, with outputImage and filterLoop
830 // ordered so that outputImage[i] pairs with filterLoop[i].
831 ConvolutionDimensions dimensions{
832 SmallVector<unsigned, 2>(batch.begin(), batch.end()),
833 SmallVector<unsigned, 2>(oi.begin(), oi.end()),
834 SmallVector<unsigned, 2>(oc.begin(), oc.end()),
835 /*filterLoop=*/SmallVector<unsigned, 2>{},
836 SmallVector<unsigned, 2>(ic.begin(), ic.end()),
837 SmallVector<unsigned, 2>(depth.begin(), depth.end()),
838 /*strides=*/SmallVector<int64_t, 2>{},
839 /*dilations=*/SmallVector<int64_t, 2>{}};
840 llvm::sort(dimensions.batch);
841 llvm::sort(dimensions.outputImage);
842 llvm::sort(dimensions.outputChannel);
843 llvm::sort(dimensions.inputChannel);
844 llvm::sort(dimensions.depth);
845 // Order filterLoop to match the pairing with outputImage. Each outputImage
846 // dimension has a corresponding filterLoop dimension from the convolution
847 // access pattern (e.g., oh + kh). This ensures outputImage[i] pairs with
848 // filterLoop[i].
849 for (unsigned oiDim : dimensions.outputImage)
850 dimensions.filterLoop.push_back(inputExprWalker.convolvedDimMapping[oiDim]);
851
852 // Use the op carried strides/dilations attribute if present.
853 if (!nativeStrides) {
854 SmallVector<AffineExpr, 2> strideExprs;
855 for (unsigned oiDim : dimensions.outputImage)
856 strideExprs.push_back(inputExprWalker.strideAndDilationMapping[oiDim]);
857 dimensions.strides = getConstantsFromExprList(strideExprs);
858 } else {
859 dimensions.strides = llvm::to_vector<2>(nativeStrides.getValues<int64_t>());
860 }
861 if (!nativeDilations) {
862 SmallVector<AffineExpr, 2> dilationExprs;
863 for (unsigned flDim : dimensions.filterLoop)
864 dilationExprs.push_back(inputExprWalker.strideAndDilationMapping[flDim]);
865 dimensions.dilations = getConstantsFromExprList(dilationExprs);
866 } else {
867 dimensions.dilations =
868 llvm::to_vector<2>(nativeDilations.getValues<int64_t>());
869 }
870 return dimensions;
871}
872
874 StringRef name) {
875 return dyn_cast_or_null<DenseIntElementsAttr>(
876 linalgOp->getInherentAttr(name).value_or(Attribute{}));
877}
878
879/// Find at least 1 parallel (output_image) and reduction (filter_loop)
880/// dimension candidates that form a convolution subcomputation within
881/// `linalgOp`. The LHS is assumed to be the convolution input while the
882/// RHS is assumed as the filter.
883/// These dimensions are such that:
884/// 1. Optional batch dimensions that appear in the input and filter.
885/// 2. The output_image dimension is involved in a cross-correlation along LHS
886/// (i.e. it is a permutation on RES and LHS and has an associated
887/// filter_loop in RHS).
888/// 3. Optional output_channel dimension is involved in an outer-product along
889/// RHS (i.e. it is a permutation on RES and RHS and does not appear in
890/// LHS).
891/// 4. Optional input_channel dimension appears as a permutation on LHS and
892/// RHS.
893/// 5. The filter_loop dimension appears as a permutation on the RHS and
894/// represents the shape of the kernel cross-correlated along a
895/// corresponding output_image dim.
896/// 6. The input_channel dimension appears as a permutation on LHS and RHS.
897/// 7. All dimensions appear only once in any given indexing map.
898/// This allows e.g. detecting that some convolution is embedded within
899/// `linalgOp` with some orthogonal heuristic.
900///
901/// The `outputImage` and `filterLoop` arrays are ordered such that
902/// `outputImage[i]` pairs with `filterLoop[i]` based on the convolution access
903/// pattern in the input indexing map (e.g., `d0 + d2` pairs dimension 0 with
904/// dimension 2). Other dimension sets are returned in sorted order.
905///
906/// Returns a failure if `output_image` (and implicitly `filter_loop`) is empty.
907FailureOr<ConvolutionDimensions>
909 if (linalgOp.getNumDpsInits() != 1 || linalgOp.getNumDpsInputs() != 2)
910 return failure();
911
912 auto indexingMaps = linalgOp.getIndexingMapsArray();
913
914 // Check the input indexing map has the right form.
915 ConvAccessExprWalker inputExprWalker;
916 for (AffineExpr expr : indexingMaps[0].getResults())
917 (void)inputExprWalker.visit(expr);
918 inputExprWalker.clearMultiUseDims(indexingMaps[0]);
919
921 indexingMaps, linalgOp.getIteratorTypesArray(), inputExprWalker,
922 /*allowEmptyConvolvedDims=*/false,
923 getInherentConvolutionAttr(linalgOp, "strides"),
924 getInherentConvolutionAttr(linalgOp, "dilations"));
925}
926
927FailureOr<ConvolutionDimensions>
929 if (indexingMaps.size() != 3)
930 return failure();
931
932 // Infer iterator types from the output map.
933 FailureOr<SmallVector<utils::IteratorType>> iterators =
934 inferIteratorsFromOutMap(indexingMaps[2]);
935 if (failed(iterators))
936 return failure();
937
938 // Check the input indexing map has the right form.
939 ConvAccessExprWalker inputExprWalker;
940 for (AffineExpr expr : indexingMaps[0].getResults())
941 (void)inputExprWalker.visit(expr);
942 inputExprWalker.clearMultiUseDims(indexingMaps[0]);
943
944 return inferConvolutionDimsImpl(indexingMaps, iterators.value(),
945 inputExprWalker,
946 /*allowEmptyConvolvedDims=*/false,
947 /*nativeStrides=*/nullptr,
948 /*nativeDilations=*/nullptr);
949}
950
951namespace mlir::linalg::detail {
963} // namespace mlir::linalg::detail
964
967 Operation *op, ConvolutionDimensions *dimensions,
968 bool allowEmptyConvolvedDims) {
969 auto linalgOp = dyn_cast<linalg::LinalgOp>(op);
970 if (!linalgOp)
972 if (linalgOp.getNumDpsInputs() < 2 || linalgOp.getNumDpsInits() != 1)
974
975 auto indexingMaps = linalgOp.getIndexingMapsArray();
976
977 // Check the input indexing map has the right form.
978 ConvAccessExprWalker inputExprWalker;
979 if (llvm::any_of(indexingMaps[0].getResults(),
980 [&inputExprWalker](AffineExpr expr) {
981 return failed(inputExprWalker.visit(expr));
982 })) {
984 }
985
986 // Filter and output maps must be projected permutation.
987 if (!indexingMaps[1].isProjectedPermutation() ||
988 !indexingMaps.back().isProjectedPermutation())
990
991 auto iteratorTypes = linalgOp.getIteratorTypesArray();
992
993 llvm::SmallDenseSet<int64_t> outputDims =
994 getPreservedDims(indexingMaps.back());
995 llvm::SmallDenseSet<int64_t> filterDims = getPreservedDims(indexingMaps[1]);
996 // Make sure all loops are characterized as one of:
997 // - Batch loop : present in output, as non-convolved in input, not present in
998 // filter.
999 // - Output image dimension : present in output, convolved dims in input, not
1000 // present in filter.
1001 // - Output channel dimension : present in output, not present in input,
1002 // present in filter.
1003 // - Filter loop dimension : present in filter, convolved in input, not
1004 // present in output.
1005 // - Input channel dimension : unconvolved in input, not present in output,
1006 // present in filter.
1007 // - Depth multiplier : unconvolved in input, present in output, present in
1008 // filter.
1009 llvm::SmallDenseSet<int64_t> allLoopDims;
1010 for (auto outputExpr : indexingMaps.back().getResults()) {
1011 int64_t outputDim = cast<AffineDimExpr>(outputExpr).getPosition();
1012 if (inputExprWalker.unConvolvedDims.count(outputDim) &&
1013 !filterDims.count(outputDim)) {
1014 // Batch dimension.
1015 if (iteratorTypes[outputDim] != utils::IteratorType::parallel)
1017 allLoopDims.insert(outputDim);
1018 continue;
1019 }
1020 if (inputExprWalker.convolvedDims.count(outputDim) &&
1021 !filterDims.count(outputDim)) {
1022 // Output image Loop dimension.
1023 if (iteratorTypes[outputDim] != utils::IteratorType::parallel)
1025 allLoopDims.insert(outputDim);
1026 continue;
1027 }
1028 if (!inputExprWalker.convolvedDims.count(outputDim) &&
1029 !inputExprWalker.unConvolvedDims.count(outputDim) &&
1030 filterDims.count(outputDim)) {
1031 // Output channel dimension.
1032 if (iteratorTypes[outputDim] != utils::IteratorType::parallel)
1034 allLoopDims.insert(outputDim);
1035 continue;
1036 }
1037 if (inputExprWalker.unConvolvedDims.count(outputDim) &&
1038 filterDims.count(outputDim)) {
1039 // Depth multiplier.
1040 if (iteratorTypes[outputDim] != utils::IteratorType::parallel)
1042 allLoopDims.insert(outputDim);
1043 continue;
1044 }
1046 }
1047 for (auto filterExpr : indexingMaps[1].getResults()) {
1048 int64_t filterDim = cast<AffineDimExpr>(filterExpr).getPosition();
1049 if (outputDims.count(filterDim) &&
1050 !inputExprWalker.unConvolvedDims.count(filterDim) &&
1051 !inputExprWalker.convolvedDims.count(filterDim)) {
1052 // Output channel dimension. This is already seen, continue;
1053 continue;
1054 }
1055 if (inputExprWalker.convolvedDims.count(filterDim) &&
1056 !outputDims.count(filterDim)) {
1057 // Filter loop dimension.
1058 if (iteratorTypes[filterDim] != utils::IteratorType::reduction)
1060 if (allLoopDims.count(filterDim))
1062 allLoopDims.insert(filterDim);
1063 continue;
1064 }
1065 if (inputExprWalker.unConvolvedDims.count(filterDim) &&
1066 !outputDims.count(filterDim)) {
1067 // Input channel dimension.
1068 if (iteratorTypes[filterDim] != utils::IteratorType::reduction)
1070 if (allLoopDims.count(filterDim))
1072 allLoopDims.insert(filterDim);
1073 continue;
1074 }
1075 if (inputExprWalker.unConvolvedDims.count(filterDim) &&
1076 outputDims.count(filterDim)) {
1077 // Depthwise loop. Already seen.
1078 continue;
1079 }
1081 }
1082 // All loops must be covered now.
1083 if (allLoopDims.size() != linalgOp.getNumLoops())
1085
1086 if (!allowEmptyConvolvedDims && inputExprWalker.convolvedDims.empty())
1088
1089 if (dimensions) {
1090 FailureOr<ConvolutionDimensions> res = inferConvolutionDimsImpl(
1091 indexingMaps, iteratorTypes, inputExprWalker, allowEmptyConvolvedDims,
1092 getInherentConvolutionAttr(linalgOp, "strides"),
1093 getInherentConvolutionAttr(linalgOp, "dilations"));
1094 assert(succeeded(res) && "unexpected failure to infer convolution dims");
1095 *dimensions = *res;
1096 }
1097
1099}
1100
1101StringRef
1103 switch (res) {
1105 return "expected a LinalgOp";
1107 return "expected op with 2 inputs and 1 output";
1109 return "unexpected input index map for convolutions";
1111 return "expected output/filter indexing maps to be projected permutations";
1113 return "unexpected loop dimension for convolution op";
1115 return "expected all iterators used to access outputs to be parallel";
1117 return "expected all iterators not used to access outputs to be reduction";
1119 return "expected convolved dim to be non-empty";
1121 return "";
1122 }
1123 llvm_unreachable("unhandled MatchConvolutionResult case");
1124}
1125
1127 bool allowEmptyConvolvedDims) {
1129 linalgOp.getOperation(), nullptr, allowEmptyConvolvedDims) ==
1131}
1132
1139
1140//===----------------------------------------------------------------------===//
1141// FillOpInterface implementation
1142//===----------------------------------------------------------------------===//
1143
1144namespace {
1145enum class MatchFillResult {
1146 Success = 0,
1147 NotLinalgOp,
1148 WrongNumOperands,
1149 NotScalarInput,
1150 TypeMismatch
1151};
1152} // namespace
1153
1154static MatchFillResult isFillInterfaceImpl(Operation *op) {
1155 auto linalgOp = dyn_cast<linalg::LinalgOp>(op);
1156 if (!linalgOp)
1157 return MatchFillResult::NotLinalgOp;
1158 if (linalgOp.getNumDpsInputs() != 1 || linalgOp.getNumDpsInits() != 1)
1159 return MatchFillResult::WrongNumOperands;
1160
1161 OpOperand *value = linalgOp.getDpsInputOperand(0);
1162 if (!linalgOp.isScalar(value))
1163 return MatchFillResult::NotScalarInput;
1164
1165 // Check that the scalar input type matches the output element type.
1166 OpOperand *output = linalgOp.getDpsInitOperand(0);
1167 Type scalarType = value->get().getType();
1168 Type outputElementType = getElementTypeOrSelf(output->get().getType());
1169 if (scalarType != outputElementType)
1170 return MatchFillResult::TypeMismatch;
1171
1172 return MatchFillResult::Success;
1173}
1174
1176 MatchFillResult res = isFillInterfaceImpl(op);
1177 if (res == MatchFillResult::NotLinalgOp)
1178 return op->emitError("expected a LinalgOp");
1179 if (res == MatchFillResult::WrongNumOperands)
1180 return op->emitError("expected op with 1 input and 1 output");
1181 if (res == MatchFillResult::NotScalarInput)
1182 return op->emitError("expected op with scalar input");
1183 if (res == MatchFillResult::TypeMismatch) {
1184 auto linalgOp = cast<linalg::LinalgOp>(op);
1185 Type scalarType = linalgOp.getDpsInputOperand(0)->get().getType();
1186 Type outputElementType =
1187 getElementTypeOrSelf(linalgOp.getDpsInitOperand(0)->get().getType());
1188 return op->emitOpError("expected fill value type (")
1189 << scalarType << ") to match output element type ("
1190 << outputElementType << ")";
1191 }
1192
1193 return success();
1194}
1195
1196//===----------------------------------------------------------------------===//
1197// StructuredOpInterface implementation
1198//===----------------------------------------------------------------------===//
1199
1200SmallVector<OpFoldResult> LinalgOp::createFlatListOfOperandDims(OpBuilder &b,
1201 Location loc) {
1203 for (OpOperand &opOperand : getOperation()->getOpOperands()) {
1204 for (int64_t i = 0, e = getRank(&opOperand); i < e; ++i)
1205 res.push_back(createFoldedDimOp(b, loc, opOperand.get(), i));
1206 }
1207 return res;
1208}
1209
1210SmallVector<int64_t, 4> LinalgOp::createFlatListOfOperandStaticDims() {
1212 assert(!hasDynamicShape() && "expected operands to have static shapes");
1213 for (OpOperand &opOperand : getOperation()->getOpOperands())
1214 llvm::append_range(res, getShape(&opOperand));
1215 return res;
1216}
1217
1218SmallVector<Range, 4> LinalgOp::createLoopRanges(OpBuilder &b, Location loc) {
1219 AffineMap map = getLoopsToShapesMap();
1220 unsigned numDims = map.getNumDims(), numRes = map.getNumResults();
1221 auto viewSizes = createFlatListOfOperandDims(b, loc);
1222 SmallVector<Range, 4> res(numDims);
1223 for (unsigned idx = 0; idx < numRes; ++idx) {
1224 auto result = map.getResult(idx);
1225 if (auto d = dyn_cast<AffineDimExpr>(result)) {
1226 if (res[d.getPosition()].offset)
1227 continue;
1228 res[d.getPosition()] =
1229 Range{b.getIndexAttr(0), viewSizes[idx], b.getIndexAttr(1)};
1230 }
1231 }
1232 return res;
1233}
1234
1235/// Visitor to check if any of the given set of positions from AffineDimExprs
1236/// are used within an AffineExpr.
1238 : public AffineExprVisitor<HasAffineDimExprVisitor, bool> {
1239 HasAffineDimExprVisitor(llvm::SmallBitVector positions)
1240 : positions(std::move(positions)) {}
1241
1243 return visit(binaryOpExpr.getLHS()) || visit(binaryOpExpr.getRHS());
1244 }
1245
1247 return positions.test(dimExpr.getPosition());
1248 }
1249
1250 bool visitConstantExpr(AffineConstantExpr constExpr) { return false; }
1251
1252 bool visitSymbolExpr(AffineSymbolExpr symbolExpr) { return false; }
1253
1254private:
1255 llvm::SmallBitVector positions;
1256};
1257
1258static std::pair<int64_t, int64_t>
1260 int64_t inputRankSum = 0;
1261 int64_t outputRankSum = 0;
1262 for (OpOperand *input : op.getDpsInputOperands())
1263 inputRankSum += op.getRank(input);
1264 for (OpOperand &output : op.getDpsInitsMutable())
1265 outputRankSum += op.getRank(&output);
1266 return {inputRankSum, inputRankSum + outputRankSum};
1267}
1268
1269LogicalResult
1270LinalgOp::reifyResultShapes(OpBuilder &b,
1271 ReifiedRankedShapedTypeDims &reifiedReturnShapes) {
1272 // An example that helps understand the logic below.
1273 // Consider the following expression O(i+j, j) += A(i,k) * B(k, j)
1274 // We want to express the shape of dim 0 of O in terms of shape of the inputs.
1275 // This is achieved as follows.
1276 // loopsToShapesMap = (d0, d1, d2) -> (d0, d2, d2, d1, d0 + d1, d1)
1277 // subMapOfResultShapes = (d0, d1, d2) -> (d0 + d1, d1)
1278 // shapesToLoopsMap = (d0, d2, d2, d3, d4, d5) -> (d0, d3, d2)
1279 // resultShapesFromInputShapes = subMapOfResultDim.compose(shapesToLoopMap)
1280 // = (d0, d1, d2, d3, d4, d5) -> (d0 + d1, d1)
1281 AffineMap loopsToShapesMap = getLoopsToShapesMap();
1282
1283 // Find the position in the above map that represents the shape of the
1284 // result:dim being inferred.
1285 auto resultShapesSubMapPos = getResultsPositionInLoopsToShapeMap(*this);
1286
1287 /// From loopsToShapesMap extract the submap that represents the shape of the
1288 /// (resultIdx, dim) needed.
1289 AffineMap loopToResultsShapeMap = loopsToShapesMap.getSliceMap(
1290 resultShapesSubMapPos.first,
1291 resultShapesSubMapPos.second - resultShapesSubMapPos.first);
1292 AffineMap resultShapesFromInputShapesMap =
1293 loopToResultsShapeMap.compose(getShapesToLoopsMap());
1294
1295 // Check that the result dim map does not contain the positions corresponding
1296 // to the outputs.
1297 llvm::SmallBitVector outputDims(resultShapesFromInputShapesMap.getNumDims());
1298 outputDims.set(resultShapesSubMapPos.first, resultShapesSubMapPos.second);
1299 HasAffineDimExprVisitor checkDimExpr(std::move(outputDims));
1300 Location loc = getOperation()->getLoc();
1301 IRRewriter rewriter(b);
1302 SmallVector<OpFoldResult> allResultDimValues =
1303 affine::makeComposedFoldedMultiResultAffineApply(
1304 rewriter, loc, resultShapesFromInputShapesMap,
1305 createFlatListOfOperandDims(b, loc));
1306 int64_t pos = 0;
1307 ArrayRef<AffineExpr> shapeExprs = resultShapesFromInputShapesMap.getResults();
1308 for (OpOperand &opOperand : getDpsInitsMutable()) {
1309 SmallVector<OpFoldResult> shapes;
1310 for (int64_t dim : llvm::seq<int64_t>(0, getRank(&opOperand))) {
1311 auto shapedType = llvm::cast<ShapedType>(opOperand.get().getType());
1312 if (!shapedType.isDynamicDim(dim)) {
1313 // Static dim: Return IntegerAttr.
1314 shapes.push_back(b.getIndexAttr(shapedType.getDimSize(dim)));
1315 } else {
1316 // Dynamic dim: Return Value.
1317 OpFoldResult ofr = checkDimExpr.visit(shapeExprs[pos])
1318 ? createOrFoldDimOp(b, loc, opOperand.get(), dim)
1319 : allResultDimValues[pos];
1320 shapes.push_back(getValueOrCreateConstantIndexOp(b, loc, ofr));
1321 }
1322 pos++;
1323 }
1324 reifiedReturnShapes.emplace_back(std::move(shapes));
1325 }
1326 return success();
1327}
1328
1329/// Return the index in the indexingMaps vector that corresponds to this
1330/// `opOperand`.
1331int64_t LinalgOp::getIndexingMapIndex(OpOperand *opOperand) {
1332 auto operandNumber = opOperand->getOperandNumber();
1333 auto dpsIface = cast<DestinationStyleOpInterface>(*this->getOperation());
1334 if (!dpsIface.isDpsInput(opOperand))
1335 return operandNumber;
1336 unsigned start = dpsIface.getDpsInits().getBeginOperandIndex();
1337 assert(!dpsIface.isDpsInit(opOperand));
1338 // Account for potential inputs that are not DPS and may not appear in
1339 // `indexingMaps`.
1340 return cast<DestinationStyleOpInterface>(*this->getOperation())
1341 .getNumDpsInputs() +
1342 operandNumber - start;
1343}
1344
1346 LinalgOp linalgOp = cast<LinalgOp>(op);
1347 // Mixed tensor/buffer operands are not allowed.
1348 if (!linalgOp.hasPureTensorSemantics() &&
1349 !linalgOp.hasPureBufferSemantics() && op->getNumOperands() > 0)
1350 return op->emitOpError("expected to have pure tensor or buffer semantics");
1351
1352 // Before checking indexing maps, we need to make sure the attributes
1353 // referenced by it are valid.
1354 if (linalgOp.hasDynamicIndexingMaps())
1355 if (failed(linalgOp.verifyIndexingMapRequiredAttributes()))
1356 return failure();
1357
1358 // Delayed calling of IndexingMapOpInterface::verifyImpl.
1359 if (failed(cast<IndexingMapOpInterface>(op).verifyImpl()))
1360 return failure();
1361
1362 // Set this flag if this op has user defined maps. This is required to guard
1363 // the below error condition which assume default indexing maps.
1364 for (OpOperand &opOperand : linalgOp->getOpOperands()) {
1365 AffineMap indexingMap = linalgOp.getMatchingIndexingMap(&opOperand);
1366 // Domain must be consistent.
1367 unsigned numLoops = linalgOp.getNumLoops();
1368 if (indexingMap.getNumDims() != numLoops)
1369 return op->emitOpError("expected indexing_map #")
1370 << opOperand.getOperandNumber() << " to have " << numLoops
1371 << " dim(s) to match the number of loops";
1372 }
1373 SmallVector<unsigned> redDims;
1374 linalgOp.getReductionDims(redDims);
1375
1376 if (!linalgOp.getShapesToLoopsMap())
1377 return op->emitOpError("expected the shape-to-loops map to be non-null");
1378
1379 // Check the region has exactly one block.
1380 if (linalgOp->getNumRegions() != 1 || !linalgOp->getRegion(0).hasOneBlock())
1381 return op->emitOpError("expects to have 1 region with 1 block");
1382
1383 // Simplifying assumption: bbargs match 1-1 with shape operands elemental
1384 // types.
1385 // TODO: once ranked shape types are plugged in, we may want to drop the
1386 // corresponding bbargs, that can never be read from. This will be subject to
1387 // consistency discussions (i.e. what to do with output tensors whose bbarg is
1388 // not used).
1389 Block &block = linalgOp->getRegion(0).front();
1390
1391 if (linalgOp.getOpOperandsMatchingBBargs().size() != block.getNumArguments())
1392 return op->emitOpError("expected as many non-induction variable region "
1393 "arguments as the number of input/output operands");
1394
1395 for (OpOperand *opOperand : linalgOp.getOpOperandsMatchingBBargs()) {
1396 Type elementType = opOperand->get().getType();
1397 if (isa<MemRefType, RankedTensorType>(elementType))
1398 elementType = getElementTypeOrSelf(opOperand->get().getType());
1399 Type argType = block.getArgument(opOperand->getOperandNumber()).getType();
1400 if (elementType != argType)
1401 return op->emitOpError("expected type of bb argument #")
1402 << opOperand->getOperandNumber() << " (" << argType << ")"
1403 << " to match element or self type of the corresponding operand ("
1404 << elementType << ")";
1405 }
1406
1407 return success();
1408}
return success()
static FailureOr< ContractionDimensions > inferContractionDimsImpl(ArrayRef< AffineMap > indexingMaps, ArrayRef< utils::IteratorType > iterators)
Find 2 parallel (m and n) and 1 reduction (k) dimension candidates that form a matmul subcomputation ...
static Value getSourceSkipUnary(Value value)
If the value is defined by a chain of unary side effect-free, go up the use-def chain until the first...
static llvm::SmallDenseSet< int64_t > getPreservedDims(AffineMap map)
static T getAffineExprOfType(AffineExpr lhs, AffineExpr rhs)
Of the given two expressions returns one that is of type T (lhs gets preference over rhs)
static FailureOr< ConvolutionDimensions > inferConvolutionDimsImpl(ArrayRef< AffineMap > indexingMaps, ArrayRef< utils::IteratorType > iterators, ConvAccessExprWalker &inputExprWalker, bool allowEmptyConvolvedDims, DenseIntElementsAttr nativeStrides, DenseIntElementsAttr nativeDilations)
Classifies dimensions in the indexingMaps used by a convolution subcomputation, as captured by inputE...
static std::pair< int64_t, int64_t > getResultsPositionInLoopsToShapeMap(LinalgOp &op)
static bool isPairTemplateImpl(Operation *add, Operation *mul)
Returns true if the two operations are of the kinds specified by a pair of consecutive template argum...
static bool isaElemwiseSingleOpInterface(linalg::GenericOp op, unsigned arity)
static MatchFillResult isFillInterfaceImpl(Operation *op)
static bool isContractionBody(Block &block)
Returns true if the block is a body of a contraction with the kinds of operations given pairwise by t...
static std::optional< Value > isaExternalFillOp(GenericOp op)
Detects if a linalg.generic operation represents an external scalar input.
static FailureOr< SmallVector< utils::IteratorType > > inferIteratorsFromOutMap(AffineMap map)
Infer the iterator types from the init affine map.
static DenseIntElementsAttr getInherentConvolutionAttr(LinalgOp linalgOp, StringRef name)
static llvm::SmallDenseSet< int64_t > findPermutationsIndexingOperand(AffineMap indexingMap, ArrayRef< utils::IteratorType > iterators, utils::IteratorType iter)
Given an indexingMap and its corresponding iterators, returns the positions of the iterators of type ...
static SmallVector< int64_t, 2 > getConstantsFromExprList(const SmallVector< AffineExpr, 2 > &exprs)
static std::optional< Value > isaInlinedFillOp(GenericOp op)
Detects if a linalg.generic operation represents a fill with an inlined constant.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Definition Traits.cpp:117
#define mul(a, b)
#define add(a, b)
Affine binary operation expression.
Definition AffineExpr.h:214
AffineExpr getLHS() const
AffineExpr getRHS() const
An integer constant appearing in affine expression.
Definition AffineExpr.h:239
A dimensional identifier appearing in an affine expression.
Definition AffineExpr.h:223
unsigned getPosition() const
See documentation for AffineExprVisitorBase.
Base type for affine expression.
Definition AffineExpr.h:68
AffineExprKind getKind() const
Return the classification for this type.
MLIRContext * getContext() const
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
AffineMap getSliceMap(unsigned start, unsigned length) const
Returns the map consisting of length expressions starting from start.
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
AffineExpr getResult(unsigned idx) const
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
A symbolic identifier appearing in an affine expression.
Definition AffineExpr.h:231
Attributes are known-constant values of operations.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
bool empty()
Definition Block.h:173
BlockArgument getArgument(unsigned i)
Definition Block.h:154
unsigned getNumArguments()
Definition Block.h:153
OpListType & getOperations()
Definition Block.h:162
Operation & front()
Definition Block.h:178
Operation & back()
Definition Block.h:177
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
An attribute that represents a reference to a dense integer vector or tensor object.
IRValueT get() const
Return the current value being used by this operand.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
This class helps build Operations.
Definition Builders.h:210
This class represents an operand of an operation.
Definition Value.h:254
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
Definition Value.cpp:226
This class provides the API for ops that are known to be terminators.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
bool mightHaveTrait()
Returns true if the operation might have the provided trait.
Definition Operation.h:809
unsigned getNumOperands()
Definition Operation.h:371
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
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
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
MatchConvolutionResult isConvolutionInterfaceImpl(Operation *op, ConvolutionDimensions *dimensions=nullptr, bool allowEmptyConvolvedDims=false)
Checks whether op conforms to ConvolutionOpInterface and populates dimensions with indexes of the dif...
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:
StringRef getMatchConvolutionMessage(MatchConvolutionResult res)
Returns the error message corresponding to the convolution checking return code.
bool canOpOperandsBeDroppedImpl(linalg::LinalgOp linalgOp, ArrayRef< OpOperand * > droppedOperands)
Implementation of the method that check if given operands can be dropped, i.e.
MatchContractionResult isContractionInterfaceImpl(Operation *op, ContractionDimensions *dimensions=nullptr)
Checks whether op conforms to ContractionOpInterface and populates dimensions with indexes of the dif...
LogicalResult verifyContractionInterface(Operation *op)
Verify that op conforms to ContractionOpInterface.
LogicalResult verifyFillInterface(Operation *op)
Verify that op conforms to the FillOpInterface.
StringRef getMatchContractionMessage(MatchContractionResult res)
Returns the error message corresponding to the contraction checking return code.
LogicalResult verifyStructuredOpInterface(Operation *op)
Verify that op conforms to the invariants of StructuredOpInterface.
LogicalResult verifyConvolutionInterface(Operation *op)
Verify that op conforms to the ConvolutionOpInterface.
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.
FailureOr< ConvolutionDimensions > inferConvolutionDims(LinalgOp linalgOp)
Find at least 1 parallel (output_image) and reduction (filter_loop) dimension candidates that form a ...
OpFoldResult createFoldedDimOp(OpBuilder &b, Location loc, Value val, int64_t dim)
Create one memref::DimOp or tensor::DimOp depending on the type of val.
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....
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...
Value createOrFoldDimOp(OpBuilder &b, Location loc, Value val, int64_t dim)
Create one memref::DimOp or tensor::DimOp depending on the type of val.
bool isaContractionOpInterface(LinalgOp linalgOp)
Checks whether linalgOp conforms to ContractionOpInterface.
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.
AffineMap concatAffineMaps(ArrayRef< AffineMap > maps, MLIRContext *context)
Concatenates a list of maps into a single AffineMap, stepping over potentially empty maps.
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Definition Utils.cpp:114
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
HasAffineDimExprVisitor(llvm::SmallBitVector positions)
bool visitDimExpr(AffineDimExpr dimExpr)
bool visitAffineBinaryOpExpr(AffineBinaryOpExpr binaryOpExpr)
bool visitSymbolExpr(AffineSymbolExpr symbolExpr)
bool visitConstantExpr(AffineConstantExpr constExpr)
Positions of a Linalg op loops that correspond to different kinds of a contraction dimension.
SmallVector< unsigned, 2 > batch
Positions of a Linalg op loops that correspond to different kinds of a convolution dimension.
SmallVector< unsigned, 2 > depth
SmallVector< unsigned, 2 > outputImage
SmallVector< unsigned, 2 > outputChannel
SmallVector< int64_t, 2 > dilations
SmallVector< int64_t, 2 > strides
SmallVector< unsigned, 2 > inputChannel
SmallVector< unsigned, 2 > batch
SmallVector< unsigned, 2 > filterLoop