MLIR 24.0.0git
AffineOps.cpp
Go to the documentation of this file.
1//===- AffineOps.cpp - MLIR Affine Operations -----------------------------===//
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
14#include "mlir/IR/AffineExpr.h"
16#include "mlir/IR/IRMapping.h"
17#include "mlir/IR/IntegerSet.h"
18#include "mlir/IR/Matchers.h"
21#include "mlir/IR/Value.h"
25#include "llvm/ADT/STLExtras.h"
26#include "llvm/ADT/SmallBitVector.h"
27#include "llvm/ADT/SmallVectorExtras.h"
28#include "llvm/ADT/TypeSwitch.h"
29#include "llvm/Support/DebugLog.h"
30#include "llvm/Support/LogicalResult.h"
31#include "llvm/Support/MathExtras.h"
32#include <numeric>
33#include <optional>
34
35using namespace mlir;
36using namespace mlir::affine;
37
38using llvm::divideCeilSigned;
39using llvm::divideFloorSigned;
40using llvm::mod;
41
42#define DEBUG_TYPE "affine-ops"
43
44#include "mlir/Dialect/Affine/IR/AffineOpsDialect.cpp.inc"
45
46/// A utility function to check if a value is defined at the top level of
47/// `region` or is an argument of `region`. A value of index type defined at the
48/// top level of a `AffineScope` region is always a valid symbol for all
49/// uses in that region.
51 if (auto arg = dyn_cast<BlockArgument>(value))
52 return arg.getParentRegion() == region;
53 return value.getDefiningOp()->getParentRegion() == region;
54}
55
56/// Checks if `value` known to be a legal affine dimension or symbol in `src`
57/// region remains legal if the operation that uses it is inlined into `dest`
58/// with the given value mapping. `legalityCheck` is either `isValidDim` or
59/// `isValidSymbol`, depending on the value being required to remain a valid
60/// dimension or symbol.
61static bool
63 const IRMapping &mapping,
64 function_ref<bool(Value, Region *)> legalityCheck) {
65 // If the value is a valid dimension for any other reason than being
66 // a top-level value, it will remain valid: constants get inlined
67 // with the function, transitive affine applies also get inlined and
68 // will be checked themselves, etc.
69 if (!isTopLevelValue(value, src))
70 return true;
71
72 // If it's a top-level value because it's a block operand, i.e. a
73 // function argument, check whether the value replacing it after
74 // inlining is a valid dimension in the new region.
75 if (llvm::isa<BlockArgument>(value))
76 return legalityCheck(mapping.lookup(value), dest);
77
78 // If it's a top-level value because it's defined in the region,
79 // it can only be inlined if the defining op is a constant or a
80 // `dim`, which can appear anywhere and be valid, since the defining
81 // op won't be top-level anymore after inlining.
82 Attribute operandCst;
83 bool isDimLikeOp = isa<ShapedDimOpInterface>(value.getDefiningOp());
84 return matchPattern(value.getDefiningOp(), m_Constant(&operandCst)) ||
85 isDimLikeOp;
86}
87
88/// Checks if all values known to be legal affine dimensions or symbols in `src`
89/// remain so if their respective users are inlined into `dest`.
90static bool
92 const IRMapping &mapping,
93 function_ref<bool(Value, Region *)> legalityCheck) {
94 return llvm::all_of(values, [&](Value v) {
95 return remainsLegalAfterInline(v, src, dest, mapping, legalityCheck);
96 });
97}
98
99/// Checks if an affine read or write operation remains legal after inlining
100/// from `src` to `dest`.
101template <typename OpTy>
102static bool remainsLegalAfterInline(OpTy op, Region *src, Region *dest,
103 const IRMapping &mapping) {
104 static_assert(llvm::is_one_of<OpTy, AffineReadOpInterface,
105 AffineWriteOpInterface>::value,
106 "only ops with affine read/write interface are supported");
107
108 AffineMap map = op.getAffineMap();
109 ValueRange dimOperands = op.getMapOperands().take_front(map.getNumDims());
110 ValueRange symbolOperands =
111 op.getMapOperands().take_back(map.getNumSymbols());
113 dimOperands, src, dest, mapping,
114 static_cast<bool (*)(Value, Region *)>(isValidDim)))
115 return false;
117 symbolOperands, src, dest, mapping,
118 static_cast<bool (*)(Value, Region *)>(isValidSymbol)))
119 return false;
120 return true;
121}
122
123/// Checks if an affine apply operation remains legal after inlining from `src`
124/// to `dest`.
125// Use "unused attribute" marker to silence clang-tidy warning stemming from
126// the inability to see through "llvm::TypeSwitch".
127template <>
128[[maybe_unused]] bool remainsLegalAfterInline(AffineApplyOp op, Region *src,
129 Region *dest,
130 const IRMapping &mapping) {
131 // If it's a valid dimension, we need to check that it remains so.
132 if (isValidDim(op.getResult(), src))
134 op.getMapOperands(), src, dest, mapping,
135 static_cast<bool (*)(Value, Region *)>(isValidDim));
136
137 // Otherwise it must be a valid symbol, check that it remains so.
139 op.getMapOperands(), src, dest, mapping,
140 static_cast<bool (*)(Value, Region *)>(isValidSymbol));
141}
142
143//===----------------------------------------------------------------------===//
144// AffineDialect Interfaces
145//===----------------------------------------------------------------------===//
146
147namespace {
148/// This class defines the interface for handling inlining with affine
149/// operations.
150struct AffineInlinerInterface : public DialectInlinerInterface {
151 using DialectInlinerInterface::DialectInlinerInterface;
152
153 //===--------------------------------------------------------------------===//
154 // Analysis Hooks
155 //===--------------------------------------------------------------------===//
156
157 /// Returns true if the given region 'src' can be inlined into the region
158 /// 'dest' that is attached to an operation registered to the current dialect.
159 /// 'wouldBeCloned' is set if the region is cloned into its new location
160 /// rather than moved, indicating there may be other users.
161 bool isLegalToInline(Region *dest, Region *src, bool wouldBeCloned,
162 IRMapping &valueMapping) const final {
163 // We can inline into affine loops and conditionals if this doesn't break
164 // affine value categorization rules.
165 Operation *destOp = dest->getParentOp();
166 if (!isa<AffineParallelOp, AffineForOp, AffineIfOp>(destOp))
167 return false;
168
169 // Multi-block regions cannot be inlined into affine constructs, all of
170 // which require single-block regions.
171 if (!src->hasOneBlock())
172 return false;
173
174 // Side-effecting operations that the affine dialect cannot understand
175 // should not be inlined.
176 Block &srcBlock = src->front();
177 for (Operation &op : srcBlock) {
178 // Ops with no side effects are fine,
179 if (auto iface = dyn_cast<MemoryEffectOpInterface>(op)) {
180 if (iface.hasNoEffect())
181 continue;
182 }
183
184 // Assuming the inlined region is valid, we only need to check if the
185 // inlining would change it.
186 bool remainsValid =
187 llvm::TypeSwitch<Operation *, bool>(&op)
188 .Case<AffineApplyOp, AffineReadOpInterface,
189 AffineWriteOpInterface>([&](auto op) {
190 return remainsLegalAfterInline(op, src, dest, valueMapping);
191 })
192 .Default([](Operation *) {
193 // Conservatively disallow inlining ops we cannot reason about.
194 return false;
195 });
196
197 if (!remainsValid)
198 return false;
199 }
200
201 return true;
202 }
203
204 /// Returns true if the given operation 'op', that is registered to this
205 /// dialect, can be inlined into the given region, false otherwise.
206 bool isLegalToInline(Operation *op, Region *region, bool wouldBeCloned,
207 IRMapping &valueMapping) const final {
208 // Always allow inlining affine operations into a region that is marked as
209 // affine scope, or into affine loops and conditionals. There are some edge
210 // cases when inlining *into* affine structures, but that is handled in the
211 // other 'isLegalToInline' hook above.
212 Operation *parentOp = region->getParentOp();
213 return parentOp->hasTrait<OpTrait::AffineScope>() ||
214 isa<AffineForOp, AffineParallelOp, AffineIfOp>(parentOp);
215 }
216
217 /// Affine regions should be analyzed recursively.
218 bool shouldAnalyzeRecursively(Operation *op) const final { return true; }
219};
220} // namespace
221
222//===----------------------------------------------------------------------===//
223// AffineDialect
224//===----------------------------------------------------------------------===//
225
226void AffineDialect::initialize() {
227 addOperations<
228#define GET_OP_LIST
229#include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc"
230 >();
231 addInterfaces<AffineInlinerInterface>();
232 declarePromisedInterfaces<ValueBoundsOpInterface, AffineApplyOp, AffineMaxOp,
233 AffineMinOp>();
234}
235
236/// Materialize a single constant operation from a given attribute value with
237/// the desired resultant type.
238Operation *AffineDialect::materializeConstant(OpBuilder &builder,
239 Attribute value, Type type,
240 Location loc) {
241 if (auto poison = dyn_cast<ub::PoisonAttr>(value))
242 return ub::PoisonOp::create(builder, loc, type, poison);
243 return arith::ConstantOp::materialize(builder, value, type, loc);
244}
245
246/// A utility function to check if a value is defined at the top level of an
247/// op with trait `AffineScope`. If the value is defined in an unlinked region,
248/// conservatively assume it is not top-level. A value of index type defined at
249/// the top level is always a valid symbol.
251 if (auto arg = dyn_cast<BlockArgument>(value)) {
252 // The block owning the argument may be unlinked, e.g. when the surrounding
253 // region has not yet been attached to an Op, at which point the parent Op
254 // is null.
255 Operation *parentOp = arg.getOwner()->getParentOp();
256 return parentOp && parentOp->hasTrait<OpTrait::AffineScope>();
257 }
258 // The defining Op may live in an unlinked block so its parent Op may be null.
259 Operation *parentOp = value.getDefiningOp()->getParentOp();
260 return parentOp && parentOp->hasTrait<OpTrait::AffineScope>();
261}
262
263/// Returns the closest region enclosing `op` that is held by an operation with
264/// trait `AffineScope`; `nullptr` if there is no such region.
266 auto *curOp = op;
267 while (auto *parentOp = curOp->getParentOp()) {
268 if (parentOp->hasTrait<OpTrait::AffineScope>())
269 return curOp->getParentRegion();
270 curOp = parentOp;
271 }
272 return nullptr;
273}
274
276 Operation *curOp = op;
277 while (auto *parentOp = curOp->getParentOp()) {
278 if (!isa<AffineForOp, AffineIfOp, AffineParallelOp>(parentOp))
279 return curOp->getParentRegion();
280 curOp = parentOp;
281 }
282 return nullptr;
283}
284
285// A Value can be used as a dimension id iff it meets one of the following
286// conditions:
287// *) It is valid as a symbol.
288// *) It is an induction variable.
289// *) It is the result of affine apply operation with dimension id arguments.
291 // The value must be an index type.
292 if (!value.getType().isIndex())
293 return false;
294
295 if (auto *defOp = value.getDefiningOp())
296 return isValidDim(value, getAffineScope(defOp));
297
298 // This value has to be a block argument for an op that has the
299 // `AffineScope` trait or an induction var of an affine.for or
300 // affine.parallel.
301 if (isAffineInductionVar(value))
302 return true;
303 auto *parentOp = llvm::cast<BlockArgument>(value).getOwner()->getParentOp();
304 return parentOp && parentOp->hasTrait<OpTrait::AffineScope>();
305}
306
307// Value can be used as a dimension id iff it meets one of the following
308// conditions:
309// *) It is valid as a symbol.
310// *) It is an induction variable.
311// *) It is the result of an affine apply operation with dimension id operands.
312// *) It is the result of a more specialized index transformation (ex.
313// delinearize_index or linearize_index) with dimension id operands.
315 // The value must be an index type.
316 if (!value.getType().isIndex())
317 return false;
318
319 // All valid symbols are okay.
320 if (isValidSymbol(value, region))
321 return true;
322
323 auto *op = value.getDefiningOp();
324 if (!op) {
325 // This value has to be an induction var for an affine.for or an
326 // affine.parallel.
327 return isAffineInductionVar(value);
328 }
329
330 // Affine apply operation is ok if all of its operands are ok.
331 if (auto applyOp = dyn_cast<AffineApplyOp>(op))
332 return applyOp.isValidDim(region);
333 // delinearize_index and linearize_index are special forms of apply
334 // and so are valid dimensions if all their arguments are valid dimensions.
335 if (isa<AffineDelinearizeIndexOp, AffineLinearizeIndexOp>(op))
336 return llvm::all_of(op->getOperands(),
337 [&](Value arg) { return ::isValidDim(arg, region); });
338 // The dim op is okay if its operand memref/tensor is defined at the top
339 // level.
340 if (auto dimOp = dyn_cast<ShapedDimOpInterface>(op))
341 return isTopLevelValue(dimOp.getShapedValue());
342 return false;
343}
344
345/// Returns true if the 'index' dimension of the `memref` defined by
346/// `memrefDefOp` is a statically shaped one or defined using a valid symbol
347/// for `region`.
348template <typename AnyMemRefDefOp>
349static bool isMemRefSizeValidSymbol(AnyMemRefDefOp memrefDefOp, unsigned index,
350 Region *region) {
351 MemRefType memRefType = memrefDefOp.getType();
352
353 // Dimension index is out of bounds.
354 if (index >= memRefType.getRank()) {
355 return false;
356 }
357
358 // Statically shaped.
359 if (!memRefType.isDynamicDim(index))
360 return true;
361 // Get the position of the dimension among dynamic dimensions;
362 unsigned dynamicDimPos = memRefType.getDynamicDimIndex(index);
363 return isValidSymbol(*(memrefDefOp.getDynamicSizes().begin() + dynamicDimPos),
364 region);
365}
366
367/// Returns true if the result of the dim op is a valid symbol for `region`.
368static bool isDimOpValidSymbol(ShapedDimOpInterface dimOp, Region *region) {
369 // The dim op is okay if its source is defined at the top level.
370 if (isTopLevelValue(dimOp.getShapedValue()))
371 return true;
372
373 // Conservatively handle remaining BlockArguments as non-valid symbols.
374 // E.g. scf.for iterArgs.
375 if (llvm::isa<BlockArgument>(dimOp.getShapedValue()))
376 return false;
377
378 // The dim op is also okay if its operand memref is a view/subview whose
379 // corresponding size is a valid symbol.
380 std::optional<int64_t> index = getConstantIntValue(dimOp.getDimension());
381
382 // Be conservative if we can't understand the dimension.
383 if (!index.has_value())
384 return false;
385
386 // Skip over all memref.cast ops (if any).
387 Operation *op = dimOp.getShapedValue().getDefiningOp();
388 while (auto castOp = dyn_cast<memref::CastOp>(op)) {
389 // Bail on unranked memrefs.
390 if (isa<UnrankedMemRefType>(castOp.getSource().getType()))
391 return false;
392 op = castOp.getSource().getDefiningOp();
393 if (!op)
394 return false;
395 }
396
397 int64_t i = index.value();
399 .Case<memref::ViewOp, memref::SubViewOp, memref::AllocOp>(
400 [&](auto op) { return isMemRefSizeValidSymbol(op, i, region); })
401 .Default([](Operation *) { return false; });
402}
403
404// A value can be used as a symbol (at all its use sites) iff it meets one of
405// the following conditions:
406// *) It is a constant.
407// *) Its defining op or block arg appearance is immediately enclosed by an op
408// with `AffineScope` trait.
409// *) It is the result of an affine.apply operation with symbol operands.
410// *) It is a result of the dim op on a memref whose corresponding size is a
411// valid symbol.
413 if (!value)
414 return false;
415
416 // The value must be an index type.
417 if (!value.getType().isIndex())
418 return false;
419
420 // Check that the value is a top level value.
421 if (isTopLevelValue(value))
422 return true;
423
424 if (auto *defOp = value.getDefiningOp())
425 return isValidSymbol(value, getAffineScope(defOp));
426
427 return false;
428}
429
430/// A utility function to check if a value is defined at the top level of
431/// `region` or is an argument of `region` or is defined above the region.
432static bool isTopLevelValueOrAbove(Value value, Region *region) {
433 Region *parentRegion = value.getParentRegion();
434 do {
435 if (parentRegion == region)
436 return true;
437 Operation *regionOp = region->getParentOp();
438 if (regionOp->hasTrait<OpTrait::IsIsolatedFromAbove>())
439 break;
440 region = region->getParentOp()->getParentRegion();
441 } while (region);
442 return false;
443}
444
445/// A value can be used as a symbol for `region` iff it meets one of the
446/// following conditions:
447/// *) It is a constant.
448/// *) It is a result of a `Pure` operation whose operands are valid symbolic
449/// *) identifiers.
450/// *) It is a result of the dim op on a memref whose corresponding size is
451/// a valid symbol.
452/// *) It is defined at the top level of 'region' or is its argument.
453/// *) It dominates `region`'s parent op.
454/// If `region` is null, conservatively assume the symbol definition scope does
455/// not exist and only accept the values that would be symbols regardless of
456/// the surrounding region structure, i.e. the first three cases above.
458 // The value must be an index type.
459 if (!value.getType().isIndex())
460 return false;
461
462 // A top-level value is a valid symbol.
463 if (region && isTopLevelValueOrAbove(value, region))
464 return true;
465
466 auto *defOp = value.getDefiningOp();
467 if (!defOp)
468 return false;
469
470 // Constant operation is ok.
471 Attribute operandCst;
472 if (matchPattern(defOp, m_Constant(&operandCst)))
473 return true;
474
475 // `Pure` operation that whose operands are valid symbolic identifiers.
476 if (isPure(defOp) && llvm::all_of(defOp->getOperands(), [&](Value operand) {
477 return affine::isValidSymbol(operand, region);
478 })) {
479 return true;
480 }
481
482 // Dim op results could be valid symbols at any level.
483 if (auto dimOp = dyn_cast<ShapedDimOpInterface>(defOp))
484 return isDimOpValidSymbol(dimOp, region);
485
486 return false;
487}
488
489// Returns true if 'value' is a valid index to an affine operation (e.g.
490// affine.load, affine.store, affine.dma_start, affine.dma_wait) where
491// `region` provides the polyhedral symbol scope. Returns false otherwise.
492static bool isValidAffineIndexOperand(Value value, Region *region) {
493 return isValidDim(value, region) || isValidSymbol(value, region);
494}
495
496/// Prints dimension and symbol list.
499 unsigned numDims, OpAsmPrinter &printer) {
500 OperandRange operands(begin, end);
501 printer << '(' << operands.take_front(numDims) << ')';
502 if (operands.size() > numDims)
503 printer << '[' << operands.drop_front(numDims) << ']';
504}
505
506/// Parses dimension and symbol list and returns true if parsing failed.
508 OpAsmParser &parser, SmallVectorImpl<Value> &operands, unsigned &numDims) {
511 return failure();
512 // Store number of dimensions for validation by caller.
513 numDims = opInfos.size();
514
515 // Parse the optional symbol operands.
516 auto indexTy = parser.getBuilder().getIndexType();
517 return failure(parser.parseOperandList(
519 parser.resolveOperands(opInfos, indexTy, operands));
520}
521
522/// Utility function to verify that a set of operands are valid dimension and
523/// symbol identifiers. The operands should be laid out such that the dimension
524/// operands are before the symbol operands. This function returns failure if
525/// there was an invalid operand. An operation is provided to emit any necessary
526/// errors.
527template <typename OpTy>
528static LogicalResult
530 unsigned numDims) {
531 unsigned opIt = 0;
532 for (auto operand : operands) {
533 if (opIt++ < numDims) {
534 if (!isValidDim(operand, getAffineScope(op)))
535 return op.emitOpError("operand cannot be used as a dimension id");
536 } else if (!isValidSymbol(operand, getAffineScope(op))) {
537 return op.emitOpError("operand cannot be used as a symbol");
538 }
539 }
540 return success();
541}
542
543//===----------------------------------------------------------------------===//
544// AffineApplyOp
545//===----------------------------------------------------------------------===//
546
547AffineValueMap AffineApplyOp::getAffineValueMap() {
548 return AffineValueMap(getAffineMap(), getOperands(), getResult());
549}
550
551ParseResult AffineApplyOp::parse(OpAsmParser &parser, OperationState &result) {
552 auto &builder = parser.getBuilder();
553 auto indexTy = builder.getIndexType();
554
555 AffineMapAttr mapAttr;
556 unsigned numDims;
557 if (parser.parseAttribute(mapAttr, "map", result.attributes) ||
558 parseDimAndSymbolList(parser, result.operands, numDims) ||
559 parser.parseOptionalAttrDict(result.attributes))
560 return failure();
561 auto map = mapAttr.getValue();
562
563 if (map.getNumDims() != numDims ||
564 numDims + map.getNumSymbols() != result.operands.size()) {
565 return parser.emitError(parser.getNameLoc(),
566 "dimension or symbol index mismatch");
567 }
568
569 result.types.append(map.getNumResults(), indexTy);
570 return success();
571}
572
573void AffineApplyOp::print(OpAsmPrinter &p) {
574 p << " " << getMapAttr();
575 printDimAndSymbolList(operand_begin(), operand_end(),
576 getAffineMap().getNumDims(), p);
577 p.printOptionalAttrDict((*this)->getAttrs(), /*elidedAttrs=*/{"map"});
578}
579
580LogicalResult AffineApplyOp::verify() {
581 // Check input and output dimensions match.
582 AffineMap affineMap = getMap();
583
584 // Verify that operand count matches affine map dimension and symbol count.
585 if (getNumOperands() != affineMap.getNumDims() + affineMap.getNumSymbols())
586 return emitOpError(
587 "operand count and affine map dimension and symbol count must match");
588
589 // Verify that the map only produces one result.
590 if (affineMap.getNumResults() != 1)
591 return emitOpError("mapping must produce one value");
592
593 // Do not allow valid dims to be used in symbol positions. We do allow
594 // affine.apply to use operands for values that may neither qualify as affine
595 // dims or affine symbols due to usage outside of affine ops, analyses, etc.
596 Region *region = getAffineScope(*this);
597 for (Value operand : getMapOperands().drop_front(affineMap.getNumDims())) {
598 if (::isValidDim(operand, region) && !::isValidSymbol(operand, region))
599 return emitError("dimensional operand cannot be used as a symbol");
600 }
601
602 return success();
603}
604
605// The result of the affine apply operation can be used as a dimension id if all
606// its operands are valid dimension ids.
607bool AffineApplyOp::isValidDim() {
608 return llvm::all_of(getOperands(),
609 [](Value op) { return affine::isValidDim(op); });
610}
611
612// The result of the affine apply operation can be used as a dimension id if all
613// its operands are valid dimension ids with the parent operation of `region`
614// defining the polyhedral scope for symbols.
615bool AffineApplyOp::isValidDim(Region *region) {
616 return llvm::all_of(getOperands(),
617 [&](Value op) { return ::isValidDim(op, region); });
618}
619
620// The result of the affine apply operation can be used as a symbol if all its
621// operands are symbols.
622bool AffineApplyOp::isValidSymbol() {
623 return llvm::all_of(getOperands(),
624 [](Value op) { return affine::isValidSymbol(op); });
625}
626
627// The result of the affine apply operation can be used as a symbol in `region`
628// if all its operands are symbols in `region`.
629bool AffineApplyOp::isValidSymbol(Region *region) {
630 return llvm::all_of(getOperands(), [&](Value operand) {
631 return affine::isValidSymbol(operand, region);
632 });
633}
634
635OpFoldResult AffineApplyOp::fold(FoldAdaptor adaptor) {
636 auto map = getAffineMap();
637
638 // Fold dims and symbols to existing values.
639 auto expr = map.getResult(0);
640 if (auto dim = dyn_cast<AffineDimExpr>(expr))
641 return getOperand(dim.getPosition());
642 if (auto sym = dyn_cast<AffineSymbolExpr>(expr))
643 return getOperand(map.getNumDims() + sym.getPosition());
644
645 // Otherwise, default to folding the map.
647 bool hasPoison = false;
648 auto foldResult =
649 map.constantFold(adaptor.getMapOperands(), result, &hasPoison);
650 if (hasPoison)
651 return ub::PoisonAttr::get(getContext());
652 if (failed(foldResult))
653 return {};
654 return result[0];
655}
656
657/// Returns the largest known divisor of `e`. Exploits information from the
658/// values in `operands`.
660 // This method isn't aware of `operands`.
662
663 // We now make use of operands for the case `e` is a dim expression.
664 // TODO: More powerful simplification would have to modify
665 // getLargestKnownDivisor to take `operands` and exploit that information as
666 // well for dim/sym expressions, but in that case, getLargestKnownDivisor
667 // can't be part of the IR library but of the `Analysis` library. The IR
668 // library can only really depend on simple O(1) checks.
669 auto dimExpr = dyn_cast<AffineDimExpr>(e);
670 // If it's not a dim expr, `div` is the best we have.
671 if (!dimExpr)
672 return div;
673
674 // We simply exploit information from loop IVs.
675 // We don't need to use mlir::getLargestKnownDivisorOfValue since the other
676 // desired simplifications are expected to be part of other
677 // canonicalizations. Also, mlir::getLargestKnownDivisorOfValue is part of the
678 // LoopAnalysis library.
679 Value operand = operands[dimExpr.getPosition()];
680 int64_t operandDivisor = 1;
681 // TODO: With the right accessors, this can be extended to
682 // LoopLikeOpInterface.
683 if (AffineForOp forOp = getForInductionVarOwner(operand)) {
684 if (forOp.hasConstantLowerBound() && forOp.getConstantLowerBound() == 0) {
685 operandDivisor = forOp.getStepAsInt();
686 } else {
687 uint64_t lbLargestKnownDivisor =
688 forOp.getLowerBoundMap().getLargestKnownDivisorOfMapExprs();
689 operandDivisor = std::gcd(lbLargestKnownDivisor, forOp.getStepAsInt());
690 }
691 }
692 return operandDivisor;
693}
694
695/// Check if `e` is known to be: 0 <= `e` < `k`. Handles the simple cases of `e`
696/// being an affine dim expression or a constant.
698 int64_t k) {
699 if (auto constExpr = dyn_cast<AffineConstantExpr>(e)) {
700 int64_t constVal = constExpr.getValue();
701 return constVal >= 0 && constVal < k;
702 }
703 auto dimExpr = dyn_cast<AffineDimExpr>(e);
704 if (!dimExpr)
705 return false;
706 Value operand = operands[dimExpr.getPosition()];
707 // TODO: With the right accessors, this can be extended to
708 // LoopLikeOpInterface.
709 if (AffineForOp forOp = getForInductionVarOwner(operand)) {
710 if (forOp.hasConstantLowerBound() && forOp.getConstantLowerBound() >= 0 &&
711 forOp.hasConstantUpperBound() && forOp.getConstantUpperBound() <= k) {
712 return true;
713 }
714 }
715
716 // We don't consider other cases like `operand` being defined by a constant or
717 // an affine.apply op since such cases will already be handled by other
718 // patterns and propagation of loop IVs or constant would happen.
719 return false;
720}
721
722/// Check if expression `e` is of the form d*e_1 + e_2 where 0 <= e_2 < d.
723/// Set `div` to `d`, `quotientTimesDiv` to e_1 and `rem` to e_2 if the
724/// expression is in that form.
726 AffineExpr &quotientTimesDiv, AffineExpr &rem) {
727 auto bin = dyn_cast<AffineBinaryOpExpr>(e);
728 if (!bin || bin.getKind() != AffineExprKind::Add)
729 return false;
730
731 AffineExpr llhs = bin.getLHS();
732 AffineExpr rlhs = bin.getRHS();
733 div = getLargestKnownDivisor(llhs, operands);
734 if (isNonNegativeBoundedBy(rlhs, operands, div)) {
735 quotientTimesDiv = llhs;
736 rem = rlhs;
737 return true;
738 }
739 div = getLargestKnownDivisor(rlhs, operands);
740 if (isNonNegativeBoundedBy(llhs, operands, div)) {
741 quotientTimesDiv = rlhs;
742 rem = llhs;
743 return true;
744 }
745 return false;
746}
747
748/// Gets the constant lower bound on an `iv`.
749static std::optional<int64_t> getLowerBound(Value iv) {
750 AffineForOp forOp = getForInductionVarOwner(iv);
751 if (forOp && forOp.hasConstantLowerBound())
752 return forOp.getConstantLowerBound();
753 return std::nullopt;
754}
755
756/// Gets the constant upper bound on an affine.for `iv`.
757static std::optional<int64_t> getUpperBound(Value iv) {
758 AffineForOp forOp = getForInductionVarOwner(iv);
759 if (!forOp || !forOp.hasConstantUpperBound())
760 return std::nullopt;
761
762 // If its lower bound is also known, we can get a more precise bound
763 // whenever the step is not one.
764 if (forOp.hasConstantLowerBound()) {
765 return forOp.getConstantUpperBound() - 1 -
766 (forOp.getConstantUpperBound() - forOp.getConstantLowerBound() - 1) %
767 forOp.getStepAsInt();
768 }
769 return forOp.getConstantUpperBound() - 1;
770}
771
772/// Determine a constant upper bound for `expr` if one exists while exploiting
773/// values in `operands`. Note that the upper bound is an inclusive one. `expr`
774/// is guaranteed to be less than or equal to it.
775static std::optional<int64_t> getUpperBound(AffineExpr expr, unsigned numDims,
776 unsigned numSymbols,
777 ArrayRef<Value> operands) {
778 if (auto constExpr = dyn_cast<AffineConstantExpr>(expr))
779 return constExpr.getValue();
780
781 // Get the constant lower or upper bounds on the operands.
782 SmallVector<std::optional<int64_t>> constLowerBounds, constUpperBounds;
783 constLowerBounds.reserve(operands.size());
784 constUpperBounds.reserve(operands.size());
785 for (Value operand : operands) {
786 constLowerBounds.push_back(getLowerBound(operand));
787 constUpperBounds.push_back(getUpperBound(operand));
788 }
789
790 return getBoundForAffineExpr(expr, numDims, numSymbols, constLowerBounds,
791 constUpperBounds,
792 /*isUpper=*/true);
793}
794
795/// Determine a constant lower bound for `expr` if one exists while exploiting
796/// values in `operands`. Note that the upper bound is an inclusive one. `expr`
797/// is guaranteed to be less than or equal to it.
798static std::optional<int64_t> getLowerBound(AffineExpr expr, unsigned numDims,
799 unsigned numSymbols,
800 ArrayRef<Value> operands) {
801 if (auto constExpr = dyn_cast<AffineConstantExpr>(expr))
802 return constExpr.getValue();
803
804 // Get the constant lower or upper bounds on the operands.
805 SmallVector<std::optional<int64_t>> constLowerBounds, constUpperBounds;
806 constLowerBounds.reserve(operands.size());
807 constUpperBounds.reserve(operands.size());
808 for (Value operand : operands) {
809 constLowerBounds.push_back(getLowerBound(operand));
810 constUpperBounds.push_back(getUpperBound(operand));
811 }
812
813 return getBoundForAffineExpr(expr, numDims, numSymbols, constLowerBounds,
814 constUpperBounds,
815 /*isUpper=*/false);
816}
817
818/// Simplify `expr` while exploiting information from the values in `operands`.
819static void simplifyExprAndOperands(AffineExpr &expr, unsigned numDims,
820 unsigned numSymbols,
821 ArrayRef<Value> operands) {
822 // We do this only for certain floordiv/mod expressions.
823 auto binExpr = dyn_cast<AffineBinaryOpExpr>(expr);
824 if (!binExpr)
825 return;
826
827 // Simplify the child expressions first.
828 AffineExpr lhs = binExpr.getLHS();
829 AffineExpr rhs = binExpr.getRHS();
830 simplifyExprAndOperands(lhs, numDims, numSymbols, operands);
831 simplifyExprAndOperands(rhs, numDims, numSymbols, operands);
832 expr = getAffineBinaryOpExpr(binExpr.getKind(), lhs, rhs);
833
834 binExpr = dyn_cast<AffineBinaryOpExpr>(expr);
835 if (!binExpr || (expr.getKind() != AffineExprKind::FloorDiv &&
837 expr.getKind() != AffineExprKind::Mod)) {
838 return;
839 }
840
841 // The `lhs` and `rhs` may be different post construction of simplified expr.
842 lhs = binExpr.getLHS();
843 rhs = binExpr.getRHS();
844 auto rhsConst = dyn_cast<AffineConstantExpr>(rhs);
845 if (!rhsConst)
846 return;
847
848 int64_t rhsConstVal = rhsConst.getValue();
849 // Undefined exprsessions aren't touched; IR can still be valid with them.
850 if (rhsConstVal <= 0)
851 return;
852
853 // Exploit constant lower/upper bounds to simplify a floordiv or mod.
854 MLIRContext *context = expr.getContext();
855 std::optional<int64_t> lhsLbConst =
856 getLowerBound(lhs, numDims, numSymbols, operands);
857 std::optional<int64_t> lhsUbConst =
858 getUpperBound(lhs, numDims, numSymbols, operands);
859 if (lhsLbConst && lhsUbConst) {
860 int64_t lhsLbConstVal = *lhsLbConst;
861 int64_t lhsUbConstVal = *lhsUbConst;
862 // lhs floordiv c is a single value lhs is bounded in a range `c` that has
863 // the same quotient.
864 if (binExpr.getKind() == AffineExprKind::FloorDiv &&
865 divideFloorSigned(lhsLbConstVal, rhsConstVal) ==
866 divideFloorSigned(lhsUbConstVal, rhsConstVal)) {
868 divideFloorSigned(lhsLbConstVal, rhsConstVal), context);
869 return;
870 }
871 // lhs ceildiv c is a single value if the entire range has the same ceil
872 // quotient.
873 if (binExpr.getKind() == AffineExprKind::CeilDiv &&
874 divideCeilSigned(lhsLbConstVal, rhsConstVal) ==
875 divideCeilSigned(lhsUbConstVal, rhsConstVal)) {
876 expr = getAffineConstantExpr(divideCeilSigned(lhsLbConstVal, rhsConstVal),
877 context);
878 return;
879 }
880 // lhs mod c is lhs if the entire range has quotient 0 w.r.t the rhs.
881 if (binExpr.getKind() == AffineExprKind::Mod && lhsLbConstVal >= 0 &&
882 lhsLbConstVal < rhsConstVal && lhsUbConstVal < rhsConstVal) {
883 expr = lhs;
884 return;
885 }
886 }
887
888 // Simplify expressions of the form e = (e_1 + e_2) floordiv c or (e_1 + e_2)
889 // mod c, where e_1 is a multiple of `k` and 0 <= e_2 < k. In such cases, if
890 // `c` % `k` == 0, (e_1 + e_2) floordiv c can be simplified to e_1 floordiv c.
891 // And when k % c == 0, (e_1 + e_2) mod c can be simplified to e_2 mod c.
892 AffineExpr quotientTimesDiv, rem;
893 int64_t divisor;
894 if (isQTimesDPlusR(lhs, operands, divisor, quotientTimesDiv, rem)) {
895 if (rhsConstVal % divisor == 0 &&
896 binExpr.getKind() == AffineExprKind::FloorDiv) {
897 expr = quotientTimesDiv.floorDiv(rhsConst);
898 } else if (divisor % rhsConstVal == 0 &&
899 binExpr.getKind() == AffineExprKind::Mod) {
900 expr = rem % rhsConst;
901 }
902 return;
903 }
904
905 // Handle the simple case when the LHS expression can be either upper
906 // bounded or is a known multiple of RHS constant.
907 // lhs floordiv c -> 0 if 0 <= lhs < c,
908 // lhs mod c -> 0 if lhs % c = 0.
909 if ((isNonNegativeBoundedBy(lhs, operands, rhsConstVal) &&
910 binExpr.getKind() == AffineExprKind::FloorDiv) ||
911 (getLargestKnownDivisor(lhs, operands) % rhsConstVal == 0 &&
912 binExpr.getKind() == AffineExprKind::Mod)) {
913 expr = getAffineConstantExpr(0, expr.getContext());
914 }
915}
916
917/// Simplify the expressions in `map` while making use of lower or upper bounds
918/// of its operands. If `isMax` is true, the map is to be treated as a max of
919/// its result expressions, and min otherwise. Eg: min (d0, d1) -> (8, 4 * d0 +
920/// d1) can be simplified to (8) if the operands are respectively lower bounded
921/// by 2 and 0 (the second expression can't be lower than 8).
923 ArrayRef<Value> operands,
924 bool isMax) {
925 // Can't simplify.
926 if (operands.empty())
927 return;
928
929 // Get the upper or lower bound on an affine.for op IV using its range.
930 // Get the constant lower or upper bounds on the operands.
931 SmallVector<std::optional<int64_t>> constLowerBounds, constUpperBounds;
932 constLowerBounds.reserve(operands.size());
933 constUpperBounds.reserve(operands.size());
934 for (Value operand : operands) {
935 constLowerBounds.push_back(getLowerBound(operand));
936 constUpperBounds.push_back(getUpperBound(operand));
937 }
938
939 // We will compute the lower and upper bounds on each of the expressions
940 // Then, we will check (depending on max or min) as to whether a specific
941 // bound is redundant by checking if its highest (in case of max) and its
942 // lowest (in the case of min) value is already lower than (or higher than)
943 // the lower bound (or upper bound in the case of min) of another bound.
944 SmallVector<std::optional<int64_t>, 4> lowerBounds, upperBounds;
945 lowerBounds.reserve(map.getNumResults());
946 upperBounds.reserve(map.getNumResults());
947 for (AffineExpr e : map.getResults()) {
948 if (auto constExpr = dyn_cast<AffineConstantExpr>(e)) {
949 lowerBounds.push_back(constExpr.getValue());
950 upperBounds.push_back(constExpr.getValue());
951 } else {
952 lowerBounds.push_back(
954 constLowerBounds, constUpperBounds,
955 /*isUpper=*/false));
956 upperBounds.push_back(
958 constLowerBounds, constUpperBounds,
959 /*isUpper=*/true));
960 }
961 }
962
963 // Collect expressions that are not redundant.
964 SmallVector<AffineExpr, 4> irredundantExprs;
965 for (auto exprEn : llvm::enumerate(map.getResults())) {
966 AffineExpr e = exprEn.value();
967 unsigned i = exprEn.index();
968 // Some expressions can be turned into constants.
969 if (lowerBounds[i] && upperBounds[i] && *lowerBounds[i] == *upperBounds[i])
970 e = getAffineConstantExpr(*lowerBounds[i], e.getContext());
971
972 // Check if the expression is redundant.
973 if (isMax) {
974 if (!upperBounds[i]) {
975 irredundantExprs.push_back(e);
976 continue;
977 }
978 // If there exists another expression such that its lower bound is greater
979 // than this expression's upper bound, it's redundant.
980 if (!llvm::any_of(llvm::enumerate(lowerBounds), [&](const auto &en) {
981 auto otherLowerBound = en.value();
982 unsigned pos = en.index();
983 if (pos == i || !otherLowerBound)
984 return false;
985 if (*otherLowerBound > *upperBounds[i])
986 return true;
987 if (*otherLowerBound < *upperBounds[i])
988 return false;
989 // Equality case. When both expressions are considered redundant, we
990 // don't want to get both of them. We keep the one that appears
991 // first.
992 if (upperBounds[pos] && lowerBounds[i] &&
993 lowerBounds[i] == upperBounds[i] &&
994 otherLowerBound == *upperBounds[pos] && i < pos)
995 return false;
996 return true;
997 }))
998 irredundantExprs.push_back(e);
999 } else {
1000 if (!lowerBounds[i]) {
1001 irredundantExprs.push_back(e);
1002 continue;
1003 }
1004 // Likewise for the `min` case. Use the complement of the condition above.
1005 if (!llvm::any_of(llvm::enumerate(upperBounds), [&](const auto &en) {
1006 auto otherUpperBound = en.value();
1007 unsigned pos = en.index();
1008 if (pos == i || !otherUpperBound)
1009 return false;
1010 if (*otherUpperBound < *lowerBounds[i])
1011 return true;
1012 if (*otherUpperBound > *lowerBounds[i])
1013 return false;
1014 if (lowerBounds[pos] && upperBounds[i] &&
1015 lowerBounds[i] == upperBounds[i] &&
1016 otherUpperBound == lowerBounds[pos] && i < pos)
1017 return false;
1018 return true;
1019 }))
1020 irredundantExprs.push_back(e);
1021 }
1022 }
1023
1024 // Create the map without the redundant expressions.
1025 map = AffineMap::get(map.getNumDims(), map.getNumSymbols(), irredundantExprs,
1026 map.getContext());
1027}
1028
1029/// Simplify the map while exploiting information on the values in `operands`.
1030// Use "unused attribute" marker to silence warning stemming from the inability
1031// to see through the template expansion.
1032[[maybe_unused]] static void simplifyMapWithOperands(AffineMap &map,
1033 ArrayRef<Value> operands) {
1034 assert(map.getNumInputs() == operands.size() && "invalid operands for map");
1035 SmallVector<AffineExpr> newResults;
1036 newResults.reserve(map.getNumResults());
1037 for (AffineExpr expr : map.getResults()) {
1039 operands);
1040 newResults.push_back(expr);
1041 }
1042 map = AffineMap::get(map.getNumDims(), map.getNumSymbols(), newResults,
1043 map.getContext());
1044}
1045
1046/// Assuming `dimOrSym` is a quantity in the apply op map `map` and defined by
1047/// `minOp = affine_min(x_1, ..., x_n)`. This function checks that:
1048/// `0 < affine_min(x_1, ..., x_n)` and proceeds with replacing the patterns:
1049/// ```
1050/// dimOrSym.ceildiv(x_k)
1051/// (dimOrSym + x_k - 1).floordiv(x_k)
1052/// ```
1053/// by `1` for all `k` in `1, ..., n`. This is possible because `x / x_k <= 1`.
1054///
1055///
1056/// Warning: ValueBoundsConstraintSet::computeConstantBound is needed to check
1057/// `minOp` is positive.
1058static LogicalResult replaceAffineMinBoundingBoxExpression(AffineMinOp minOp,
1059 AffineExpr dimOrSym,
1060 AffineMap *map,
1061 ValueRange dims,
1062 ValueRange syms) {
1063 LDBG() << "replaceAffineMinBoundingBoxExpression: `" << minOp << "`";
1064 AffineMap affineMinMap = minOp.getAffineMap();
1065
1066 // Check the value is positive.
1067 for (unsigned i = 0, e = affineMinMap.getNumResults(); i < e; ++i) {
1068 // Compare each expression in the minimum against 0.
1070 getAsIndexOpFoldResult(minOp.getContext(), 0),
1073 minOp.getOperands())))
1074 return failure();
1075 }
1076
1077 /// Convert affine symbols and dimensions in minOp to symbols or dimensions in
1078 /// the apply op affine map.
1079 DenseMap<AffineExpr, AffineExpr> dimSymConversionTable;
1080 SmallVector<unsigned> unmappedDims, unmappedSyms;
1081 for (auto [i, dim] : llvm::enumerate(minOp.getDimOperands())) {
1082 auto it = llvm::find(dims, dim);
1083 if (it == dims.end()) {
1084 unmappedDims.push_back(i);
1085 continue;
1086 }
1087 dimSymConversionTable[getAffineDimExpr(i, minOp.getContext())] =
1088 getAffineDimExpr(it.getIndex(), minOp.getContext());
1089 }
1090 for (auto [i, sym] : llvm::enumerate(minOp.getSymbolOperands())) {
1091 auto it = llvm::find(syms, sym);
1092 if (it == syms.end()) {
1093 unmappedSyms.push_back(i);
1094 continue;
1095 }
1096 dimSymConversionTable[getAffineSymbolExpr(i, minOp.getContext())] =
1097 getAffineSymbolExpr(it.getIndex(), minOp.getContext());
1098 }
1099
1100 // Create the replacement map.
1102 AffineExpr c1 = getAffineConstantExpr(1, minOp.getContext());
1103 for (AffineExpr expr : affineMinMap.getResults()) {
1104 // If we cannot express the result in terms of the apply map symbols and
1105 // sims then continue.
1106 if (llvm::any_of(unmappedDims,
1107 [&](unsigned i) { return expr.isFunctionOfDim(i); }) ||
1108 llvm::any_of(unmappedSyms,
1109 [&](unsigned i) { return expr.isFunctionOfSymbol(i); }))
1110 continue;
1111
1112 AffineExpr convertedExpr = expr.replace(dimSymConversionTable);
1113
1114 // dimOrSym.ceilDiv(expr) -> 1
1115 repl[dimOrSym.ceilDiv(convertedExpr)] = c1;
1116 // (dimOrSym + expr - 1).floorDiv(expr) -> 1
1117 repl[(dimOrSym + convertedExpr - 1).floorDiv(convertedExpr)] = c1;
1118 }
1119 AffineMap initialMap = *map;
1120 *map = initialMap.replace(repl, initialMap.getNumDims(),
1121 initialMap.getNumSymbols());
1122 return success(*map != initialMap);
1123}
1124
1125/// Recursively traverse `e`. If `e` or one of its sub-expressions has the form
1126/// e1 + e2 + ... + eK, where the e_i are a super(multi)set of `exprsToRemove`,
1127/// place a map between e and `newVal` + sum({e1, e2, .. eK} - exprsToRemove)
1128/// into `replacementsMap`. If no entries were added to `replacementsMap`,
1129/// nothing was found.
1131 AffineExpr e, const llvm::SmallDenseSet<AffineExpr, 4> &exprsToRemove,
1132 AffineExpr newVal, DenseMap<AffineExpr, AffineExpr> &replacementsMap) {
1133 auto binOp = dyn_cast<AffineBinaryOpExpr>(e);
1134 if (!binOp)
1135 return;
1136 AffineExpr lhs = binOp.getLHS();
1137 AffineExpr rhs = binOp.getRHS();
1138 if (binOp.getKind() != AffineExprKind::Add) {
1139 shortenAddChainsContainingAll(lhs, exprsToRemove, newVal, replacementsMap);
1140 shortenAddChainsContainingAll(rhs, exprsToRemove, newVal, replacementsMap);
1141 return;
1142 }
1143 SmallVector<AffineExpr> toPreserve;
1144 llvm::SmallDenseSet<AffineExpr, 4> ourTracker(exprsToRemove);
1145 AffineExpr thisTerm = rhs;
1146 AffineExpr nextTerm = lhs;
1147
1148 while (thisTerm) {
1149 if (!ourTracker.erase(thisTerm)) {
1150 toPreserve.push_back(thisTerm);
1151 shortenAddChainsContainingAll(thisTerm, exprsToRemove, newVal,
1152 replacementsMap);
1153 }
1154 auto nextBinOp = dyn_cast_if_present<AffineBinaryOpExpr>(nextTerm);
1155 if (!nextBinOp || nextBinOp.getKind() != AffineExprKind::Add) {
1156 thisTerm = nextTerm;
1157 nextTerm = AffineExpr();
1158 } else {
1159 thisTerm = nextBinOp.getRHS();
1160 nextTerm = nextBinOp.getLHS();
1161 }
1162 }
1163 if (!ourTracker.empty())
1164 return;
1165 // We reverse the terms to be preserved here in order to preserve
1166 // associativity between them.
1167 AffineExpr newExpr = newVal;
1168 for (AffineExpr preserved : llvm::reverse(toPreserve))
1169 newExpr = newExpr + preserved;
1170 replacementsMap.insert({e, newExpr});
1171}
1172
1173/// If this map contains of the expression `x_1 + x_1 * C_1 + ... x_n * C_N +
1174/// ...` (not necessarily in order) where the set of the `x_i` is the set of
1175/// outputs of an `affine.delinearize_index` whos inverse is that expression,
1176/// replace that expression with the input of that delinearize_index op.
1177///
1178/// `unitDimInput` is the input that was detected as the potential start to this
1179/// replacement chain - if it isn't the rightmost result of the delinearization,
1180/// this method fails. (This is intended to ensure we don't have redundant scans
1181/// over the same expression).
1182///
1183/// While this currently only handles delinearizations with a constant basis,
1184/// that isn't a fundamental limitation.
1185///
1186/// This is a utility function for `replaceDimOrSym` below.
1188 AffineDelinearizeIndexOp delinOp, Value resultToReplace, AffineMap *map,
1190 if (!delinOp.getDynamicBasis().empty())
1191 return failure();
1192 if (resultToReplace != delinOp.getMultiIndex().back())
1193 return failure();
1194
1195 MLIRContext *ctx = delinOp.getContext();
1196 SmallVector<AffineExpr> resToExpr(delinOp.getNumResults(), AffineExpr());
1197 for (auto [pos, dim] : llvm::enumerate(dims)) {
1198 auto asResult = dyn_cast_if_present<OpResult>(dim);
1199 if (!asResult)
1200 continue;
1201 if (asResult.getOwner() == delinOp.getOperation())
1202 resToExpr[asResult.getResultNumber()] = getAffineDimExpr(pos, ctx);
1203 }
1204 for (auto [pos, sym] : llvm::enumerate(syms)) {
1205 auto asResult = dyn_cast_if_present<OpResult>(sym);
1206 if (!asResult)
1207 continue;
1208 if (asResult.getOwner() == delinOp.getOperation())
1209 resToExpr[asResult.getResultNumber()] = getAffineSymbolExpr(pos, ctx);
1210 }
1211 if (llvm::is_contained(resToExpr, AffineExpr()))
1212 return failure();
1213
1214 bool isDimReplacement = llvm::all_of(resToExpr, llvm::IsaPred<AffineDimExpr>);
1215 int64_t stride = 1;
1216 llvm::SmallDenseSet<AffineExpr, 4> expectedExprs;
1217 // This isn't zip_equal since sometimes the delinearize basis is missing a
1218 // size for the first result.
1219 for (auto [binding, size] : llvm::zip(
1220 llvm::reverse(resToExpr), llvm::reverse(delinOp.getStaticBasis()))) {
1221 expectedExprs.insert(binding * getAffineConstantExpr(stride, ctx));
1222 stride *= size;
1223 }
1224 if (resToExpr.size() != delinOp.getStaticBasis().size())
1225 expectedExprs.insert(resToExpr[0] * stride);
1226
1228 AffineExpr delinInExpr = isDimReplacement
1229 ? getAffineDimExpr(dims.size(), ctx)
1230 : getAffineSymbolExpr(syms.size(), ctx);
1231
1232 for (AffineExpr e : map->getResults())
1233 shortenAddChainsContainingAll(e, expectedExprs, delinInExpr, replacements);
1234 if (replacements.empty())
1235 return failure();
1236
1237 AffineMap origMap = *map;
1238 if (isDimReplacement)
1239 dims.push_back(delinOp.getLinearIndex());
1240 else
1241 syms.push_back(delinOp.getLinearIndex());
1242 *map = origMap.replace(replacements, dims.size(), syms.size());
1243
1244 // Blank out dead dimensions and symbols
1245 for (AffineExpr e : resToExpr) {
1246 if (auto d = dyn_cast<AffineDimExpr>(e)) {
1247 unsigned pos = d.getPosition();
1248 if (!map->isFunctionOfDim(pos))
1249 dims[pos] = nullptr;
1250 }
1251 if (auto s = dyn_cast<AffineSymbolExpr>(e)) {
1252 unsigned pos = s.getPosition();
1253 if (!map->isFunctionOfSymbol(pos))
1254 syms[pos] = nullptr;
1255 }
1256 }
1257 return success();
1258}
1259
1260/// Replace all occurrences of AffineExpr at position `pos` in `map` by the
1261/// defining AffineApplyOp expression and operands.
1262/// When `dimOrSymbolPosition < dims.size()`, AffineDimExpr@[pos] is replaced.
1263/// When `dimOrSymbolPosition >= dims.size()`,
1264/// AffineSymbolExpr@[pos - dims.size()] is replaced.
1265/// Mutate `map`,`dims` and `syms` in place as follows:
1266/// 1. `dims` and `syms` are only appended to.
1267/// 2. `map` dim and symbols are gradually shifted to higher positions.
1268/// 3. Old `dim` and `sym` entries are replaced by nullptr
1269/// This avoids the need for any bookkeeping.
1270/// If `replaceAffineMin` is set to true, additionally triggers more expensive
1271/// replacements involving affine_min operations.
1272static LogicalResult replaceDimOrSym(AffineMap *map,
1273 unsigned dimOrSymbolPosition,
1276 bool replaceAffineMin) {
1277 MLIRContext *ctx = map->getContext();
1278 bool isDimReplacement = (dimOrSymbolPosition < dims.size());
1279 unsigned pos = isDimReplacement ? dimOrSymbolPosition
1280 : dimOrSymbolPosition - dims.size();
1281 Value &v = isDimReplacement ? dims[pos] : syms[pos];
1282 if (!v)
1283 return failure();
1284
1285 if (auto minOp = v.getDefiningOp<AffineMinOp>(); minOp && replaceAffineMin) {
1286 AffineExpr dimOrSym = isDimReplacement ? getAffineDimExpr(pos, ctx)
1287 : getAffineSymbolExpr(pos, ctx);
1288 return replaceAffineMinBoundingBoxExpression(minOp, dimOrSym, map, dims,
1289 syms);
1290 }
1291
1292 if (auto delinOp = v.getDefiningOp<affine::AffineDelinearizeIndexOp>()) {
1293 return replaceAffineDelinearizeIndexInverseExpression(delinOp, v, map, dims,
1294 syms);
1295 }
1296
1297 auto affineApply = v.getDefiningOp<AffineApplyOp>();
1298 if (!affineApply)
1299 return failure();
1300
1301 // At this point we will perform a replacement of `v`, set the entry in `dim`
1302 // or `sym` to nullptr immediately.
1303 v = nullptr;
1304
1305 // Compute the map, dims and symbols coming from the AffineApplyOp.
1306 AffineMap composeMap = affineApply.getAffineMap();
1307 assert(composeMap.getNumResults() == 1 && "affine.apply with >1 results");
1308 SmallVector<Value> composeOperands(affineApply.getMapOperands().begin(),
1309 affineApply.getMapOperands().end());
1310 // Canonicalize the map to promote dims to symbols when possible. This is to
1311 // avoid generating invalid maps.
1312 canonicalizeMapAndOperands(&composeMap, &composeOperands);
1313 AffineExpr replacementExpr =
1314 composeMap.shiftDims(dims.size()).shiftSymbols(syms.size()).getResult(0);
1315 ValueRange composeDims =
1316 ArrayRef<Value>(composeOperands).take_front(composeMap.getNumDims());
1317 ValueRange composeSyms =
1318 ArrayRef<Value>(composeOperands).take_back(composeMap.getNumSymbols());
1319 AffineExpr toReplace = isDimReplacement ? getAffineDimExpr(pos, ctx)
1320 : getAffineSymbolExpr(pos, ctx);
1321
1322 // Append the dims and symbols where relevant and perform the replacement.
1323 dims.append(composeDims.begin(), composeDims.end());
1324 syms.append(composeSyms.begin(), composeSyms.end());
1325 *map = map->replace(toReplace, replacementExpr, dims.size(), syms.size());
1326
1327 return success();
1328}
1329
1330/// Iterate over `operands` and fold away all those produced by an AffineApplyOp
1331/// iteratively. Perform canonicalization of map and operands as well as
1332/// AffineMap simplification. `map` and `operands` are mutated in place.
1334 SmallVectorImpl<Value> *operands,
1335 bool composeAffineMin = false) {
1336 if (map->getNumResults() == 0) {
1337 canonicalizeMapAndOperands(map, operands);
1338 *map = simplifyAffineMap(*map);
1339 return;
1340 }
1341
1342 MLIRContext *ctx = map->getContext();
1343 SmallVector<Value, 4> dims(operands->begin(),
1344 operands->begin() + map->getNumDims());
1345 SmallVector<Value, 4> syms(operands->begin() + map->getNumDims(),
1346 operands->end());
1347
1348 // Iterate over dims and symbols coming from AffineApplyOp and replace until
1349 // exhaustion. This iteratively mutates `map`, `dims` and `syms`. Both `dims`
1350 // and `syms` can only increase by construction.
1351 // The implementation uses a `while` loop to support the case of symbols
1352 // that may be constructed from dims ;this may be overkill.
1353 while (true) {
1354 bool changed = false;
1355 for (unsigned pos = 0; pos != dims.size() + syms.size(); ++pos)
1356 if ((changed |=
1357 succeeded(replaceDimOrSym(map, pos, dims, syms, composeAffineMin))))
1358 break;
1359 if (!changed)
1360 break;
1361 }
1362
1363 // Clear operands so we can fill them anew.
1364 operands->clear();
1365
1366 // At this point we may have introduced null operands, prune them out before
1367 // canonicalizing map and operands.
1368 unsigned nDims = 0, nSyms = 0;
1369 SmallVector<AffineExpr, 4> dimReplacements, symReplacements;
1370 dimReplacements.reserve(dims.size());
1371 symReplacements.reserve(syms.size());
1372 for (auto *container : {&dims, &syms}) {
1373 bool isDim = (container == &dims);
1374 auto &repls = isDim ? dimReplacements : symReplacements;
1375 for (const auto &en : llvm::enumerate(*container)) {
1376 Value v = en.value();
1377 if (!v) {
1378 assert(isDim ? !map->isFunctionOfDim(en.index())
1379 : !map->isFunctionOfSymbol(en.index()) &&
1380 "map is function of unexpected expr@pos");
1381 repls.push_back(getAffineConstantExpr(0, ctx));
1382 continue;
1383 }
1384 repls.push_back(isDim ? getAffineDimExpr(nDims++, ctx)
1385 : getAffineSymbolExpr(nSyms++, ctx));
1386 operands->push_back(v);
1387 }
1388 }
1389 *map = map->replaceDimsAndSymbols(dimReplacements, symReplacements, nDims,
1390 nSyms);
1391
1392 // Canonicalize and simplify before returning.
1393 canonicalizeMapAndOperands(map, operands);
1394 *map = simplifyAffineMap(*map);
1395}
1396
1398 AffineMap *map, SmallVectorImpl<Value> *operands, bool composeAffineMin) {
1399 while (llvm::any_of(*operands, [](Value v) {
1400 return isa_and_nonnull<AffineApplyOp>(v.getDefiningOp());
1401 })) {
1402 composeAffineMapAndOperands(map, operands, composeAffineMin);
1403 }
1404 // Additional trailing step for AffineMinOps in case no chains of AffineApply.
1405 if (composeAffineMin && llvm::any_of(*operands, [](Value v) {
1406 return isa_and_nonnull<AffineMinOp>(v.getDefiningOp());
1407 })) {
1408 composeAffineMapAndOperands(map, operands, composeAffineMin);
1409 }
1410}
1411
1412AffineApplyOp
1414 ArrayRef<OpFoldResult> operands,
1415 bool composeAffineMin) {
1416 SmallVector<Value> valueOperands;
1417 map = foldAttributesIntoMap(b, map, operands, valueOperands);
1418 composeAffineMapAndOperands(&map, &valueOperands, composeAffineMin);
1419 assert(map);
1420 return AffineApplyOp::create(b, loc, map, valueOperands);
1421}
1422
1423AffineApplyOp
1425 ArrayRef<OpFoldResult> operands,
1426 bool composeAffineMin) {
1428 b, loc,
1430 .front(),
1431 operands, composeAffineMin);
1432}
1433
1434/// Composes the given affine map with the given list of operands, pulling in
1435/// the maps from any affine.apply operations that supply the operands.
1437 SmallVectorImpl<Value> &operands,
1438 bool composeAffineMin = false) {
1439 // Compose and canonicalize each expression in the map individually because
1440 // composition only applies to single-result maps, collecting potentially
1441 // duplicate operands in a single list with shifted dimensions and symbols.
1442 SmallVector<Value> dims, symbols;
1444 for (unsigned i : llvm::seq<unsigned>(0, map.getNumResults())) {
1445 SmallVector<Value> submapOperands(operands.begin(), operands.end());
1446 AffineMap submap = map.getSubMap({i});
1447 fullyComposeAffineMapAndOperands(&submap, &submapOperands,
1448 composeAffineMin);
1449 canonicalizeMapAndOperands(&submap, &submapOperands);
1450 unsigned numNewDims = submap.getNumDims();
1451 submap = submap.shiftDims(dims.size()).shiftSymbols(symbols.size());
1452 llvm::append_range(dims,
1453 ArrayRef<Value>(submapOperands).take_front(numNewDims));
1454 llvm::append_range(symbols,
1455 ArrayRef<Value>(submapOperands).drop_front(numNewDims));
1456 exprs.push_back(submap.getResult(0));
1457 }
1458
1459 // Canonicalize the map created from composed expressions to deduplicate the
1460 // dimension and symbol operands.
1461 operands = llvm::to_vector(llvm::concat<Value>(dims, symbols));
1462 map = AffineMap::get(dims.size(), symbols.size(), exprs, map.getContext());
1463 canonicalizeMapAndOperands(&map, &operands);
1464}
1465
1468 bool composeAffineMin) {
1469 assert(map.getNumResults() == 1 && "building affine.apply with !=1 result");
1470
1471 // Create new builder without a listener, so that no notification is
1472 // triggered if the op is folded.
1473 // TODO: OpBuilder::createOrFold should return OpFoldResults, then this
1474 // workaround is no longer needed.
1475 OpBuilder newBuilder(b.getContext());
1476 newBuilder.setInsertionPoint(b.getInsertionBlock(), b.getInsertionPoint());
1477
1478 // Create op.
1479 AffineApplyOp applyOp =
1480 makeComposedAffineApply(newBuilder, loc, map, operands, composeAffineMin);
1481
1482 // Get constant operands.
1483 SmallVector<Attribute> constOperands(applyOp->getNumOperands());
1484 for (unsigned i = 0, e = constOperands.size(); i != e; ++i)
1485 matchPattern(applyOp->getOperand(i), m_Constant(&constOperands[i]));
1486
1487 // Try to fold the operation.
1488 SmallVector<OpFoldResult> foldResults;
1489 if (failed(applyOp->fold(constOperands, foldResults)) ||
1490 foldResults.empty()) {
1491 if (OpBuilder::Listener *listener = b.getListener())
1492 listener->notifyOperationInserted(applyOp, /*previous=*/{});
1493 return applyOp.getResult();
1494 }
1495
1496 applyOp->erase();
1497 return llvm::getSingleElement(foldResults);
1498}
1499
1501 OpBuilder &b, Location loc, AffineExpr expr,
1502 ArrayRef<OpFoldResult> operands, bool composeAffineMin) {
1504 b, loc,
1506 .front(),
1507 operands, composeAffineMin);
1508}
1509
1513 bool composeAffineMin) {
1514 return llvm::map_to_vector(
1515 llvm::seq<unsigned>(0, map.getNumResults()), [&](unsigned i) {
1516 return makeComposedFoldedAffineApply(b, loc, map.getSubMap({i}),
1517 operands, composeAffineMin);
1518 });
1519}
1520
1521template <typename OpTy>
1523 ArrayRef<OpFoldResult> operands) {
1524 SmallVector<Value> valueOperands;
1525 map = foldAttributesIntoMap(b, map, operands, valueOperands);
1526 composeMultiResultAffineMap(map, valueOperands);
1527 return OpTy::create(b, loc, b.getIndexType(), map, valueOperands);
1528}
1529
1530AffineMinOp
1535
1536template <typename OpTy>
1538 AffineMap map,
1539 ArrayRef<OpFoldResult> operands) {
1540 // Create new builder without a listener, so that no notification is
1541 // triggered if the op is folded.
1542 // TODO: OpBuilder::createOrFold should return OpFoldResults, then this
1543 // workaround is no longer needed.
1544 OpBuilder newBuilder(b.getContext());
1545 newBuilder.setInsertionPoint(b.getInsertionBlock(), b.getInsertionPoint());
1546
1547 // Create op.
1548 auto minMaxOp = makeComposedMinMax<OpTy>(newBuilder, loc, map, operands);
1549
1550 // Get constant operands.
1551 SmallVector<Attribute> constOperands(minMaxOp->getNumOperands());
1552 for (unsigned i = 0, e = constOperands.size(); i != e; ++i)
1553 matchPattern(minMaxOp->getOperand(i), m_Constant(&constOperands[i]));
1554
1555 // Try to fold the operation.
1556 SmallVector<OpFoldResult> foldResults;
1557 if (failed(minMaxOp->fold(constOperands, foldResults)) ||
1558 foldResults.empty()) {
1559 if (OpBuilder::Listener *listener = b.getListener())
1560 listener->notifyOperationInserted(minMaxOp, /*previous=*/{});
1561 return minMaxOp.getResult();
1562 }
1563
1564 minMaxOp->erase();
1565 return llvm::getSingleElement(foldResults);
1566}
1567
1574
1581
1582// A symbol may appear as a dim in affine.apply operations. This function
1583// canonicalizes dims that are valid symbols into actual symbols.
1584template <class MapOrSet>
1585static void canonicalizePromotedSymbols(MapOrSet *mapOrSet,
1586 SmallVectorImpl<Value> *operands) {
1587 if (!mapOrSet || operands->empty())
1588 return;
1589
1590 assert(mapOrSet->getNumInputs() == operands->size() &&
1591 "map/set inputs must match number of operands");
1592
1593 auto *context = mapOrSet->getContext();
1594 SmallVector<Value, 8> resultOperands;
1595 resultOperands.reserve(operands->size());
1596 SmallVector<Value, 8> remappedSymbols;
1597 remappedSymbols.reserve(operands->size());
1598 unsigned nextDim = 0;
1599 unsigned nextSym = 0;
1600 unsigned oldNumSyms = mapOrSet->getNumSymbols();
1601 SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims());
1602 for (unsigned i = 0, e = mapOrSet->getNumInputs(); i != e; ++i) {
1603 if (i < mapOrSet->getNumDims()) {
1604 if (isValidSymbol((*operands)[i])) {
1605 // This is a valid symbol that appears as a dim, canonicalize it.
1606 dimRemapping[i] = getAffineSymbolExpr(oldNumSyms + nextSym++, context);
1607 remappedSymbols.push_back((*operands)[i]);
1608 } else {
1609 dimRemapping[i] = getAffineDimExpr(nextDim++, context);
1610 resultOperands.push_back((*operands)[i]);
1611 }
1612 } else {
1613 resultOperands.push_back((*operands)[i]);
1614 }
1615 }
1616
1617 resultOperands.append(remappedSymbols.begin(), remappedSymbols.end());
1618 *operands = resultOperands;
1619 *mapOrSet = mapOrSet->replaceDimsAndSymbols(
1620 dimRemapping, /*symReplacements=*/{}, nextDim, oldNumSyms + nextSym);
1621
1622 assert(mapOrSet->getNumInputs() == operands->size() &&
1623 "map/set inputs must match number of operands");
1624}
1625
1626/// A valid affine dimension may appear as a symbol in affine.apply operations.
1627/// Given an application of `operands` to an affine map or integer set
1628/// `mapOrSet`, this function canonicalizes symbols of `mapOrSet` that are valid
1629/// dims, but not valid symbols into actual dims. Without such a legalization,
1630/// the affine.apply will be invalid. This method is the exact inverse of
1631/// canonicalizePromotedSymbols.
1632template <class MapOrSet>
1633static void legalizeDemotedDims(MapOrSet &mapOrSet,
1634 SmallVectorImpl<Value> &operands) {
1635 if (!mapOrSet || operands.empty())
1636 return;
1637
1638 unsigned numOperands = operands.size();
1639
1640 assert(mapOrSet.getNumInputs() == numOperands &&
1641 "map/set inputs must match number of operands");
1642
1643 auto *context = mapOrSet.getContext();
1644 SmallVector<Value, 8> resultOperands;
1645 resultOperands.reserve(numOperands);
1646 SmallVector<Value, 8> remappedDims;
1647 remappedDims.reserve(numOperands);
1648 SmallVector<Value, 8> symOperands;
1649 symOperands.reserve(mapOrSet.getNumSymbols());
1650 unsigned nextSym = 0;
1651 unsigned nextDim = 0;
1652 unsigned oldNumDims = mapOrSet.getNumDims();
1653 SmallVector<AffineExpr, 8> symRemapping(mapOrSet.getNumSymbols());
1654 resultOperands.assign(operands.begin(), operands.begin() + oldNumDims);
1655 for (unsigned i = oldNumDims, e = mapOrSet.getNumInputs(); i != e; ++i) {
1656 if (operands[i] && isValidDim(operands[i]) && !isValidSymbol(operands[i])) {
1657 // This is a valid dim that appears as a symbol, legalize it.
1658 symRemapping[i - oldNumDims] =
1659 getAffineDimExpr(oldNumDims + nextDim++, context);
1660 remappedDims.push_back(operands[i]);
1661 } else {
1662 symRemapping[i - oldNumDims] = getAffineSymbolExpr(nextSym++, context);
1663 symOperands.push_back(operands[i]);
1664 }
1665 }
1666
1667 append_range(resultOperands, remappedDims);
1668 append_range(resultOperands, symOperands);
1669 operands = resultOperands;
1670 mapOrSet = mapOrSet.replaceDimsAndSymbols(
1671 /*dimReplacements=*/{}, symRemapping, oldNumDims + nextDim, nextSym);
1672
1673 assert(mapOrSet.getNumInputs() == operands.size() &&
1674 "map/set inputs must match number of operands");
1675}
1676
1677// Works for either an affine map or an integer set.
1678template <class MapOrSet>
1679static void canonicalizeMapOrSetAndOperands(MapOrSet *mapOrSet,
1680 SmallVectorImpl<Value> *operands) {
1681 static_assert(llvm::is_one_of<MapOrSet, AffineMap, IntegerSet>::value,
1682 "Argument must be either of AffineMap or IntegerSet type");
1683
1684 if (!mapOrSet || operands->empty())
1685 return;
1686
1687 assert(mapOrSet->getNumInputs() == operands->size() &&
1688 "map/set inputs must match number of operands");
1689
1690 canonicalizePromotedSymbols<MapOrSet>(mapOrSet, operands);
1691 legalizeDemotedDims<MapOrSet>(*mapOrSet, *operands);
1692
1693 // Check to see what dims are used.
1694 llvm::SmallBitVector usedDims(mapOrSet->getNumDims());
1695 llvm::SmallBitVector usedSyms(mapOrSet->getNumSymbols());
1696 mapOrSet->walkExprs([&](AffineExpr expr) {
1697 if (auto dimExpr = dyn_cast<AffineDimExpr>(expr))
1698 usedDims[dimExpr.getPosition()] = true;
1699 else if (auto symExpr = dyn_cast<AffineSymbolExpr>(expr))
1700 usedSyms[symExpr.getPosition()] = true;
1701 });
1702
1703 auto *context = mapOrSet->getContext();
1704
1705 SmallVector<Value, 8> resultOperands;
1706 resultOperands.reserve(operands->size());
1707
1708 llvm::SmallDenseMap<Value, AffineExpr, 8> seenDims;
1709 SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims());
1710 unsigned nextDim = 0;
1711 for (unsigned i = 0, e = mapOrSet->getNumDims(); i != e; ++i) {
1712 if (usedDims[i]) {
1713 // Remap dim positions for duplicate operands.
1714 auto it = seenDims.find((*operands)[i]);
1715 if (it == seenDims.end()) {
1716 dimRemapping[i] = getAffineDimExpr(nextDim++, context);
1717 resultOperands.push_back((*operands)[i]);
1718 seenDims.insert(std::make_pair((*operands)[i], dimRemapping[i]));
1719 } else {
1720 dimRemapping[i] = it->second;
1721 }
1722 }
1723 }
1724 llvm::SmallDenseMap<Value, AffineExpr, 8> seenSymbols;
1725 SmallVector<AffineExpr, 8> symRemapping(mapOrSet->getNumSymbols());
1726 unsigned nextSym = 0;
1727 for (unsigned i = 0, e = mapOrSet->getNumSymbols(); i != e; ++i) {
1728 if (!usedSyms[i])
1729 continue;
1730 // Handle constant operands (only needed for symbolic operands since
1731 // constant operands in dimensional positions would have already been
1732 // promoted to symbolic positions above).
1733 IntegerAttr operandCst;
1734 if (matchPattern((*operands)[i + mapOrSet->getNumDims()],
1735 m_Constant(&operandCst))) {
1736 symRemapping[i] =
1737 getAffineConstantExpr(operandCst.getValue().getSExtValue(), context);
1738 continue;
1739 }
1740 // Remap symbol positions for duplicate operands.
1741 auto it = seenSymbols.find((*operands)[i + mapOrSet->getNumDims()]);
1742 if (it == seenSymbols.end()) {
1743 symRemapping[i] = getAffineSymbolExpr(nextSym++, context);
1744 resultOperands.push_back((*operands)[i + mapOrSet->getNumDims()]);
1745 seenSymbols.insert(std::make_pair((*operands)[i + mapOrSet->getNumDims()],
1746 symRemapping[i]));
1747 } else {
1748 symRemapping[i] = it->second;
1749 }
1750 }
1751 *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, symRemapping,
1752 nextDim, nextSym);
1753 *operands = resultOperands;
1754}
1755
1760
1765
1766namespace {
1767/// Simplify AffineApply, AffineLoad, and AffineStore operations by composing
1768/// maps that supply results into them.
1769///
1770template <typename AffineOpTy>
1771struct SimplifyAffineOp : public OpRewritePattern<AffineOpTy> {
1772 using OpRewritePattern<AffineOpTy>::OpRewritePattern;
1773
1774 /// Replace the affine op with another instance of it with the supplied
1775 /// map and mapOperands.
1776 void replaceAffineOp(PatternRewriter &rewriter, AffineOpTy affineOp,
1777 AffineMap map, ArrayRef<Value> mapOperands) const;
1778
1779 LogicalResult matchAndRewrite(AffineOpTy affineOp,
1780 PatternRewriter &rewriter) const override {
1781 static_assert(
1782 llvm::is_one_of<AffineOpTy, AffineLoadOp, AffinePrefetchOp,
1783 AffineStoreOp, AffineApplyOp, AffineMinOp, AffineMaxOp,
1784 AffineVectorStoreOp, AffineVectorLoadOp>::value,
1785 "affine load/store/vectorstore/vectorload/apply/prefetch/min/max op "
1786 "expected");
1787 auto map = affineOp.getAffineMap();
1788 AffineMap oldMap = map;
1789 auto oldOperands = affineOp.getMapOperands();
1790 SmallVector<Value, 8> resultOperands(oldOperands);
1791 composeAffineMapAndOperands(&map, &resultOperands);
1792 canonicalizeMapAndOperands(&map, &resultOperands);
1793 simplifyMapWithOperands(map, resultOperands);
1794 if (map == oldMap && std::equal(oldOperands.begin(), oldOperands.end(),
1795 resultOperands.begin()))
1796 return failure();
1797
1798 replaceAffineOp(rewriter, affineOp, map, resultOperands);
1799 return success();
1800 }
1801};
1802
1803// Specialize the template to account for the different build signatures for
1804// affine load, store, and apply ops.
1805template <>
1806void SimplifyAffineOp<AffineLoadOp>::replaceAffineOp(
1807 PatternRewriter &rewriter, AffineLoadOp load, AffineMap map,
1808 ArrayRef<Value> mapOperands) const {
1809 rewriter.replaceOpWithNewOp<AffineLoadOp>(load, load.getMemRef(), map,
1810 mapOperands, load.getMaybeAlign());
1811}
1812template <>
1813void SimplifyAffineOp<AffinePrefetchOp>::replaceAffineOp(
1814 PatternRewriter &rewriter, AffinePrefetchOp prefetch, AffineMap map,
1815 ArrayRef<Value> mapOperands) const {
1816 rewriter.replaceOpWithNewOp<AffinePrefetchOp>(
1817 prefetch, prefetch.getMemref(), map, mapOperands, prefetch.getIsWrite(),
1818 prefetch.getLocalityHint(), prefetch.getIsDataCache());
1819}
1820template <>
1821void SimplifyAffineOp<AffineStoreOp>::replaceAffineOp(
1822 PatternRewriter &rewriter, AffineStoreOp store, AffineMap map,
1823 ArrayRef<Value> mapOperands) const {
1824 rewriter.replaceOpWithNewOp<AffineStoreOp>(
1825 store, store.getValueToStore(), store.getMemRef(), map, mapOperands,
1826 store.getMaybeAlign());
1827}
1828template <>
1829void SimplifyAffineOp<AffineVectorLoadOp>::replaceAffineOp(
1830 PatternRewriter &rewriter, AffineVectorLoadOp vectorload, AffineMap map,
1831 ArrayRef<Value> mapOperands) const {
1832 rewriter.replaceOpWithNewOp<AffineVectorLoadOp>(
1833 vectorload, vectorload.getVectorType(), vectorload.getMemRef(), map,
1834 mapOperands, vectorload.getMaybeAlign());
1835}
1836template <>
1837void SimplifyAffineOp<AffineVectorStoreOp>::replaceAffineOp(
1838 PatternRewriter &rewriter, AffineVectorStoreOp vectorstore, AffineMap map,
1839 ArrayRef<Value> mapOperands) const {
1840 rewriter.replaceOpWithNewOp<AffineVectorStoreOp>(
1841 vectorstore, vectorstore.getValueToStore(), vectorstore.getMemRef(), map,
1842 mapOperands, vectorstore.getMaybeAlign());
1843}
1844
1845// Generic version for ops that don't have extra operands.
1846template <typename AffineOpTy>
1847void SimplifyAffineOp<AffineOpTy>::replaceAffineOp(
1848 PatternRewriter &rewriter, AffineOpTy op, AffineMap map,
1849 ArrayRef<Value> mapOperands) const {
1850 rewriter.replaceOpWithNewOp<AffineOpTy>(op, map, mapOperands);
1851}
1852} // namespace
1853
1854void AffineApplyOp::getCanonicalizationPatterns(RewritePatternSet &results,
1855 MLIRContext *context) {
1856 results.add<SimplifyAffineOp<AffineApplyOp>>(context);
1857}
1858
1859//===----------------------------------------------------------------------===//
1860// AffineDmaStartOp
1861//===----------------------------------------------------------------------===//
1862
1863// TODO: Check that map operands are loop IVs or symbols.
1864void AffineDmaStartOp::build(OpBuilder &builder, OperationState &result,
1865 Value srcMemRef, AffineMap srcMap,
1866 ValueRange srcIndices, Value destMemRef,
1867 AffineMap dstMap, ValueRange destIndices,
1868 Value tagMemRef, AffineMap tagMap,
1869 ValueRange tagIndices, Value numElements,
1870 Value stride, Value elementsPerStride) {
1871 result.addOperands(srcMemRef);
1872 result.addAttribute(getSrcMapAttrStrName(), AffineMapAttr::get(srcMap));
1873 result.addOperands(srcIndices);
1874 result.addOperands(destMemRef);
1875 result.addAttribute(getDstMapAttrStrName(), AffineMapAttr::get(dstMap));
1876 result.addOperands(destIndices);
1877 result.addOperands(tagMemRef);
1878 result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap));
1879 result.addOperands(tagIndices);
1880 result.addOperands(numElements);
1881 if (stride) {
1882 result.addOperands({stride, elementsPerStride});
1883 }
1884}
1885
1886void AffineDmaStartOp::print(OpAsmPrinter &p) {
1887 p << " " << getSrcMemRef() << '[';
1888 p.printAffineMapOfSSAIds(getSrcMapAttr(), getSrcIndices());
1889 p << "], " << getDstMemRef() << '[';
1890 p.printAffineMapOfSSAIds(getDstMapAttr(), getDstIndices());
1891 p << "], " << getTagMemRef() << '[';
1892 p.printAffineMapOfSSAIds(getTagMapAttr(), getTagIndices());
1893 p << "], " << getNumElements();
1894 if (isStrided()) {
1895 p << ", " << getStride();
1896 p << ", " << getNumElementsPerStride();
1897 }
1898 p << " : " << getSrcMemRefType() << ", " << getDstMemRefType() << ", "
1899 << getTagMemRefType();
1900}
1901
1902// Parse AffineDmaStartOp.
1903// Ex:
1904// affine.dma_start %src[%i, %j], %dst[%k, %l], %tag[%index], %size,
1905// %stride, %num_elt_per_stride
1906// : memref<3076 x f32, 0>, memref<1024 x f32, 2>, memref<1 x i32>
1907//
1908ParseResult AffineDmaStartOp::parse(OpAsmParser &parser,
1910 OpAsmParser::UnresolvedOperand srcMemRefInfo;
1911 AffineMapAttr srcMapAttr;
1913 OpAsmParser::UnresolvedOperand dstMemRefInfo;
1914 AffineMapAttr dstMapAttr;
1916 OpAsmParser::UnresolvedOperand tagMemRefInfo;
1917 AffineMapAttr tagMapAttr;
1919 OpAsmParser::UnresolvedOperand numElementsInfo;
1921
1923 auto indexType = parser.getBuilder().getIndexType();
1924
1925 // Parse and resolve the following list of operands:
1926 // *) dst memref followed by its affine maps operands (in square brackets).
1927 // *) src memref followed by its affine map operands (in square brackets).
1928 // *) tag memref followed by its affine map operands (in square brackets).
1929 // *) number of elements transferred by DMA operation.
1930 if (parser.parseOperand(srcMemRefInfo) ||
1931 parser.parseAffineMapOfSSAIds(srcMapOperands, srcMapAttr,
1932 getSrcMapAttrStrName(),
1933 result.attributes) ||
1934 parser.parseComma() || parser.parseOperand(dstMemRefInfo) ||
1935 parser.parseAffineMapOfSSAIds(dstMapOperands, dstMapAttr,
1936 getDstMapAttrStrName(),
1937 result.attributes) ||
1938 parser.parseComma() || parser.parseOperand(tagMemRefInfo) ||
1939 parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr,
1940 getTagMapAttrStrName(),
1941 result.attributes) ||
1942 parser.parseComma() || parser.parseOperand(numElementsInfo))
1943 return failure();
1944
1945 // Parse optional stride and elements per stride.
1946 if (parser.parseTrailingOperandList(strideInfo))
1947 return failure();
1948
1949 if (!strideInfo.empty() && strideInfo.size() != 2) {
1950 return parser.emitError(parser.getNameLoc(),
1951 "expected two stride related operands");
1952 }
1953 bool isStrided = strideInfo.size() == 2;
1954
1955 if (parser.parseColonTypeList(types))
1956 return failure();
1957
1958 if (types.size() != 3)
1959 return parser.emitError(parser.getNameLoc(), "expected three types");
1960
1961 if (parser.resolveOperand(srcMemRefInfo, types[0], result.operands) ||
1962 parser.resolveOperands(srcMapOperands, indexType, result.operands) ||
1963 parser.resolveOperand(dstMemRefInfo, types[1], result.operands) ||
1964 parser.resolveOperands(dstMapOperands, indexType, result.operands) ||
1965 parser.resolveOperand(tagMemRefInfo, types[2], result.operands) ||
1966 parser.resolveOperands(tagMapOperands, indexType, result.operands) ||
1967 parser.resolveOperand(numElementsInfo, indexType, result.operands))
1968 return failure();
1969
1970 if (isStrided) {
1971 if (parser.resolveOperands(strideInfo, indexType, result.operands))
1972 return failure();
1973 }
1974
1975 // Check that src/dst/tag operand counts match their map.numInputs.
1976 if (srcMapOperands.size() != srcMapAttr.getValue().getNumInputs() ||
1977 dstMapOperands.size() != dstMapAttr.getValue().getNumInputs() ||
1978 tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
1979 return parser.emitError(parser.getNameLoc(),
1980 "memref operand count not equal to map.numInputs");
1981 return success();
1982}
1983
1984LogicalResult AffineDmaStartOp::verify() {
1985 if (!llvm::isa<MemRefType>(getOperand(getSrcMemRefOperandIndex()).getType()))
1986 return emitOpError("expected DMA source to be of memref type");
1987 if (!llvm::isa<MemRefType>(getOperand(getDstMemRefOperandIndex()).getType()))
1988 return emitOpError("expected DMA destination to be of memref type");
1989 if (!llvm::isa<MemRefType>(getOperand(getTagMemRefOperandIndex()).getType()))
1990 return emitOpError("expected DMA tag to be of memref type");
1991
1992 unsigned numInputsAllMaps = getSrcMap().getNumInputs() +
1993 getDstMap().getNumInputs() +
1994 getTagMap().getNumInputs();
1995 if (getNumOperands() != numInputsAllMaps + 3 + 1 &&
1996 getNumOperands() != numInputsAllMaps + 3 + 1 + 2) {
1997 return emitOpError("incorrect number of operands");
1998 }
1999
2000 Region *scope = getAffineScope(*this);
2001 for (auto idx : getSrcIndices()) {
2002 if (!idx.getType().isIndex())
2003 return emitOpError("src index to dma_start must have 'index' type");
2004 if (!isValidAffineIndexOperand(idx, scope))
2005 return emitOpError(
2006 "src index must be a valid dimension or symbol identifier");
2007 }
2008 for (auto idx : getDstIndices()) {
2009 if (!idx.getType().isIndex())
2010 return emitOpError("dst index to dma_start must have 'index' type");
2011 if (!isValidAffineIndexOperand(idx, scope))
2012 return emitOpError(
2013 "dst index must be a valid dimension or symbol identifier");
2014 }
2015 for (auto idx : getTagIndices()) {
2016 if (!idx.getType().isIndex())
2017 return emitOpError("tag index to dma_start must have 'index' type");
2018 if (!isValidAffineIndexOperand(idx, scope))
2019 return emitOpError(
2020 "tag index must be a valid dimension or symbol identifier");
2021 }
2022 return success();
2023}
2024
2025LogicalResult AffineDmaStartOp::fold(FoldAdaptor adaptor,
2027 /// dma_start(memrefcast) -> dma_start
2028 return memref::foldMemRefCast(*this);
2029}
2030
2031void AffineDmaStartOp::getEffects(
2033 &effects) {
2034 effects.emplace_back(MemoryEffects::Read::get(), &getSrcMemRefMutable(),
2036 effects.emplace_back(MemoryEffects::Write::get(), &getDstMemRefMutable(),
2038 effects.emplace_back(MemoryEffects::Read::get(), &getTagMemRefMutable(),
2040}
2041
2042//===----------------------------------------------------------------------===//
2043// AffineDmaWaitOp
2044//===----------------------------------------------------------------------===//
2045
2046// TODO: Check that map operands are loop IVs or symbols.
2047void AffineDmaWaitOp::build(OpBuilder &builder, OperationState &result,
2048 Value tagMemRef, AffineMap tagMap,
2049 ValueRange tagIndices, Value numElements) {
2050 result.addOperands(tagMemRef);
2051 result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap));
2052 result.addOperands(tagIndices);
2053 result.addOperands(numElements);
2054}
2055
2056void AffineDmaWaitOp::print(OpAsmPrinter &p) {
2057 p << " " << getTagMemRef() << '[';
2058 SmallVector<Value, 2> operands(getTagIndices());
2059 p.printAffineMapOfSSAIds(getTagMapAttr(), operands);
2060 p << "], ";
2062 p << " : " << getTagMemRef().getType();
2063}
2064
2065// Parse AffineDmaWaitOp.
2066// Eg:
2067// affine.dma_wait %tag[%index], %num_elements
2068// : memref<1 x i32, (d0) -> (d0), 4>
2069//
2070ParseResult AffineDmaWaitOp::parse(OpAsmParser &parser,
2072 OpAsmParser::UnresolvedOperand tagMemRefInfo;
2073 AffineMapAttr tagMapAttr;
2075 Type type;
2076 auto indexType = parser.getBuilder().getIndexType();
2077 OpAsmParser::UnresolvedOperand numElementsInfo;
2078
2079 // Parse tag memref, its map operands, and dma size.
2080 if (parser.parseOperand(tagMemRefInfo) ||
2081 parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr,
2082 getTagMapAttrStrName(),
2083 result.attributes) ||
2084 parser.parseComma() || parser.parseOperand(numElementsInfo) ||
2085 parser.parseColonType(type) ||
2086 parser.resolveOperand(tagMemRefInfo, type, result.operands) ||
2087 parser.resolveOperands(tagMapOperands, indexType, result.operands) ||
2088 parser.resolveOperand(numElementsInfo, indexType, result.operands))
2089 return failure();
2090
2091 if (!llvm::isa<MemRefType>(type))
2092 return parser.emitError(parser.getNameLoc(),
2093 "expected tag to be of memref type");
2094
2095 if (tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
2096 return parser.emitError(parser.getNameLoc(),
2097 "tag memref operand count != to map.numInputs");
2098 return success();
2099}
2100
2101LogicalResult AffineDmaWaitOp::verify() {
2102 if (!llvm::isa<MemRefType>(getOperand(0).getType()))
2103 return emitOpError("expected DMA tag to be of memref type");
2104 Region *scope = getAffineScope(*this);
2105 for (auto idx : getTagIndices()) {
2106 if (!idx.getType().isIndex())
2107 return emitOpError("index to dma_wait must have 'index' type");
2108 if (!isValidAffineIndexOperand(idx, scope))
2109 return emitOpError(
2110 "index must be a valid dimension or symbol identifier");
2111 }
2112 return success();
2113}
2114
2115LogicalResult AffineDmaWaitOp::fold(FoldAdaptor adaptor,
2117 /// dma_wait(memrefcast) -> dma_wait
2118 return memref::foldMemRefCast(*this);
2119}
2120
2121void AffineDmaWaitOp::getEffects(
2123 &effects) {
2124 effects.emplace_back(MemoryEffects::Read::get(), &getTagMemRefMutable(),
2126}
2127
2128//===----------------------------------------------------------------------===//
2129// AffineForOp
2130//===----------------------------------------------------------------------===//
2131
2132/// 'bodyBuilder' is used to build the body of affine.for. If iterArgs and
2133/// bodyBuilder are empty/null, we include default terminator op.
2134void AffineForOp::build(OpBuilder &builder, OperationState &result,
2135 ValueRange lbOperands, AffineMap lbMap,
2136 ValueRange ubOperands, AffineMap ubMap, int64_t step,
2137 ValueRange iterArgs, BodyBuilderFn bodyBuilder) {
2138 assert(((!lbMap && lbOperands.empty()) ||
2139 lbOperands.size() == lbMap.getNumInputs()) &&
2140 "lower bound operand count does not match the affine map");
2141 assert(((!ubMap && ubOperands.empty()) ||
2142 ubOperands.size() == ubMap.getNumInputs()) &&
2143 "upper bound operand count does not match the affine map");
2144 assert(step > 0 && "step has to be a positive integer constant");
2145
2146 OpBuilder::InsertionGuard guard(builder);
2147
2148 // Set variadic segment sizes.
2149 result.addAttribute(
2150 getOperandSegmentSizeAttr(),
2151 builder.getDenseI32ArrayAttr({static_cast<int32_t>(lbOperands.size()),
2152 static_cast<int32_t>(ubOperands.size()),
2153 static_cast<int32_t>(iterArgs.size())}));
2154
2155 for (Value val : iterArgs)
2156 result.addTypes(val.getType());
2157
2158 // Add an attribute for the step.
2159 result.addAttribute(getStepAttrName(result.name),
2160 builder.getIntegerAttr(builder.getIndexType(), step));
2161
2162 // Add the lower bound.
2163 result.addAttribute(getLowerBoundMapAttrName(result.name),
2164 AffineMapAttr::get(lbMap));
2165 result.addOperands(lbOperands);
2166
2167 // Add the upper bound.
2168 result.addAttribute(getUpperBoundMapAttrName(result.name),
2169 AffineMapAttr::get(ubMap));
2170 result.addOperands(ubOperands);
2171
2172 result.addOperands(iterArgs);
2173 // Create a region and a block for the body. The argument of the region is
2174 // the loop induction variable.
2175 Region *bodyRegion = result.addRegion();
2176 Block *bodyBlock = builder.createBlock(bodyRegion);
2177 Value inductionVar =
2178 bodyBlock->addArgument(builder.getIndexType(), result.location);
2179 for (Value val : iterArgs)
2180 bodyBlock->addArgument(val.getType(), val.getLoc());
2181
2182 // Create the default terminator if the builder is not provided and if the
2183 // iteration arguments are not provided. Otherwise, leave this to the caller
2184 // because we don't know which values to return from the loop.
2185 if (iterArgs.empty() && !bodyBuilder) {
2186 ensureTerminator(*bodyRegion, builder, result.location);
2187 } else if (bodyBuilder) {
2188 OpBuilder::InsertionGuard guard(builder);
2189 builder.setInsertionPointToStart(bodyBlock);
2190 bodyBuilder(builder, result.location, inductionVar,
2191 bodyBlock->getArguments().drop_front());
2192 }
2193}
2194
2195void AffineForOp::build(OpBuilder &builder, OperationState &result, int64_t lb,
2196 int64_t ub, int64_t step, ValueRange iterArgs,
2197 BodyBuilderFn bodyBuilder) {
2198 auto lbMap = AffineMap::getConstantMap(lb, builder.getContext());
2199 auto ubMap = AffineMap::getConstantMap(ub, builder.getContext());
2200 return build(builder, result, {}, lbMap, {}, ubMap, step, iterArgs,
2201 bodyBuilder);
2202}
2203
2204LogicalResult AffineForOp::verify() {
2205 auto *body = getBody();
2206 if (body->getNumArguments() == 0 || !getInductionVar().getType().isIndex())
2207 return emitOpError("expected body to have an index argument for the "
2208 "induction variable");
2209
2210 return success();
2211}
2212
2213LogicalResult AffineForOp::verifyRegions() {
2214 // Step must be a strictly positive integer.
2215 if (getStepAsInt() <= 0)
2216 return emitOpError("expected step to be a positive integer, got ")
2217 << getStepAsInt();
2218
2219 // Verify that the bound operands are valid dimension/symbols.
2220 /// Lower bound.
2221 if (getLowerBoundMap().getNumInputs() > 0)
2223 getLowerBoundMap().getNumDims())))
2224 return failure();
2225 /// Upper bound.
2226 if (getUpperBoundMap().getNumInputs() > 0)
2228 getUpperBoundMap().getNumDims())))
2229 return failure();
2230 if (getLowerBoundMap().getNumResults() < 1)
2231 return emitOpError("expected lower bound map to have at least one result");
2232 if (getUpperBoundMap().getNumResults() < 1)
2233 return emitOpError("expected upper bound map to have at least one result");
2234
2235 unsigned opNumResults = getNumResults();
2236 if (opNumResults == 0)
2237 return success();
2238
2239 // If ForOp defines values, check that the number and types of the defined
2240 // values match ForOp initial iter operands and backedge basic block
2241 // arguments.
2242 if (getNumIterOperands() != opNumResults)
2243 return emitOpError(
2244 "mismatch between the number of loop-carried values and results");
2245 if (getNumRegionIterArgs() != opNumResults)
2246 return emitOpError(
2247 "mismatch between the number of basic block args and results");
2248
2249 return success();
2250}
2251
2252/// Parse a for operation loop bounds.
2253static ParseResult parseBound(bool isLower, OperationState &result,
2254 OpAsmParser &p) {
2255 // 'min' / 'max' prefixes are generally syntactic sugar, but are required if
2256 // the map has multiple results.
2257 bool failedToParsedMinMax =
2258 failed(p.parseOptionalKeyword(isLower ? "max" : "min"));
2259
2260 auto &builder = p.getBuilder();
2261 auto boundAttrStrName =
2262 isLower ? AffineForOp::getLowerBoundMapAttrName(result.name)
2263 : AffineForOp::getUpperBoundMapAttrName(result.name);
2264
2265 // Parse ssa-id as identity map.
2267 if (p.parseOperandList(boundOpInfos))
2268 return failure();
2269
2270 if (!boundOpInfos.empty()) {
2271 // Check that only one operand was parsed.
2272 if (boundOpInfos.size() > 1)
2273 return p.emitError(p.getNameLoc(),
2274 "expected only one loop bound operand");
2275
2276 // TODO: improve error message when SSA value is not of index type.
2277 // Currently it is 'use of value ... expects different type than prior uses'
2278 if (p.resolveOperand(boundOpInfos.front(), builder.getIndexType(),
2279 result.operands))
2280 return failure();
2281
2282 // Create an identity map using symbol id. This representation is optimized
2283 // for storage. Analysis passes may expand it into a multi-dimensional map
2284 // if desired.
2285 AffineMap map = builder.getSymbolIdentityMap();
2286 result.addAttribute(boundAttrStrName, AffineMapAttr::get(map));
2287 return success();
2288 }
2289
2290 // Get the attribute location.
2291 SMLoc attrLoc = p.getCurrentLocation();
2292
2293 Attribute boundAttr;
2294 if (p.parseAttribute(boundAttr, builder.getIndexType(), boundAttrStrName,
2295 result.attributes))
2296 return failure();
2297
2298 // Parse full form - affine map followed by dim and symbol list.
2299 if (auto affineMapAttr = dyn_cast<AffineMapAttr>(boundAttr)) {
2300 unsigned currentNumOperands = result.operands.size();
2301 unsigned numDims;
2302 if (parseDimAndSymbolList(p, result.operands, numDims))
2303 return failure();
2304
2305 auto map = affineMapAttr.getValue();
2306 if (map.getNumDims() != numDims)
2307 return p.emitError(
2308 p.getNameLoc(),
2309 "dim operand count and affine map dim count must match");
2310
2311 unsigned numDimAndSymbolOperands =
2312 result.operands.size() - currentNumOperands;
2313 if (numDims + map.getNumSymbols() != numDimAndSymbolOperands)
2314 return p.emitError(
2315 p.getNameLoc(),
2316 "symbol operand count and affine map symbol count must match");
2317
2318 // If the map has multiple results, make sure that we parsed the min/max
2319 // prefix.
2320 if (map.getNumResults() > 1 && failedToParsedMinMax) {
2321 if (isLower) {
2322 return p.emitError(attrLoc, "lower loop bound affine map with "
2323 "multiple results requires 'max' prefix");
2324 }
2325 return p.emitError(attrLoc, "upper loop bound affine map with multiple "
2326 "results requires 'min' prefix");
2327 }
2328 return success();
2329 }
2330
2331 // Parse custom assembly form.
2332 if (auto integerAttr = dyn_cast<IntegerAttr>(boundAttr)) {
2333 result.attributes.pop_back();
2334 result.addAttribute(
2335 boundAttrStrName,
2336 AffineMapAttr::get(builder.getConstantAffineMap(integerAttr.getInt())));
2337 return success();
2338 }
2339
2340 return p.emitError(
2341 p.getNameLoc(),
2342 "expected valid affine map representation for loop bounds");
2343}
2344
2345ParseResult AffineForOp::parse(OpAsmParser &parser, OperationState &result) {
2346 auto &builder = parser.getBuilder();
2347 OpAsmParser::Argument inductionVariable;
2348 inductionVariable.type = builder.getIndexType();
2349 // Parse the induction variable followed by '='.
2350 if (parser.parseArgument(inductionVariable) || parser.parseEqual())
2351 return failure();
2352
2353 // Parse loop bounds.
2354 int64_t numOperands = result.operands.size();
2355 if (parseBound(/*isLower=*/true, result, parser))
2356 return failure();
2357 int64_t numLbOperands = result.operands.size() - numOperands;
2358 if (parser.parseKeyword("to", " between bounds"))
2359 return failure();
2360 numOperands = result.operands.size();
2361 if (parseBound(/*isLower=*/false, result, parser))
2362 return failure();
2363 int64_t numUbOperands = result.operands.size() - numOperands;
2364
2365 // Parse the optional loop step, we default to 1 if one is not present.
2366 if (parser.parseOptionalKeyword("step")) {
2367 result.addAttribute(
2368 getStepAttrName(result.name),
2369 builder.getIntegerAttr(builder.getIndexType(), /*value=*/1));
2370 } else {
2371 SMLoc stepLoc = parser.getCurrentLocation();
2372 IntegerAttr stepAttr;
2373 if (parser.parseAttribute(stepAttr, builder.getIndexType(),
2374 getStepAttrName(result.name).data(),
2375 result.attributes))
2376 return failure();
2377
2378 if (!stepAttr.getValue().isStrictlyPositive())
2379 return parser.emitError(
2380 stepLoc,
2381 "expected step to be representable as a positive signed integer");
2382 }
2383
2384 // Parse the optional initial iteration arguments.
2385 SmallVector<OpAsmParser::Argument, 4> regionArgs;
2386 SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
2387
2388 // Induction variable.
2389 regionArgs.push_back(inductionVariable);
2390
2391 if (succeeded(parser.parseOptionalKeyword("iter_args"))) {
2392 // Parse assignment list and results type list.
2393 if (parser.parseAssignmentList(regionArgs, operands) ||
2394 parser.parseArrowTypeList(result.types))
2395 return failure();
2396 // Resolve input operands.
2397 for (auto argOperandType :
2398 llvm::zip(llvm::drop_begin(regionArgs), operands, result.types)) {
2399 Type type = std::get<2>(argOperandType);
2400 std::get<0>(argOperandType).type = type;
2401 if (parser.resolveOperand(std::get<1>(argOperandType), type,
2402 result.operands))
2403 return failure();
2404 }
2405 }
2406
2407 result.addAttribute(
2408 getOperandSegmentSizeAttr(),
2409 builder.getDenseI32ArrayAttr({static_cast<int32_t>(numLbOperands),
2410 static_cast<int32_t>(numUbOperands),
2411 static_cast<int32_t>(operands.size())}));
2412
2413 // Parse the body region.
2414 Region *body = result.addRegion();
2415 if (regionArgs.size() != result.types.size() + 1)
2416 return parser.emitError(
2417 parser.getNameLoc(),
2418 "mismatch between the number of loop-carried values and results");
2419 if (parser.parseRegion(*body, regionArgs))
2420 return failure();
2421
2422 AffineForOp::ensureTerminator(*body, builder, result.location);
2423
2424 // Parse the optional attribute list.
2425 return parser.parseOptionalAttrDict(result.attributes);
2426}
2427
2428static void printBound(AffineMapAttr boundMap,
2429 Operation::operand_range boundOperands,
2430 const char *prefix, OpAsmPrinter &p) {
2431 AffineMap map = boundMap.getValue();
2432
2433 // Check if this bound should be printed using custom assembly form.
2434 // The decision to restrict printing custom assembly form to trivial cases
2435 // comes from the will to roundtrip MLIR binary -> text -> binary in a
2436 // lossless way.
2437 // Therefore, custom assembly form parsing and printing is only supported for
2438 // zero-operand constant maps and single symbol operand identity maps.
2439 if (map.getNumResults() == 1) {
2440 AffineExpr expr = map.getResult(0);
2441
2442 // Print constant bound.
2443 if (map.getNumDims() == 0 && map.getNumSymbols() == 0) {
2444 if (auto constExpr = dyn_cast<AffineConstantExpr>(expr)) {
2445 p << constExpr.getValue();
2446 return;
2447 }
2448 }
2449
2450 // Print bound that consists of a single SSA symbol if the map is over a
2451 // single symbol.
2452 if (map.getNumDims() == 0 && map.getNumSymbols() == 1) {
2453 if (isa<AffineSymbolExpr>(expr)) {
2454 p.printOperand(*boundOperands.begin());
2455 return;
2456 }
2457 }
2458 } else {
2459 // Map has multiple results. Print 'min' or 'max' prefix.
2460 p << prefix << ' ';
2461 }
2462
2463 // Print the map and its operands.
2464 p << boundMap;
2465 printDimAndSymbolList(boundOperands.begin(), boundOperands.end(),
2466 map.getNumDims(), p);
2467}
2468
2469unsigned AffineForOp::getNumIterOperands() {
2470 AffineMap lbMap = getLowerBoundMapAttr().getValue();
2471 AffineMap ubMap = getUpperBoundMapAttr().getValue();
2472
2473 return getNumOperands() - lbMap.getNumInputs() - ubMap.getNumInputs();
2474}
2475
2476std::optional<MutableArrayRef<OpOperand>>
2477AffineForOp::getYieldedValuesMutable() {
2478 return cast<AffineYieldOp>(getBody()->getTerminator()).getOperandsMutable();
2479}
2480
2481void AffineForOp::print(OpAsmPrinter &p) {
2482 p << ' ';
2483 p.printRegionArgument(getBody()->getArgument(0), /*argAttrs=*/{},
2484 /*omitType=*/true);
2485 p << " = ";
2486 printBound(getLowerBoundMapAttr(), getLowerBoundOperands(), "max", p);
2487 p << " to ";
2488 printBound(getUpperBoundMapAttr(), getUpperBoundOperands(), "min", p);
2489
2490 if (getStepAsInt() != 1)
2491 p << " step " << getStepAsInt();
2492
2493 bool printBlockTerminators = false;
2494 if (getNumIterOperands() > 0) {
2495 p << " iter_args(";
2496 auto regionArgs = getRegionIterArgs();
2497 auto operands = getInits();
2498
2499 llvm::interleaveComma(llvm::zip(regionArgs, operands), p, [&](auto it) {
2500 p << std::get<0>(it) << " = " << std::get<1>(it);
2501 });
2502 p << ") -> (" << getResultTypes() << ")";
2503 printBlockTerminators = true;
2504 }
2505
2506 p << ' ';
2507 p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
2508 printBlockTerminators);
2510 (*this)->getAttrs(),
2511 /*elidedAttrs=*/{getLowerBoundMapAttrName(getOperation()->getName()),
2512 getUpperBoundMapAttrName(getOperation()->getName()),
2513 getStepAttrName(getOperation()->getName()),
2514 getOperandSegmentSizeAttr()});
2515}
2516
2517/// Fold the constant bounds of a loop.
2518static LogicalResult foldLoopBounds(AffineForOp forOp) {
2519 auto foldLowerOrUpperBound = [&forOp](bool lower) {
2520 // Check to see if each of the operands is the result of a constant. If
2521 // so, get the value. If not, ignore it.
2522 SmallVector<Attribute, 8> operandConstants;
2523 auto boundOperands =
2524 lower ? forOp.getLowerBoundOperands() : forOp.getUpperBoundOperands();
2525 for (auto operand : boundOperands) {
2526 Attribute operandCst;
2527 matchPattern(operand, m_Constant(&operandCst));
2528 operandConstants.push_back(operandCst);
2529 }
2530
2531 AffineMap boundMap =
2532 lower ? forOp.getLowerBoundMap() : forOp.getUpperBoundMap();
2533 assert(boundMap.getNumResults() >= 1 &&
2534 "bound maps should have at least one result");
2535 SmallVector<Attribute, 4> foldedResults;
2536 if (failed(boundMap.constantFold(operandConstants, foldedResults)))
2537 return failure();
2538
2539 // Compute the max or min as applicable over the results.
2540 assert(!foldedResults.empty() && "bounds should have at least one result");
2541 auto maxOrMin = llvm::cast<IntegerAttr>(foldedResults[0]).getValue();
2542 for (unsigned i = 1, e = foldedResults.size(); i < e; i++) {
2543 auto foldedResult = llvm::cast<IntegerAttr>(foldedResults[i]).getValue();
2544 maxOrMin = lower ? llvm::APIntOps::smax(maxOrMin, foldedResult)
2545 : llvm::APIntOps::smin(maxOrMin, foldedResult);
2546 }
2547 lower ? forOp.setConstantLowerBound(maxOrMin.getSExtValue())
2548 : forOp.setConstantUpperBound(maxOrMin.getSExtValue());
2549 return success();
2550 };
2551
2552 // Try to fold the lower bound.
2553 bool folded = false;
2554 if (!forOp.hasConstantLowerBound())
2555 folded |= succeeded(foldLowerOrUpperBound(/*lower=*/true));
2556
2557 // Try to fold the upper bound.
2558 if (!forOp.hasConstantUpperBound())
2559 folded |= succeeded(foldLowerOrUpperBound(/*lower=*/false));
2560 return success(folded);
2561}
2562
2563/// Returns constant trip count in trivial cases.
2564static std::optional<uint64_t> getTrivialConstantTripCount(AffineForOp forOp) {
2565 int64_t step = forOp.getStepAsInt();
2566 if (!forOp.hasConstantBounds() || step <= 0)
2567 return std::nullopt;
2568 int64_t lb = forOp.getConstantLowerBound();
2569 int64_t ub = forOp.getConstantUpperBound();
2570 return ub - lb <= 0 ? 0 : (ub - lb + step - 1) / step;
2571}
2572
2573/// Fold the empty loop.
2575 if (!llvm::hasSingleElement(*forOp.getBody()))
2576 return {};
2577 if (forOp.getNumResults() == 0)
2578 return {};
2579 std::optional<uint64_t> tripCount = getTrivialConstantTripCount(forOp);
2580 if (tripCount == 0) {
2581 // The initial values of the iteration arguments would be the op's
2582 // results.
2583 return forOp.getInits();
2584 }
2585 SmallVector<Value, 4> replacements;
2586 auto yieldOp = cast<AffineYieldOp>(forOp.getBody()->getTerminator());
2587 auto iterArgs = forOp.getRegionIterArgs();
2588 bool hasValDefinedOutsideLoop = false;
2589 bool iterArgsNotInOrder = false;
2590 for (unsigned i = 0, e = yieldOp->getNumOperands(); i < e; ++i) {
2591 Value val = yieldOp.getOperand(i);
2592 BlockArgument *iterArgIt = llvm::find(iterArgs, val);
2593 // TODO: It should be possible to perform a replacement by computing the
2594 // last value of the IV based on the bounds and the step.
2595 if (val == forOp.getInductionVar())
2596 return {};
2597 if (iterArgIt == iterArgs.end()) {
2598 // `val` is defined outside of the loop.
2599 assert(forOp.isDefinedOutsideOfLoop(val) &&
2600 "must be defined outside of the loop");
2601 hasValDefinedOutsideLoop = true;
2602 replacements.push_back(val);
2603 } else {
2604 unsigned pos = std::distance(iterArgs.begin(), iterArgIt);
2605 if (pos != i)
2606 iterArgsNotInOrder = true;
2607 replacements.push_back(forOp.getInits()[pos]);
2608 }
2609 }
2610 // Bail out when the trip count is unknown and the loop returns any value
2611 // defined outside of the loop or any iterArg out of order.
2612 if (!tripCount.has_value() &&
2613 (hasValDefinedOutsideLoop || iterArgsNotInOrder))
2614 return {};
2615 // Bail out when the loop iterates more than once and it returns any iterArg
2616 // out of order.
2617 if (tripCount.has_value() && tripCount.value() >= 2 && iterArgsNotInOrder)
2618 return {};
2619 return llvm::to_vector_of<OpFoldResult>(replacements);
2620}
2621
2622/// Canonicalize the bounds of the given loop.
2623static LogicalResult canonicalizeLoopBounds(AffineForOp forOp) {
2624 SmallVector<Value, 4> lbOperands(forOp.getLowerBoundOperands());
2625 SmallVector<Value, 4> ubOperands(forOp.getUpperBoundOperands());
2626
2627 auto lbMap = forOp.getLowerBoundMap();
2628 auto ubMap = forOp.getUpperBoundMap();
2629 auto prevLbMap = lbMap;
2630 auto prevUbMap = ubMap;
2631
2632 composeAffineMapAndOperands(&lbMap, &lbOperands);
2633 canonicalizeMapAndOperands(&lbMap, &lbOperands);
2634 simplifyMinOrMaxExprWithOperands(lbMap, lbOperands, /*isMax=*/true);
2635 simplifyMinOrMaxExprWithOperands(ubMap, ubOperands, /*isMax=*/false);
2636 lbMap = removeDuplicateExprs(lbMap);
2637
2638 composeAffineMapAndOperands(&ubMap, &ubOperands);
2639 canonicalizeMapAndOperands(&ubMap, &ubOperands);
2640 ubMap = removeDuplicateExprs(ubMap);
2641
2642 // Any canonicalization change always leads to updated map(s).
2643 if (lbMap == prevLbMap && ubMap == prevUbMap)
2644 return failure();
2645
2646 if (lbMap != prevLbMap)
2647 forOp.setLowerBound(lbOperands, lbMap);
2648 if (ubMap != prevUbMap)
2649 forOp.setUpperBound(ubOperands, ubMap);
2650 return success();
2651}
2652
2653/// Returns true if the affine.for has zero iterations in trivial cases.
2654static bool hasTrivialZeroTripCount(AffineForOp op) {
2655 return getTrivialConstantTripCount(op) == 0;
2656}
2657
2658LogicalResult AffineForOp::fold(FoldAdaptor adaptor,
2660 bool folded = succeeded(foldLoopBounds(*this));
2661 folded |= succeeded(canonicalizeLoopBounds(*this));
2662 if (hasTrivialZeroTripCount(*this) && getNumResults() != 0) {
2663 // The initial values of the loop-carried variables (iter_args) are the
2664 // results of the op. But this must be avoided for an affine.for op that
2665 // does not return any results. Since ops that do not return results cannot
2666 // be folded away, we would enter an infinite loop of folds on the same
2667 // affine.for op.
2668 results.assign(getInits().begin(), getInits().end());
2669 folded = true;
2670 }
2671 SmallVector<OpFoldResult> foldResults = AffineForEmptyLoopFolder(*this);
2672 if (!foldResults.empty()) {
2673 results.assign(foldResults);
2674 folded = true;
2675 }
2676 return success(folded);
2677}
2678
2679OperandRange AffineForOp::getEntrySuccessorOperands(RegionSuccessor successor) {
2680 assert(
2681 (successor.isOperation() || successor.getSuccessor() == &getRegion()) &&
2682 "invalid region point");
2683
2684 // The initial operands map to the loop arguments after the induction
2685 // variable or are forwarded to the results when the trip count is zero.
2686 return getInits();
2687}
2688
2689void AffineForOp::getSuccessorRegions(
2691 assert((point.isParent() ||
2692 point.getTerminatorPredecessorOrNull()->getParentRegion() ==
2693 &getRegion()) &&
2694 "expected loop region");
2695 // The loop may typically branch back to its body or to the parent operation.
2696 // If the predecessor is the parent op and the trip count is known to be at
2697 // least one, branch into the body using the iterator arguments. And in cases
2698 // we know the trip count is zero, it can only branch back to its parent.
2699 std::optional<uint64_t> tripCount = getTrivialConstantTripCount(*this);
2700 if (tripCount.has_value()) {
2701 if (!point.isParent()) {
2702 // From the loop body, if the trip count is one, we can only branch back
2703 // to the parent.
2704 if (tripCount == 1) {
2705 regions.push_back(RegionSuccessor(getOperation()));
2706 return;
2707 }
2708 if (tripCount == 0)
2709 return;
2710 } else {
2711 if (tripCount.value() > 0) {
2712 regions.push_back(RegionSuccessor(&getRegion()));
2713 return;
2714 }
2715 if (tripCount.value() == 0) {
2716 regions.push_back(RegionSuccessor(getOperation()));
2717 return;
2718 }
2719 }
2720 }
2721
2722 // In all other cases, the loop may branch back to itself or the parent
2723 // operation.
2724 regions.push_back(RegionSuccessor(&getRegion()));
2725 regions.push_back(RegionSuccessor(getOperation()));
2726}
2727
2728ValueRange AffineForOp::getSuccessorInputs(RegionSuccessor successor) {
2729 if (successor.isOperation())
2730 return getResults();
2731 return getRegionIterArgs();
2732}
2733
2734AffineBound AffineForOp::getLowerBound() {
2735 return AffineBound(*this, getLowerBoundOperands(), getLowerBoundMap());
2736}
2737
2738AffineBound AffineForOp::getUpperBound() {
2739 return AffineBound(*this, getUpperBoundOperands(), getUpperBoundMap());
2740}
2741
2742void AffineForOp::setLowerBound(ValueRange lbOperands, AffineMap map) {
2743 assert(lbOperands.size() == map.getNumInputs());
2744 assert(map.getNumResults() >= 1 && "bound map has at least one result");
2745 getLowerBoundOperandsMutable().assign(lbOperands);
2746 setLowerBoundMap(map);
2747}
2748
2749void AffineForOp::setUpperBound(ValueRange ubOperands, AffineMap map) {
2750 assert(ubOperands.size() == map.getNumInputs());
2751 assert(map.getNumResults() >= 1 && "bound map has at least one result");
2752 getUpperBoundOperandsMutable().assign(ubOperands);
2753 setUpperBoundMap(map);
2754}
2755
2756bool AffineForOp::hasConstantLowerBound() {
2757 return getLowerBoundMap().isSingleConstant();
2758}
2759
2760bool AffineForOp::hasConstantUpperBound() {
2761 return getUpperBoundMap().isSingleConstant();
2762}
2763
2764int64_t AffineForOp::getConstantLowerBound() {
2765 return getLowerBoundMap().getSingleConstantResult();
2766}
2767
2768int64_t AffineForOp::getConstantUpperBound() {
2769 return getUpperBoundMap().getSingleConstantResult();
2770}
2771
2772void AffineForOp::setConstantLowerBound(int64_t value) {
2773 setLowerBound({}, AffineMap::getConstantMap(value, getContext()));
2774}
2775
2776void AffineForOp::setConstantUpperBound(int64_t value) {
2777 setUpperBound({}, AffineMap::getConstantMap(value, getContext()));
2778}
2779
2780AffineForOp::operand_range AffineForOp::getControlOperands() {
2781 return {operand_begin(), operand_begin() + getLowerBoundOperands().size() +
2782 getUpperBoundOperands().size()};
2783}
2784
2785bool AffineForOp::matchingBoundOperandList() {
2786 auto lbMap = getLowerBoundMap();
2787 auto ubMap = getUpperBoundMap();
2788 if (lbMap.getNumDims() != ubMap.getNumDims() ||
2789 lbMap.getNumSymbols() != ubMap.getNumSymbols())
2790 return false;
2791
2792 unsigned numOperands = lbMap.getNumInputs();
2793 for (unsigned i = 0, e = lbMap.getNumInputs(); i < e; i++) {
2794 // Compare Value 's.
2795 if (getOperand(i) != getOperand(numOperands + i))
2796 return false;
2797 }
2798 return true;
2799}
2800
2801SmallVector<Region *> AffineForOp::getLoopRegions() { return {&getRegion()}; }
2802
2803std::optional<SmallVector<Value>> AffineForOp::getLoopInductionVars() {
2804 return SmallVector<Value>{getInductionVar()};
2805}
2806
2807std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopLowerBounds() {
2808 if (!hasConstantLowerBound())
2809 return std::nullopt;
2810 OpBuilder b(getContext());
2811 return SmallVector<OpFoldResult>{
2812 OpFoldResult(b.getI64IntegerAttr(getConstantLowerBound()))};
2813}
2814
2815std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopSteps() {
2816 OpBuilder b(getContext());
2817 return SmallVector<OpFoldResult>{
2818 OpFoldResult(b.getI64IntegerAttr(getStepAsInt()))};
2819}
2820
2821std::optional<SmallVector<OpFoldResult>> AffineForOp::getLoopUpperBounds() {
2822 if (!hasConstantUpperBound())
2823 return {};
2824 OpBuilder b(getContext());
2825 return SmallVector<OpFoldResult>{
2826 OpFoldResult(b.getI64IntegerAttr(getConstantUpperBound()))};
2827}
2828
2829std::optional<APInt> AffineForOp::getStaticTripCount() {
2830 MLIRContext *context = getContext();
2831 int64_t step = getStepAsInt();
2832 if (step <= 0)
2833 return std::nullopt;
2834
2835 if (hasConstantBounds()) {
2836 int64_t lb = getConstantLowerBound();
2837 int64_t ub = getConstantUpperBound();
2838 int64_t loopSpan = ub - lb;
2839 if (loopSpan < 0)
2840 loopSpan = 0;
2841 return APInt(64, llvm::divideCeilSigned(loopSpan, step));
2842 }
2843
2844 auto lbMap = getLowerBoundMap();
2845 auto ubMap = getUpperBoundMap();
2846 if (lbMap.getNumResults() != 1)
2847 return std::nullopt;
2848
2849 // Difference of each upper bound expression from the single lower bound
2850 // expression (divided by the step) provides the expressions for the trip
2851 // count map.
2852 AffineValueMap ubValueMap(ubMap, getUpperBoundOperands());
2853
2854 SmallVector<AffineExpr, 4> lbSplatExpr(ubValueMap.getNumResults(),
2855 lbMap.getResult(0));
2856 auto lbMapSplat = AffineMap::get(lbMap.getNumDims(), lbMap.getNumSymbols(),
2857 lbSplatExpr, context);
2858 AffineValueMap lbSplatValueMap(lbMapSplat, getLowerBoundOperands());
2859
2860 AffineValueMap tripCountValueMap;
2861 AffineValueMap::difference(ubValueMap, lbSplatValueMap, &tripCountValueMap);
2862
2863 // Take the min if all trip counts are constant.
2864 std::optional<uint64_t> tripCount;
2865 for (unsigned i = 0, e = tripCountValueMap.getNumResults(); i < e; ++i) {
2866 AffineExpr expr = tripCountValueMap.getResult(i).ceilDiv(step);
2867 if (auto constExpr = llvm::dyn_cast<AffineConstantExpr>(expr)) {
2868 uint64_t value = constExpr.getValue();
2869 if (tripCount.has_value())
2870 tripCount = std::min(*tripCount, value);
2871 else
2872 tripCount = value;
2873 } else {
2874 return std::nullopt;
2875 }
2876 }
2877
2878 if (tripCount.has_value())
2879 return APInt(64, *tripCount);
2880
2881 return std::nullopt;
2882}
2883
2884FailureOr<LoopLikeOpInterface> AffineForOp::replaceWithAdditionalYields(
2885 RewriterBase &rewriter, ValueRange newInitOperands,
2886 bool replaceInitOperandUsesInLoop,
2887 const NewYieldValuesFn &newYieldValuesFn) {
2888 // Create a new loop before the existing one, with the extra operands.
2889 OpBuilder::InsertionGuard g(rewriter);
2890 rewriter.setInsertionPoint(getOperation());
2891 auto inits = llvm::to_vector(getInits());
2892 inits.append(newInitOperands.begin(), newInitOperands.end());
2893 AffineForOp newLoop = AffineForOp::create(
2894 rewriter, getLoc(), getLowerBoundOperands(), getLowerBoundMap(),
2895 getUpperBoundOperands(), getUpperBoundMap(), getStepAsInt(), inits);
2896 // Existing operands, results, and region arguments retain their positions;
2897 // only new loop-carried values are appended.
2898 newLoop->setDiscardableAttrs(getOperation()->getDiscardableAttrDictionary());
2899
2900 // Generate the new yield values and append them to the scf.yield operation.
2901 auto yieldOp = cast<AffineYieldOp>(getBody()->getTerminator());
2902 ArrayRef<BlockArgument> newIterArgs =
2903 newLoop.getBody()->getArguments().take_back(newInitOperands.size());
2904 {
2905 OpBuilder::InsertionGuard g(rewriter);
2906 rewriter.setInsertionPoint(yieldOp);
2907 SmallVector<Value> newYieldedValues =
2908 newYieldValuesFn(rewriter, getLoc(), newIterArgs);
2909 assert(newInitOperands.size() == newYieldedValues.size() &&
2910 "expected as many new yield values as new iter operands");
2911 rewriter.modifyOpInPlace(yieldOp, [&]() {
2912 yieldOp.getOperandsMutable().append(newYieldedValues);
2913 });
2914 }
2915
2916 // Move the loop body to the new op.
2917 rewriter.mergeBlocks(getBody(), newLoop.getBody(),
2918 newLoop.getBody()->getArguments().take_front(
2919 getBody()->getNumArguments()));
2920
2921 if (replaceInitOperandUsesInLoop) {
2922 // Replace all uses of `newInitOperands` with the corresponding basic block
2923 // arguments.
2924 for (auto it : llvm::zip(newInitOperands, newIterArgs)) {
2925 rewriter.replaceUsesWithIf(std::get<0>(it), std::get<1>(it),
2926 [&](OpOperand &use) {
2927 Operation *user = use.getOwner();
2928 return newLoop->isProperAncestor(user);
2929 });
2930 }
2931 }
2932
2933 // Replace the old loop.
2934 rewriter.replaceOp(getOperation(),
2935 newLoop->getResults().take_front(getNumResults()));
2936 return cast<LoopLikeOpInterface>(newLoop.getOperation());
2937}
2938
2939Speculation::Speculatability AffineForOp::getSpeculatability() {
2940 // `affine.for (I = Start; I < End; I += 1)` terminates for all values of
2941 // Start and End.
2942 //
2943 // For Step != 1, the loop may not terminate. We can add more smarts here if
2944 // needed.
2945 return getStepAsInt() == 1 ? Speculation::RecursivelySpeculatable
2947}
2948
2949/// Returns true if the provided value is the induction variable of a
2950/// AffineForOp.
2952 return getForInductionVarOwner(val) != AffineForOp();
2953}
2954
2958
2962
2964 auto ivArg = dyn_cast<BlockArgument>(val);
2965 if (!ivArg || !ivArg.getOwner() || !ivArg.getOwner()->getParent())
2966 return AffineForOp();
2967 if (auto forOp =
2968 ivArg.getOwner()->getParent()->getParentOfType<AffineForOp>())
2969 // Check to make sure `val` is the induction variable, not an iter_arg.
2970 return forOp.getInductionVar() == val ? forOp : AffineForOp();
2971 return AffineForOp();
2972}
2973
2975 auto ivArg = dyn_cast<BlockArgument>(val);
2976 if (!ivArg || !ivArg.getOwner())
2977 return nullptr;
2978 Operation *containingOp = ivArg.getOwner()->getParentOp();
2979 auto parallelOp = dyn_cast_if_present<AffineParallelOp>(containingOp);
2980 if (parallelOp && llvm::is_contained(parallelOp.getIVs(), val))
2981 return parallelOp;
2982 return nullptr;
2983}
2984
2985/// Extracts the induction variables from a list of AffineForOps and returns
2986/// them.
2989 ivs->reserve(forInsts.size());
2990 for (auto forInst : forInsts)
2991 ivs->push_back(forInst.getInductionVar());
2992}
2993
2996 ivs.reserve(affineOps.size());
2997 for (Operation *op : affineOps) {
2998 // Add constraints from forOp's bounds.
2999 if (auto forOp = dyn_cast<AffineForOp>(op))
3000 ivs.push_back(forOp.getInductionVar());
3001 else if (auto parallelOp = dyn_cast<AffineParallelOp>(op))
3002 for (size_t i = 0; i < parallelOp.getBody()->getNumArguments(); i++)
3003 ivs.push_back(parallelOp.getBody()->getArgument(i));
3004 }
3005}
3006
3007/// Builds an affine loop nest, using "loopCreatorFn" to create individual loop
3008/// operations.
3009template <typename BoundListTy, typename LoopCreatorTy>
3011 OpBuilder &builder, Location loc, BoundListTy lbs, BoundListTy ubs,
3012 ArrayRef<int64_t> steps,
3013 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn,
3014 LoopCreatorTy &&loopCreatorFn) {
3015 assert(lbs.size() == ubs.size() && "Mismatch in number of arguments");
3016 assert(lbs.size() == steps.size() && "Mismatch in number of arguments");
3017
3018 // If there are no loops to be constructed, construct the body anyway.
3019 OpBuilder::InsertionGuard guard(builder);
3020 if (lbs.empty()) {
3021 if (bodyBuilderFn)
3022 bodyBuilderFn(builder, loc, ValueRange());
3023 return;
3024 }
3025
3026 // Create the loops iteratively and store the induction variables.
3028 ivs.reserve(lbs.size());
3029 for (unsigned i = 0, e = lbs.size(); i < e; ++i) {
3030 // Callback for creating the loop body, always creates the terminator.
3031 auto loopBody = [&](OpBuilder &nestedBuilder, Location nestedLoc, Value iv,
3032 ValueRange iterArgs) {
3033 ivs.push_back(iv);
3034 // In the innermost loop, call the body builder.
3035 if (i == e - 1 && bodyBuilderFn) {
3036 OpBuilder::InsertionGuard nestedGuard(nestedBuilder);
3037 bodyBuilderFn(nestedBuilder, nestedLoc, ivs);
3038 }
3039 AffineYieldOp::create(nestedBuilder, nestedLoc);
3040 };
3041
3042 // Delegate actual loop creation to the callback in order to dispatch
3043 // between constant- and variable-bound loops.
3044 auto loop = loopCreatorFn(builder, loc, lbs[i], ubs[i], steps[i], loopBody);
3045 builder.setInsertionPointToStart(loop.getBody());
3046 }
3047}
3048
3049/// Creates an affine loop from the bounds known to be constants.
3050static AffineForOp
3052 int64_t ub, int64_t step,
3053 AffineForOp::BodyBuilderFn bodyBuilderFn) {
3054 return AffineForOp::create(builder, loc, lb, ub, step,
3055 /*iterArgs=*/ValueRange(), bodyBuilderFn);
3056}
3057
3058/// Creates an affine loop from the bounds that may or may not be constants.
3059static AffineForOp
3061 int64_t step,
3062 AffineForOp::BodyBuilderFn bodyBuilderFn) {
3063 std::optional<int64_t> lbConst = getConstantIntValue(lb);
3064 std::optional<int64_t> ubConst = getConstantIntValue(ub);
3065 if (lbConst && ubConst)
3066 return buildAffineLoopFromConstants(builder, loc, lbConst.value(),
3067 ubConst.value(), step, bodyBuilderFn);
3068 return AffineForOp::create(builder, loc, lb, builder.getDimIdentityMap(), ub,
3069 builder.getDimIdentityMap(), step,
3070 /*iterArgs=*/ValueRange(), bodyBuilderFn);
3071}
3072
3074 OpBuilder &builder, Location loc, ArrayRef<int64_t> lbs,
3076 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn) {
3077 buildAffineLoopNestImpl(builder, loc, lbs, ubs, steps, bodyBuilderFn,
3079}
3080
3082 OpBuilder &builder, Location loc, ValueRange lbs, ValueRange ubs,
3083 ArrayRef<int64_t> steps,
3084 function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn) {
3085 buildAffineLoopNestImpl(builder, loc, lbs, ubs, steps, bodyBuilderFn,
3087}
3088
3089//===----------------------------------------------------------------------===//
3090// AffineIfOp
3091//===----------------------------------------------------------------------===//
3092
3093namespace {
3094/// Remove else blocks that have nothing other than a zero value yield.
3095struct SimplifyDeadElse : public OpRewritePattern<AffineIfOp> {
3096 using OpRewritePattern<AffineIfOp>::OpRewritePattern;
3097
3098 LogicalResult matchAndRewrite(AffineIfOp ifOp,
3099 PatternRewriter &rewriter) const override {
3100 if (ifOp.getElseRegion().empty() ||
3101 !llvm::hasSingleElement(*ifOp.getElseBlock()) || ifOp.getNumResults())
3102 return failure();
3103
3104 rewriter.startOpModification(ifOp);
3105 rewriter.eraseBlock(ifOp.getElseBlock());
3106 rewriter.finalizeOpModification(ifOp);
3107 return success();
3108 }
3109};
3110
3111/// Removes affine.if cond if the condition is always true or false in certain
3112/// trivial cases. Promotes the then/else block in the parent operation block.
3113struct AlwaysTrueOrFalseIf : public OpRewritePattern<AffineIfOp> {
3114 using OpRewritePattern<AffineIfOp>::OpRewritePattern;
3115
3116 LogicalResult matchAndRewrite(AffineIfOp op,
3117 PatternRewriter &rewriter) const override {
3118
3119 auto isTriviallyFalse = [](IntegerSet iSet) {
3120 return iSet.isEmptyIntegerSet();
3121 };
3122
3123 auto isTriviallyTrue = [](IntegerSet iSet) {
3124 return (iSet.getNumEqualities() == 1 && iSet.getNumInequalities() == 0 &&
3125 iSet.getConstraint(0) == 0);
3126 };
3127
3128 IntegerSet affineIfConditions = op.getIntegerSet();
3129 Block *blockToMove;
3130 if (isTriviallyFalse(affineIfConditions)) {
3131 // The absence, or equivalently, the emptiness of the else region need not
3132 // be checked when affine.if is returning results because if an affine.if
3133 // operation is returning results, it always has a non-empty else region.
3134 if (op.getNumResults() == 0 && !op.hasElse()) {
3135 // If the else region is absent, or equivalently, empty, remove the
3136 // affine.if operation (which is not returning any results).
3137 rewriter.eraseOp(op);
3138 return success();
3139 }
3140 blockToMove = op.getElseBlock();
3141 } else if (isTriviallyTrue(affineIfConditions)) {
3142 blockToMove = op.getThenBlock();
3143 } else {
3144 return failure();
3145 }
3146 Operation *blockToMoveTerminator = blockToMove->getTerminator();
3147 // Promote the "blockToMove" block to the parent operation block between the
3148 // prologue and epilogue of "op".
3149 rewriter.inlineBlockBefore(blockToMove, op);
3150 // Replace the "op" operation with the operands of the
3151 // "blockToMoveTerminator" operation. Note that "blockToMoveTerminator" is
3152 // the affine.yield operation present in the "blockToMove" block. It has no
3153 // operands when affine.if is not returning results and therefore, in that
3154 // case, replaceOp just erases "op". When affine.if is not returning
3155 // results, the affine.yield operation can be omitted. It gets inserted
3156 // implicitly.
3157 rewriter.replaceOp(op, blockToMoveTerminator->getOperands());
3158 // Erase the "blockToMoveTerminator" operation since it is now in the parent
3159 // operation block, which already has its own terminator.
3160 rewriter.eraseOp(blockToMoveTerminator);
3161 return success();
3162 }
3163};
3164} // namespace
3165
3166/// AffineIfOp has two regions -- `then` and `else`. The flow of data should be
3167/// as follows: AffineIfOp -> `then`/`else` -> AffineIfOp
3168void AffineIfOp::getSuccessorRegions(
3170 // If the predecessor is an AffineIfOp, then branching into both `then` and
3171 // `else` region is valid.
3172 if (point.isParent()) {
3173 regions.reserve(2);
3174 regions.push_back(RegionSuccessor(&getThenRegion()));
3175 // If the "else" region is empty, branch bach into parent.
3176 if (getElseRegion().empty()) {
3177 regions.push_back(RegionSuccessor(getOperation()));
3178 } else {
3179 regions.push_back(RegionSuccessor(&getElseRegion()));
3180 }
3181 return;
3182 }
3183
3184 // If the predecessor is the `else`/`then` region, then branching into parent
3185 // op is valid.
3186 regions.push_back(RegionSuccessor(getOperation()));
3187}
3188
3189ValueRange AffineIfOp::getSuccessorInputs(RegionSuccessor successor) {
3190 if (successor.isOperation())
3191 return getResults();
3192 if (successor == &getThenRegion())
3193 return getThenRegion().getArguments();
3194 if (successor == &getElseRegion())
3195 return getElseRegion().getArguments();
3196 llvm_unreachable("invalid region successor");
3197}
3198
3199LogicalResult AffineIfOp::verify() {
3200 // Verify that we have a condition attribute.
3201 // FIXME: This should be specified in the arguments list in ODS.
3202 auto conditionAttr =
3203 (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
3204 if (!conditionAttr)
3205 return emitOpError("requires an integer set attribute named 'condition'");
3206
3207 // Verify that there are enough operands for the condition.
3208 IntegerSet condition = conditionAttr.getValue();
3209 if (getNumOperands() != condition.getNumInputs())
3210 return emitOpError("operand count and condition integer set dimension and "
3211 "symbol count must match");
3212
3213 // Verify that the operands are valid dimension/symbols.
3214 if (failed(verifyDimAndSymbolIdentifiers(*this, getOperands(),
3215 condition.getNumDims())))
3216 return failure();
3217
3218 return success();
3219}
3220
3221ParseResult AffineIfOp::parse(OpAsmParser &parser, OperationState &result) {
3222 // Parse the condition attribute set.
3223 IntegerSetAttr conditionAttr;
3224 unsigned numDims;
3225 if (parser.parseAttribute(conditionAttr,
3226 AffineIfOp::getConditionAttrStrName(),
3227 result.attributes) ||
3228 parseDimAndSymbolList(parser, result.operands, numDims))
3229 return failure();
3230
3231 // Verify the condition operands.
3232 auto set = conditionAttr.getValue();
3233 if (set.getNumDims() != numDims)
3234 return parser.emitError(
3235 parser.getNameLoc(),
3236 "dim operand count and integer set dim count must match");
3237 if (numDims + set.getNumSymbols() != result.operands.size())
3238 return parser.emitError(
3239 parser.getNameLoc(),
3240 "symbol operand count and integer set symbol count must match");
3241
3242 if (parser.parseOptionalArrowTypeList(result.types))
3243 return failure();
3244
3245 // Create the regions for 'then' and 'else'. The latter must be created even
3246 // if it remains empty for the validity of the operation.
3247 result.regions.reserve(2);
3248 Region *thenRegion = result.addRegion();
3249 Region *elseRegion = result.addRegion();
3250
3251 // Parse the 'then' region.
3252 if (parser.parseRegion(*thenRegion, {}, {}))
3253 return failure();
3254 AffineIfOp::ensureTerminator(*thenRegion, parser.getBuilder(),
3255 result.location);
3256
3257 // If we find an 'else' keyword then parse the 'else' region.
3258 if (!parser.parseOptionalKeyword("else")) {
3259 if (parser.parseRegion(*elseRegion, {}, {}))
3260 return failure();
3261 AffineIfOp::ensureTerminator(*elseRegion, parser.getBuilder(),
3262 result.location);
3263 }
3264
3265 // Parse the optional attribute list.
3266 if (parser.parseOptionalAttrDict(result.attributes))
3267 return failure();
3268
3269 return success();
3270}
3271
3272void AffineIfOp::print(OpAsmPrinter &p) {
3273 auto conditionAttr =
3274 (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
3275 p << " " << conditionAttr;
3276 printDimAndSymbolList(operand_begin(), operand_end(),
3277 conditionAttr.getValue().getNumDims(), p);
3278 p.printOptionalArrowTypeList(getResultTypes());
3279 p << ' ';
3280 p.printRegion(getThenRegion(), /*printEntryBlockArgs=*/false,
3281 /*printBlockTerminators=*/getNumResults());
3282
3283 // Print the 'else' regions if it has any blocks.
3284 auto &elseRegion = this->getElseRegion();
3285 if (!elseRegion.empty()) {
3286 p << " else ";
3287 p.printRegion(elseRegion,
3288 /*printEntryBlockArgs=*/false,
3289 /*printBlockTerminators=*/getNumResults());
3290 }
3291
3292 // Print the attribute list.
3293 p.printOptionalAttrDict((*this)->getAttrs(),
3294 /*elidedAttrs=*/getConditionAttrStrName());
3295}
3296
3297IntegerSet AffineIfOp::getIntegerSet() {
3298 return (*this)
3299 ->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName())
3300 .getValue();
3301}
3302
3303void AffineIfOp::setIntegerSet(IntegerSet newSet) {
3304 (*this)->setAttr(getConditionAttrStrName(), IntegerSetAttr::get(newSet));
3305}
3306
3307void AffineIfOp::setConditional(IntegerSet set, ValueRange operands) {
3308 setIntegerSet(set);
3309 (*this)->setOperands(operands);
3310}
3311
3312void AffineIfOp::build(OpBuilder &builder, OperationState &result,
3313 TypeRange resultTypes, IntegerSet set, ValueRange args,
3314 bool withElseRegion) {
3315 assert(resultTypes.empty() || withElseRegion);
3316 OpBuilder::InsertionGuard guard(builder);
3317
3318 result.addTypes(resultTypes);
3319 result.addOperands(args);
3320 result.addAttribute(getConditionAttrStrName(), IntegerSetAttr::get(set));
3321
3322 Region *thenRegion = result.addRegion();
3323 builder.createBlock(thenRegion);
3324 if (resultTypes.empty())
3325 AffineIfOp::ensureTerminator(*thenRegion, builder, result.location);
3326
3327 Region *elseRegion = result.addRegion();
3328 if (withElseRegion) {
3329 builder.createBlock(elseRegion);
3330 if (resultTypes.empty())
3331 AffineIfOp::ensureTerminator(*elseRegion, builder, result.location);
3332 }
3333}
3334
3335void AffineIfOp::build(OpBuilder &builder, OperationState &result,
3336 IntegerSet set, ValueRange args, bool withElseRegion) {
3337 AffineIfOp::build(builder, result, /*resultTypes=*/{}, set, args,
3338 withElseRegion);
3339}
3340
3341/// Compose any affine.apply ops feeding into `operands` of the integer set
3342/// `set` by composing the maps of such affine.apply ops with the integer
3343/// set constraints.
3345 SmallVectorImpl<Value> &operands,
3346 bool composeAffineMin = false) {
3347 // We will simply reuse the API of the map composition by viewing the LHSs of
3348 // the equalities and inequalities of `set` as the affine exprs of an affine
3349 // map. Convert to equivalent map, compose, and convert back to set.
3350 auto map = AffineMap::get(set.getNumDims(), set.getNumSymbols(),
3351 set.getConstraints(), set.getContext());
3352 // Check if any composition is possible.
3353 if (llvm::none_of(operands,
3354 [](Value v) { return v.getDefiningOp<AffineApplyOp>(); }))
3355 return;
3356
3357 composeAffineMapAndOperands(&map, &operands, composeAffineMin);
3358 set = IntegerSet::get(map.getNumDims(), map.getNumSymbols(), map.getResults(),
3359 set.getEqFlags());
3360}
3361
3362/// Canonicalize an affine if op's conditional (integer set + operands).
3363LogicalResult AffineIfOp::fold(FoldAdaptor, SmallVectorImpl<OpFoldResult> &) {
3364 auto set = getIntegerSet();
3365 SmallVector<Value, 4> operands(getOperands());
3366 composeSetAndOperands(set, operands);
3367 canonicalizeSetAndOperands(&set, &operands);
3368
3369 // Check if the canonicalization or composition led to any change.
3370 if (getIntegerSet() == set && llvm::equal(operands, getOperands()))
3371 return failure();
3372
3373 setConditional(set, operands);
3374 return success();
3375}
3376
3377void AffineIfOp::getCanonicalizationPatterns(RewritePatternSet &results,
3378 MLIRContext *context) {
3379 results.add<SimplifyDeadElse, AlwaysTrueOrFalseIf>(context);
3380}
3381
3382/// Adds the optional `alignment` attribute to `result`, if one is given.
3384 StringAttr attrName, llvm::MaybeAlign alignment) {
3385 if (alignment)
3386 result.addAttribute(attrName,
3387 builder.getI64IntegerAttr(alignment->value()));
3388}
3389
3390//===----------------------------------------------------------------------===//
3391// AffineLoadOp
3392//===----------------------------------------------------------------------===//
3393
3394void AffineLoadOp::build(OpBuilder &builder, OperationState &result,
3395 AffineMap map, ValueRange operands,
3396 llvm::MaybeAlign alignment) {
3397 assert(operands.size() == 1 + map.getNumInputs() && "inconsistent operands");
3398 result.addOperands(operands);
3399 if (map)
3400 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3401 addAlignmentAttr(builder, result, getAlignmentAttrName(result.name),
3402 alignment);
3403 auto memrefType = llvm::cast<MemRefType>(operands[0].getType());
3404 result.types.push_back(memrefType.getElementType());
3405}
3406
3407void AffineLoadOp::build(OpBuilder &builder, OperationState &result,
3408 Value memref, AffineMap map, ValueRange mapOperands,
3409 llvm::MaybeAlign alignment) {
3410 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
3411 result.addOperands(memref);
3412 result.addOperands(mapOperands);
3413 auto memrefType = llvm::cast<MemRefType>(memref.getType());
3414 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3415 addAlignmentAttr(builder, result, getAlignmentAttrName(result.name),
3416 alignment);
3417 result.types.push_back(memrefType.getElementType());
3418}
3419
3420void AffineLoadOp::build(OpBuilder &builder, OperationState &result,
3422 llvm::MaybeAlign alignment) {
3423 auto memrefType = llvm::cast<MemRefType>(memref.getType());
3424 int64_t rank = memrefType.getRank();
3425 // Create identity map for memrefs with at least one dimension or () -> ()
3426 // for zero-dimensional memrefs.
3427 auto map =
3428 rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap();
3429 build(builder, result, memref, map, indices, alignment);
3430}
3431
3432ParseResult AffineLoadOp::parse(OpAsmParser &parser, OperationState &result) {
3433 auto &builder = parser.getBuilder();
3434 auto indexTy = builder.getIndexType();
3435
3436 MemRefType type;
3438 AffineMapAttr mapAttr;
3440 return failure(
3441 parser.parseOperand(memrefInfo) ||
3442 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
3443 AffineLoadOp::getMapAttrStrName(),
3444 result.attributes) ||
3445 parser.parseOptionalAttrDict(result.attributes) ||
3446 parser.parseColonType(type) ||
3447 parser.resolveOperand(memrefInfo, type, result.operands) ||
3448 parser.resolveOperands(mapOperands, indexTy, result.operands) ||
3449 parser.addTypeToList(type.getElementType(), result.types));
3450}
3451
3452void AffineLoadOp::print(OpAsmPrinter &p) {
3453 p << " " << getMemRef() << '[';
3454 if (AffineMapAttr mapAttr =
3455 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
3456 p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
3457 p << ']';
3458 p.printOptionalAttrDict((*this)->getAttrs(),
3459 /*elidedAttrs=*/{getMapAttrStrName()});
3460 p << " : " << getMemRefType();
3461}
3462
3463/// Verify common indexing invariants of affine.load, affine.store,
3464/// affine.vector_load and affine.vector_store.
3465template <typename AffineMemOpTy>
3466static LogicalResult
3467verifyMemoryOpIndexing(AffineMemOpTy op, AffineMapAttr mapAttr,
3468 Operation::operand_range mapOperands,
3469 MemRefType memrefType, unsigned numIndexOperands) {
3470 AffineMap map = mapAttr.getValue();
3471 if (map.getNumResults() != memrefType.getRank())
3472 return op->emitOpError("affine map num results must equal memref rank");
3473 if (map.getNumInputs() != numIndexOperands)
3474 return op->emitOpError("expects as many subscripts as affine map inputs");
3475
3476 for (auto idx : mapOperands) {
3477 if (!idx.getType().isIndex())
3478 return op->emitOpError("index to load must have 'index' type");
3479 }
3480 if (failed(verifyDimAndSymbolIdentifiers(op, mapOperands, map.getNumDims())))
3481 return failure();
3482
3483 return success();
3484}
3485
3486LogicalResult AffineLoadOp::verify() {
3487 auto memrefType = getMemRefType();
3488 if (getType() != memrefType.getElementType())
3489 return emitOpError("result type must match element type of memref");
3490
3492 *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
3493 getMapOperands(), memrefType,
3494 /*numIndexOperands=*/getNumOperands() - 1)))
3495 return failure();
3496
3497 return success();
3498}
3499
3500void AffineLoadOp::getCanonicalizationPatterns(RewritePatternSet &results,
3501 MLIRContext *context) {
3502 results.add<SimplifyAffineOp<AffineLoadOp>>(context);
3503}
3504
3505OpFoldResult AffineLoadOp::fold(FoldAdaptor adaptor) {
3506 /// load(memrefcast) -> load
3507 if (succeeded(memref::foldMemRefCast(*this)))
3508 return getResult();
3509
3510 // Fold load from a global constant memref.
3511 auto getGlobalOp = getMemref().getDefiningOp<memref::GetGlobalOp>();
3512 if (!getGlobalOp)
3513 return {};
3514 // Get to the memref.global defining the symbol.
3516 getGlobalOp, getGlobalOp.getNameAttr());
3517 if (!global)
3518 return {};
3519
3520 // Check if the global memref is a constant.
3521 auto cstAttr =
3522 dyn_cast_or_null<DenseElementsAttr>(global.getConstantInitValue());
3523 if (!cstAttr)
3524 return {};
3525 // If it's a splat constant, we can fold irrespective of indices.
3526 if (auto splatAttr = dyn_cast<SplatElementsAttr>(cstAttr))
3527 return splatAttr.getSplatValue<Attribute>();
3528 // Otherwise, we can fold only if we know the indices.
3529 if (!getAffineMap().isConstant())
3530 return {};
3531 auto indices =
3532 llvm::map_to_vector<4>(getAffineMap().getConstantResults(),
3533 [](int64_t v) -> uint64_t { return v; });
3534 return cstAttr.getValues<Attribute>()[indices];
3535}
3536
3537//===----------------------------------------------------------------------===//
3538// AffineStoreOp
3539//===----------------------------------------------------------------------===//
3540
3541void AffineStoreOp::build(OpBuilder &builder, OperationState &result,
3542 Value valueToStore, Value memref, AffineMap map,
3543 ValueRange mapOperands, llvm::MaybeAlign alignment) {
3544 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
3545 result.addOperands(valueToStore);
3546 result.addOperands(memref);
3547 result.addOperands(mapOperands);
3548 result.getOrAddProperties<Properties>().map = AffineMapAttr::get(map);
3549 addAlignmentAttr(builder, result, getAlignmentAttrName(result.name),
3550 alignment);
3551}
3552
3553// Use identity map.
3554void AffineStoreOp::build(OpBuilder &builder, OperationState &result,
3555 Value valueToStore, Value memref, ValueRange indices,
3556 llvm::MaybeAlign alignment) {
3557 auto memrefType = llvm::cast<MemRefType>(memref.getType());
3558 int64_t rank = memrefType.getRank();
3559 // Create identity map for memrefs with at least one dimension or () -> ()
3560 // for zero-dimensional memrefs.
3561 auto map =
3562 rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap();
3563 build(builder, result, valueToStore, memref, map, indices, alignment);
3564}
3565
3566ParseResult AffineStoreOp::parse(OpAsmParser &parser, OperationState &result) {
3567 auto indexTy = parser.getBuilder().getIndexType();
3568
3569 MemRefType type;
3570 OpAsmParser::UnresolvedOperand storeValueInfo;
3572 AffineMapAttr mapAttr;
3574 return failure(parser.parseOperand(storeValueInfo) || parser.parseComma() ||
3575 parser.parseOperand(memrefInfo) ||
3577 mapOperands, mapAttr, AffineStoreOp::getMapAttrStrName(),
3578 result.attributes) ||
3579 parser.parseOptionalAttrDict(result.attributes) ||
3580 parser.parseColonType(type) ||
3581 parser.resolveOperand(storeValueInfo, type.getElementType(),
3582 result.operands) ||
3583 parser.resolveOperand(memrefInfo, type, result.operands) ||
3584 parser.resolveOperands(mapOperands, indexTy, result.operands));
3585}
3586
3587void AffineStoreOp::print(OpAsmPrinter &p) {
3588 p << " " << getValueToStore();
3589 p << ", " << getMemRef() << '[';
3590 if (AffineMapAttr mapAttr =
3591 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
3592 p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
3593 p << ']';
3594 p.printOptionalAttrDict((*this)->getAttrs(),
3595 /*elidedAttrs=*/{getMapAttrStrName()});
3596 p << " : " << getMemRefType();
3597}
3598
3599LogicalResult AffineStoreOp::verify() {
3600 // The value to store must have the same type as memref element type.
3601 auto memrefType = getMemRefType();
3602 if (getValueToStore().getType() != memrefType.getElementType())
3603 return emitOpError(
3604 "value to store must have the same type as memref element type");
3605
3607 *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
3608 getMapOperands(), memrefType,
3609 /*numIndexOperands=*/getNumOperands() - 2)))
3610 return failure();
3611
3612 return success();
3613}
3614
3615void AffineStoreOp::getCanonicalizationPatterns(RewritePatternSet &results,
3616 MLIRContext *context) {
3617 results.add<SimplifyAffineOp<AffineStoreOp>>(context);
3618}
3619
3620LogicalResult AffineStoreOp::fold(FoldAdaptor adaptor,
3622 /// store(memrefcast) -> store
3623 return memref::foldMemRefCast(*this, getValueToStore());
3624}
3625
3626//===----------------------------------------------------------------------===//
3627// AffineMinMaxOpBase
3628//===----------------------------------------------------------------------===//
3629
3630template <typename T>
3631static LogicalResult verifyAffineMinMaxOp(T op) {
3632 // Verify that operand count matches affine map dimension and symbol count.
3633 if (op.getNumOperands() !=
3634 op.getMap().getNumDims() + op.getMap().getNumSymbols())
3635 return op.emitOpError(
3636 "operand count and affine map dimension and symbol count must match");
3637
3638 if (op.getMap().getNumResults() == 0)
3639 return op.emitOpError("affine map expect at least one result");
3640 return success();
3641}
3642
3643template <typename T>
3644static void printAffineMinMaxOp(OpAsmPrinter &p, T op) {
3645 p << ' ' << op->getAttr(T::getMapAttrStrName());
3646 auto operands = op.getOperands();
3647 unsigned numDims = op.getMap().getNumDims();
3648 p << '(' << operands.take_front(numDims) << ')';
3649
3650 if (operands.size() != numDims)
3651 p << '[' << operands.drop_front(numDims) << ']';
3652 p.printOptionalAttrDict(op->getAttrs(),
3653 /*elidedAttrs=*/{T::getMapAttrStrName()});
3654}
3655
3656template <typename T>
3657static ParseResult parseAffineMinMaxOp(OpAsmParser &parser,
3659 auto &builder = parser.getBuilder();
3660 auto indexType = builder.getIndexType();
3663 AffineMapAttr mapAttr;
3664 return failure(
3665 parser.parseAttribute(mapAttr, T::getMapAttrStrName(),
3666 result.attributes) ||
3668 parser.parseOperandList(symInfos,
3670 parser.parseOptionalAttrDict(result.attributes) ||
3671 parser.resolveOperands(dimInfos, indexType, result.operands) ||
3672 parser.resolveOperands(symInfos, indexType, result.operands) ||
3673 parser.addTypeToList(indexType, result.types));
3674}
3675
3676/// Fold an affine min or max operation with the given operands. The operand
3677/// list may contain nulls, which are interpreted as the operand not being a
3678/// constant.
3679template <typename T>
3681 static_assert(llvm::is_one_of<T, AffineMinOp, AffineMaxOp>::value,
3682 "expected affine min or max op");
3683
3684 // Fold the affine map.
3685 // TODO: Fold more cases:
3686 // min(some_affine, some_affine + constant, ...), etc.
3688 auto foldedMap = op.getMap().partialConstantFold(operands, &results);
3689
3690 if (foldedMap.getNumSymbols() == 1 && foldedMap.isSymbolIdentity())
3691 return op.getOperand(0);
3692
3693 // If some of the map results are not constant, try changing the map in-place.
3694 if (results.empty()) {
3695 // If the map is the same, report that folding did not happen.
3696 if (foldedMap == op.getMap())
3697 return {};
3698 op->setAttr("map", AffineMapAttr::get(foldedMap));
3699 return op.getResult();
3700 }
3701
3702 // Otherwise, completely fold the op into a constant.
3703 auto resultIt = std::is_same<T, AffineMinOp>::value
3704 ? llvm::min_element(results)
3705 : llvm::max_element(results);
3706 if (resultIt == results.end())
3707 return {};
3708 return IntegerAttr::get(IndexType::get(op.getContext()), *resultIt);
3709}
3710
3711/// Remove duplicated expressions in affine min/max ops.
3712template <typename T>
3715
3716 LogicalResult matchAndRewrite(T affineOp,
3717 PatternRewriter &rewriter) const override {
3718 AffineMap oldMap = affineOp.getAffineMap();
3719
3721 for (AffineExpr expr : oldMap.getResults()) {
3722 // This is a linear scan over newExprs, but it should be fine given that
3723 // we typically just have a few expressions per op.
3724 if (!llvm::is_contained(newExprs, expr))
3725 newExprs.push_back(expr);
3726 }
3727
3728 if (newExprs.size() == oldMap.getNumResults())
3729 return failure();
3730
3731 auto newMap = AffineMap::get(oldMap.getNumDims(), oldMap.getNumSymbols(),
3732 newExprs, rewriter.getContext());
3733 rewriter.replaceOpWithNewOp<T>(affineOp, newMap, affineOp.getMapOperands());
3734
3735 return success();
3736 }
3737};
3738
3739/// Merge an affine min/max op to its consumers if its consumer is also an
3740/// affine min/max op.
3741///
3742/// This pattern requires the producer affine min/max op is bound to a
3743/// dimension/symbol that is used as a standalone expression in the consumer
3744/// affine op's map.
3745///
3746/// For example, a pattern like the following:
3747///
3748/// %0 = affine.min affine_map<()[s0] -> (s0 + 16, s0 * 8)> ()[%sym1]
3749/// %1 = affine.min affine_map<(d0)[s0] -> (s0 + 4, d0)> (%0)[%sym2]
3750///
3751/// Can be turned into:
3752///
3753/// %1 = affine.min affine_map<
3754/// ()[s0, s1] -> (s0 + 4, s1 + 16, s1 * 8)> ()[%sym2, %sym1]
3755template <typename T>
3758
3759 LogicalResult matchAndRewrite(T affineOp,
3760 PatternRewriter &rewriter) const override {
3761 AffineMap oldMap = affineOp.getAffineMap();
3762 ValueRange dimOperands =
3763 affineOp.getMapOperands().take_front(oldMap.getNumDims());
3764 ValueRange symOperands =
3765 affineOp.getMapOperands().take_back(oldMap.getNumSymbols());
3766
3767 auto newDimOperands = llvm::to_vector<8>(dimOperands);
3768 auto newSymOperands = llvm::to_vector<8>(symOperands);
3770 SmallVector<T, 4> producerOps;
3771
3772 // Go over each expression to see whether it's a single dimension/symbol
3773 // with the corresponding operand which is the result of another affine
3774 // min/max op. If So it can be merged into this affine op.
3775 for (AffineExpr expr : oldMap.getResults()) {
3776 if (auto symExpr = dyn_cast<AffineSymbolExpr>(expr)) {
3777 Value symValue = symOperands[symExpr.getPosition()];
3778 if (auto producerOp = symValue.getDefiningOp<T>()) {
3779 producerOps.push_back(producerOp);
3780 continue;
3781 }
3782 } else if (auto dimExpr = dyn_cast<AffineDimExpr>(expr)) {
3783 Value dimValue = dimOperands[dimExpr.getPosition()];
3784 if (auto producerOp = dimValue.getDefiningOp<T>()) {
3785 producerOps.push_back(producerOp);
3786 continue;
3787 }
3788 }
3789 // For the above cases we will remove the expression by merging the
3790 // producer affine min/max's affine expressions. Otherwise we need to
3791 // keep the existing expression.
3792 newExprs.push_back(expr);
3793 }
3794
3795 if (producerOps.empty())
3796 return failure();
3797
3798 unsigned numUsedDims = oldMap.getNumDims();
3799 unsigned numUsedSyms = oldMap.getNumSymbols();
3800
3801 // Now go over all producer affine ops and merge their expressions.
3802 for (T producerOp : producerOps) {
3803 AffineMap producerMap = producerOp.getAffineMap();
3804 unsigned numProducerDims = producerMap.getNumDims();
3805 unsigned numProducerSyms = producerMap.getNumSymbols();
3806
3807 // Collect all dimension/symbol values.
3808 ValueRange dimValues =
3809 producerOp.getMapOperands().take_front(numProducerDims);
3810 ValueRange symValues =
3811 producerOp.getMapOperands().take_back(numProducerSyms);
3812 newDimOperands.append(dimValues.begin(), dimValues.end());
3813 newSymOperands.append(symValues.begin(), symValues.end());
3814
3815 // For expressions we need to shift to avoid overlap.
3816 for (AffineExpr expr : producerMap.getResults()) {
3817 newExprs.push_back(expr.shiftDims(numProducerDims, numUsedDims)
3818 .shiftSymbols(numProducerSyms, numUsedSyms));
3819 }
3820
3821 numUsedDims += numProducerDims;
3822 numUsedSyms += numProducerSyms;
3823 }
3824
3825 auto newMap = AffineMap::get(numUsedDims, numUsedSyms, newExprs,
3826 rewriter.getContext());
3827 auto newOperands =
3828 llvm::to_vector<8>(llvm::concat<Value>(newDimOperands, newSymOperands));
3829 rewriter.replaceOpWithNewOp<T>(affineOp, newMap, newOperands);
3830
3831 return success();
3832 }
3833};
3834
3835/// Canonicalize the result expression order of an affine map and return success
3836/// if the order changed.
3837///
3838/// The function flattens the map's affine expressions to coefficient arrays and
3839/// sorts them in lexicographic order. A coefficient array contains a multiplier
3840/// for every dimension/symbol and a constant term. The canonicalization fails
3841/// if a result expression is not pure or if the flattening requires local
3842/// variables that, unlike dimensions and symbols, have no global order.
3843static LogicalResult canonicalizeMapExprAndTermOrder(AffineMap &map) {
3844 SmallVector<SmallVector<int64_t>> flattenedExprs;
3845 for (const AffineExpr &resultExpr : map.getResults()) {
3846 // Fail if the expression is not pure.
3847 if (!resultExpr.isPureAffine())
3848 return failure();
3849
3850 SimpleAffineExprFlattener flattener(map.getNumDims(), map.getNumSymbols());
3851 auto flattenResult = flattener.walkPostOrder(resultExpr);
3852 if (failed(flattenResult))
3853 return failure();
3854
3855 // Fail if the flattened expression has local variables.
3856 if (flattener.operandExprStack.back().size() !=
3857 map.getNumDims() + map.getNumSymbols() + 1)
3858 return failure();
3859
3860 flattenedExprs.emplace_back(flattener.operandExprStack.back().begin(),
3861 flattener.operandExprStack.back().end());
3862 }
3863
3864 // Fail if sorting is not necessary.
3865 if (llvm::is_sorted(flattenedExprs))
3866 return failure();
3867
3868 // Reorder the result expressions according to their flattened form.
3869 SmallVector<unsigned> resultPermutation =
3870 llvm::to_vector(llvm::seq<unsigned>(0, map.getNumResults()));
3871 llvm::sort(resultPermutation, [&](unsigned lhs, unsigned rhs) {
3872 return flattenedExprs[lhs] < flattenedExprs[rhs];
3873 });
3874 SmallVector<AffineExpr> newExprs;
3875 for (unsigned idx : resultPermutation)
3876 newExprs.push_back(map.getResult(idx));
3877
3878 map = AffineMap::get(map.getNumDims(), map.getNumSymbols(), newExprs,
3879 map.getContext());
3880 return success();
3881}
3882
3883/// Canonicalize the affine map result expression order of an affine min/max
3884/// operation.
3885///
3886/// The pattern calls `canonicalizeMapExprAndTermOrder` to order the result
3887/// expressions and replaces the operation if the order changed.
3888///
3889/// For example, the following operation:
3890///
3891/// %0 = affine.min affine_map<(d0, d1) -> (d0 + d1, d1 + 16, 32)> (%i0, %i1)
3892///
3893/// Turns into:
3894///
3895/// %0 = affine.min affine_map<(d0, d1) -> (32, d1 + 16, d0 + d1)> (%i0, %i1)
3896template <typename T>
3899
3900 LogicalResult matchAndRewrite(T affineOp,
3901 PatternRewriter &rewriter) const override {
3902 AffineMap map = affineOp.getAffineMap();
3903 if (failed(canonicalizeMapExprAndTermOrder(map)))
3904 return failure();
3905 rewriter.replaceOpWithNewOp<T>(affineOp, map, affineOp.getMapOperands());
3906 return success();
3907 }
3908};
3909
3910template <typename T>
3913
3914 LogicalResult matchAndRewrite(T affineOp,
3915 PatternRewriter &rewriter) const override {
3916 if (affineOp.getMap().getNumResults() != 1)
3917 return failure();
3918 rewriter.replaceOpWithNewOp<AffineApplyOp>(affineOp, affineOp.getMap(),
3919 affineOp.getOperands());
3920 return success();
3921 }
3922};
3923
3924//===----------------------------------------------------------------------===//
3925// AffineMinOp
3926//===----------------------------------------------------------------------===//
3927//
3928// %0 = affine.min (d0) -> (1000, d0 + 512) (%i0)
3929//
3930
3931OpFoldResult AffineMinOp::fold(FoldAdaptor adaptor) {
3932 return foldMinMaxOp(*this, adaptor.getOperands());
3933}
3934
3935void AffineMinOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
3936 MLIRContext *context) {
3939 MergeAffineMinMaxOp<AffineMinOp>, SimplifyAffineOp<AffineMinOp>,
3941 context);
3942}
3943
3944LogicalResult AffineMinOp::verify() { return verifyAffineMinMaxOp(*this); }
3945
3946ParseResult AffineMinOp::parse(OpAsmParser &parser, OperationState &result) {
3948}
3949
3950void AffineMinOp::print(OpAsmPrinter &p) { printAffineMinMaxOp(p, *this); }
3951
3952//===----------------------------------------------------------------------===//
3953// AffineMaxOp
3954//===----------------------------------------------------------------------===//
3955//
3956// %0 = affine.max (d0) -> (1000, d0 + 512) (%i0)
3957//
3958
3959OpFoldResult AffineMaxOp::fold(FoldAdaptor adaptor) {
3960 return foldMinMaxOp(*this, adaptor.getOperands());
3961}
3962
3963void AffineMaxOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
3964 MLIRContext *context) {
3967 MergeAffineMinMaxOp<AffineMaxOp>, SimplifyAffineOp<AffineMaxOp>,
3969 context);
3970}
3971
3972LogicalResult AffineMaxOp::verify() { return verifyAffineMinMaxOp(*this); }
3973
3974ParseResult AffineMaxOp::parse(OpAsmParser &parser, OperationState &result) {
3976}
3977
3978void AffineMaxOp::print(OpAsmPrinter &p) { printAffineMinMaxOp(p, *this); }
3979
3980//===----------------------------------------------------------------------===//
3981// AffinePrefetchOp
3982//===----------------------------------------------------------------------===//
3983
3984//
3985// affine.prefetch %0[%i, %j + 5], read, locality<3>, data : memref<400x400xi32>
3986//
3987ParseResult AffinePrefetchOp::parse(OpAsmParser &parser,
3989 auto &builder = parser.getBuilder();
3990 auto indexTy = builder.getIndexType();
3991
3992 MemRefType type;
3994 IntegerAttr hintInfo;
3995 auto i32Type = parser.getBuilder().getIntegerType(32);
3996 StringRef readOrWrite, cacheType;
3997
3998 AffineMapAttr mapAttr;
4000 if (parser.parseOperand(memrefInfo) ||
4001 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
4002 AffinePrefetchOp::getMapAttrStrName(),
4003 result.attributes) ||
4004 parser.parseComma() || parser.parseKeyword(&readOrWrite) ||
4005 parser.parseComma() || parser.parseKeyword("locality") ||
4006 parser.parseLess() ||
4007 parser.parseAttribute(hintInfo, i32Type,
4008 AffinePrefetchOp::getLocalityHintAttrStrName(),
4009 result.attributes) ||
4010 parser.parseGreater() || parser.parseComma() ||
4011 parser.parseKeyword(&cacheType) ||
4012 parser.parseOptionalAttrDict(result.attributes) ||
4013 parser.parseColonType(type) ||
4014 parser.resolveOperand(memrefInfo, type, result.operands) ||
4015 parser.resolveOperands(mapOperands, indexTy, result.operands))
4016 return failure();
4017
4018 if (readOrWrite != "read" && readOrWrite != "write")
4019 return parser.emitError(parser.getNameLoc(),
4020 "rw specifier has to be 'read' or 'write'");
4021 result.addAttribute(AffinePrefetchOp::getIsWriteAttrStrName(),
4022 parser.getBuilder().getBoolAttr(readOrWrite == "write"));
4023
4024 if (cacheType != "data" && cacheType != "instr")
4025 return parser.emitError(parser.getNameLoc(),
4026 "cache type has to be 'data' or 'instr'");
4027
4028 result.addAttribute(AffinePrefetchOp::getIsDataCacheAttrStrName(),
4029 parser.getBuilder().getBoolAttr(cacheType == "data"));
4030
4031 return success();
4032}
4033
4034void AffinePrefetchOp::print(OpAsmPrinter &p) {
4035 p << " " << getMemref() << '[';
4036 AffineMapAttr mapAttr =
4037 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
4038 if (mapAttr)
4039 p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
4040 p << ']' << ", " << (getIsWrite() ? "write" : "read") << ", " << "locality<"
4041 << getLocalityHint() << ">, " << (getIsDataCache() ? "data" : "instr");
4043 (*this)->getAttrs(),
4044 /*elidedAttrs=*/{getMapAttrStrName(), getLocalityHintAttrStrName(),
4045 getIsDataCacheAttrStrName(), getIsWriteAttrStrName()});
4046 p << " : " << getMemRefType();
4047}
4048
4049LogicalResult AffinePrefetchOp::verify() {
4050 auto mapAttr = (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
4051 if (mapAttr) {
4052 AffineMap map = mapAttr.getValue();
4053 if (map.getNumResults() != getMemRefType().getRank())
4054 return emitOpError("affine.prefetch affine map num results must equal"
4055 " memref rank");
4056 if (map.getNumInputs() + 1 != getNumOperands())
4057 return emitOpError("too few operands");
4058 } else {
4059 if (getNumOperands() != 1)
4060 return emitOpError("too few operands");
4061 }
4062
4063 Region *scope = getAffineScope(*this);
4064 for (auto idx : getMapOperands()) {
4065 if (!isValidAffineIndexOperand(idx, scope))
4066 return emitOpError(
4067 "index must be a valid dimension or symbol identifier");
4068 }
4069 return success();
4070}
4071
4072void AffinePrefetchOp::getCanonicalizationPatterns(RewritePatternSet &results,
4073 MLIRContext *context) {
4074 // prefetch(memrefcast) -> prefetch
4075 results.add<SimplifyAffineOp<AffinePrefetchOp>>(context);
4076}
4077
4078LogicalResult AffinePrefetchOp::fold(FoldAdaptor adaptor,
4080 /// prefetch(memrefcast) -> prefetch
4081 return memref::foldMemRefCast(*this);
4082}
4083
4084//===----------------------------------------------------------------------===//
4085// AffineParallelOp
4086//===----------------------------------------------------------------------===//
4087
4088void AffineParallelOp::build(OpBuilder &builder, OperationState &result,
4089 TypeRange resultTypes,
4091 ArrayRef<int64_t> ranges) {
4092 SmallVector<AffineMap> lbs(ranges.size(), builder.getConstantAffineMap(0));
4093 auto ubs = llvm::map_to_vector<4>(ranges, [&](int64_t value) {
4094 return builder.getConstantAffineMap(value);
4095 });
4096 SmallVector<int64_t> steps(ranges.size(), 1);
4097 build(builder, result, resultTypes, reductions, lbs, /*lbArgs=*/{}, ubs,
4098 /*ubArgs=*/{}, steps);
4099}
4100
4101void AffineParallelOp::build(OpBuilder &builder, OperationState &result,
4102 TypeRange resultTypes,
4104 ArrayRef<AffineMap> lbMaps, ValueRange lbArgs,
4105 ArrayRef<AffineMap> ubMaps, ValueRange ubArgs,
4106 ArrayRef<int64_t> steps) {
4107 assert(llvm::all_of(lbMaps,
4108 [lbMaps](AffineMap m) {
4109 return m.getNumDims() == lbMaps[0].getNumDims() &&
4110 m.getNumSymbols() == lbMaps[0].getNumSymbols();
4111 }) &&
4112 "expected all lower bounds maps to have the same number of dimensions "
4113 "and symbols");
4114 assert(llvm::all_of(ubMaps,
4115 [ubMaps](AffineMap m) {
4116 return m.getNumDims() == ubMaps[0].getNumDims() &&
4117 m.getNumSymbols() == ubMaps[0].getNumSymbols();
4118 }) &&
4119 "expected all upper bounds maps to have the same number of dimensions "
4120 "and symbols");
4121 assert((lbMaps.empty() || lbMaps[0].getNumInputs() == lbArgs.size()) &&
4122 "expected lower bound maps to have as many inputs as lower bound "
4123 "operands");
4124 assert((ubMaps.empty() || ubMaps[0].getNumInputs() == ubArgs.size()) &&
4125 "expected upper bound maps to have as many inputs as upper bound "
4126 "operands");
4127
4128 OpBuilder::InsertionGuard guard(builder);
4129 result.addTypes(resultTypes);
4130
4131 // Convert the reductions to integer attributes.
4132 SmallVector<Attribute, 4> reductionAttrs;
4133 for (arith::AtomicRMWKind reduction : reductions)
4134 reductionAttrs.push_back(
4135 builder.getI64IntegerAttr(static_cast<int64_t>(reduction)));
4136 result.addAttribute(getReductionsAttrStrName(),
4137 builder.getArrayAttr(reductionAttrs));
4138
4139 // Concatenates maps defined in the same input space (same dimensions and
4140 // symbols), assumes there is at least one map.
4141 auto concatMapsSameInput = [&builder](ArrayRef<AffineMap> maps,
4142 SmallVectorImpl<int32_t> &groups) {
4143 if (maps.empty())
4144 return AffineMap::get(builder.getContext());
4146 groups.reserve(groups.size() + maps.size());
4147 exprs.reserve(maps.size());
4148 for (AffineMap m : maps) {
4149 llvm::append_range(exprs, m.getResults());
4150 groups.push_back(m.getNumResults());
4151 }
4152 return AffineMap::get(maps[0].getNumDims(), maps[0].getNumSymbols(), exprs,
4153 maps[0].getContext());
4154 };
4155
4156 // Set up the bounds.
4157 SmallVector<int32_t> lbGroups, ubGroups;
4158 AffineMap lbMap = concatMapsSameInput(lbMaps, lbGroups);
4159 AffineMap ubMap = concatMapsSameInput(ubMaps, ubGroups);
4160 result.addAttribute(getLowerBoundsMapAttrStrName(),
4161 AffineMapAttr::get(lbMap));
4162 result.addAttribute(getLowerBoundsGroupsAttrStrName(),
4163 builder.getI32TensorAttr(lbGroups));
4164 result.addAttribute(getUpperBoundsMapAttrStrName(),
4165 AffineMapAttr::get(ubMap));
4166 result.addAttribute(getUpperBoundsGroupsAttrStrName(),
4167 builder.getI32TensorAttr(ubGroups));
4168 result.addAttribute(getStepsAttrStrName(), builder.getI64ArrayAttr(steps));
4169 result.addOperands(lbArgs);
4170 result.addOperands(ubArgs);
4171
4172 // Create a region and a block for the body.
4173 auto *bodyRegion = result.addRegion();
4174 Block *body = builder.createBlock(bodyRegion);
4175
4176 // Add all the block arguments.
4177 for (unsigned i = 0, e = steps.size(); i < e; ++i)
4178 body->addArgument(IndexType::get(builder.getContext()), result.location);
4179 if (resultTypes.empty())
4180 ensureTerminator(*bodyRegion, builder, result.location);
4181}
4182
4183SmallVector<Region *> AffineParallelOp::getLoopRegions() {
4184 return {&getRegion()};
4185}
4186
4187unsigned AffineParallelOp::getNumDims() { return getSteps().size(); }
4188
4189AffineParallelOp::operand_range AffineParallelOp::getLowerBoundsOperands() {
4190 return getOperands().take_front(getLowerBoundsMap().getNumInputs());
4191}
4192
4193AffineParallelOp::operand_range AffineParallelOp::getUpperBoundsOperands() {
4194 return getOperands().drop_front(getLowerBoundsMap().getNumInputs());
4195}
4196
4197AffineMap AffineParallelOp::getLowerBoundMap(unsigned pos) {
4198 auto values = getLowerBoundsGroups().getValues<int32_t>();
4199 unsigned start = 0;
4200 for (unsigned i = 0; i < pos; ++i)
4201 start += values[i];
4202 return getLowerBoundsMap().getSliceMap(start, values[pos]);
4203}
4204
4205AffineMap AffineParallelOp::getUpperBoundMap(unsigned pos) {
4206 auto values = getUpperBoundsGroups().getValues<int32_t>();
4207 unsigned start = 0;
4208 for (unsigned i = 0; i < pos; ++i)
4209 start += values[i];
4210 return getUpperBoundsMap().getSliceMap(start, values[pos]);
4211}
4212
4213AffineValueMap AffineParallelOp::getLowerBoundsValueMap() {
4214 return AffineValueMap(getLowerBoundsMap(), getLowerBoundsOperands());
4215}
4216
4217AffineValueMap AffineParallelOp::getUpperBoundsValueMap() {
4218 return AffineValueMap(getUpperBoundsMap(), getUpperBoundsOperands());
4219}
4220
4221std::optional<SmallVector<int64_t, 8>> AffineParallelOp::getConstantRanges() {
4222 if (hasMinMaxBounds())
4223 return std::nullopt;
4224
4225 // Try to convert all the ranges to constant expressions.
4227 AffineValueMap rangesValueMap;
4228 AffineValueMap::difference(getUpperBoundsValueMap(), getLowerBoundsValueMap(),
4229 &rangesValueMap);
4230 out.reserve(rangesValueMap.getNumResults());
4231 for (unsigned i = 0, e = rangesValueMap.getNumResults(); i < e; ++i) {
4232 auto expr = rangesValueMap.getResult(i);
4233 auto cst = dyn_cast<AffineConstantExpr>(expr);
4234 if (!cst)
4235 return std::nullopt;
4236 out.push_back(cst.getValue());
4237 }
4238 return out;
4239}
4240
4241Block *AffineParallelOp::getBody() { return &getRegion().front(); }
4242
4243OpBuilder AffineParallelOp::getBodyBuilder() {
4244 return OpBuilder(getBody(), std::prev(getBody()->end()));
4245}
4246
4247void AffineParallelOp::setLowerBounds(ValueRange lbOperands, AffineMap map) {
4248 assert(lbOperands.size() == map.getNumInputs() &&
4249 "operands to map must match number of inputs");
4250
4251 auto ubOperands = getUpperBoundsOperands();
4252
4253 SmallVector<Value, 4> newOperands(lbOperands);
4254 newOperands.append(ubOperands.begin(), ubOperands.end());
4255 (*this)->setOperands(newOperands);
4256
4257 setLowerBoundsMapAttr(AffineMapAttr::get(map));
4258}
4259
4260void AffineParallelOp::setUpperBounds(ValueRange ubOperands, AffineMap map) {
4261 assert(ubOperands.size() == map.getNumInputs() &&
4262 "operands to map must match number of inputs");
4263
4264 SmallVector<Value, 4> newOperands(getLowerBoundsOperands());
4265 newOperands.append(ubOperands.begin(), ubOperands.end());
4266 (*this)->setOperands(newOperands);
4267
4268 setUpperBoundsMapAttr(AffineMapAttr::get(map));
4269}
4270
4271void AffineParallelOp::setSteps(ArrayRef<int64_t> newSteps) {
4272 setStepsAttr(getBodyBuilder().getI64ArrayAttr(newSteps));
4273}
4274
4275// check whether resultType match op or not in affine.parallel
4277 arith::AtomicRMWKind op) {
4278 switch (op) {
4279 case arith::AtomicRMWKind::addf:
4280 return isa<FloatType>(resultType);
4281 case arith::AtomicRMWKind::addi:
4282 return isa<IntegerType>(resultType);
4283 case arith::AtomicRMWKind::assign:
4284 return true;
4285 case arith::AtomicRMWKind::mulf:
4286 return isa<FloatType>(resultType);
4287 case arith::AtomicRMWKind::muli:
4288 return isa<IntegerType>(resultType);
4289 case arith::AtomicRMWKind::maximumf:
4290 case arith::AtomicRMWKind::maxnumf:
4291 case arith::AtomicRMWKind::minimumf:
4292 case arith::AtomicRMWKind::minnumf:
4293 return isa<FloatType>(resultType);
4294 case arith::AtomicRMWKind::maxs: {
4295 auto intType = dyn_cast<IntegerType>(resultType);
4296 return intType && intType.isSigned();
4297 }
4298 case arith::AtomicRMWKind::mins: {
4299 auto intType = dyn_cast<IntegerType>(resultType);
4300 return intType && intType.isSigned();
4301 }
4302 case arith::AtomicRMWKind::maxu: {
4303 auto intType = dyn_cast<IntegerType>(resultType);
4304 return intType && intType.isUnsigned();
4305 }
4306 case arith::AtomicRMWKind::minu: {
4307 auto intType = dyn_cast<IntegerType>(resultType);
4308 return intType && intType.isUnsigned();
4309 }
4310 case arith::AtomicRMWKind::ori:
4311 case arith::AtomicRMWKind::andi:
4312 case arith::AtomicRMWKind::xori:
4313 return isa<IntegerType>(resultType);
4314 }
4315 llvm_unreachable("Unhandled atomic rmw kind");
4316}
4317
4318LogicalResult AffineParallelOp::verify() {
4319 auto numDims = getNumDims();
4320 if (getLowerBoundsGroups().getNumElements() != numDims ||
4321 getUpperBoundsGroups().getNumElements() != numDims ||
4322 getSteps().size() != numDims || getBody()->getNumArguments() != numDims) {
4323 return emitOpError() << "the number of region arguments ("
4324 << getBody()->getNumArguments()
4325 << ") and the number of map groups for lower ("
4326 << getLowerBoundsGroups().getNumElements()
4327 << ") and upper bound ("
4328 << getUpperBoundsGroups().getNumElements()
4329 << "), and the number of steps (" << getSteps().size()
4330 << ") must all match";
4331 }
4332
4333 unsigned expectedNumLBResults = 0;
4334 for (APInt v : getLowerBoundsGroups()) {
4335 unsigned results = v.getZExtValue();
4336 if (results == 0)
4337 return emitOpError()
4338 << "expected lower bound map to have at least one result";
4339 expectedNumLBResults += results;
4340 }
4341 if (expectedNumLBResults != getLowerBoundsMap().getNumResults())
4342 return emitOpError() << "expected lower bounds map to have "
4343 << expectedNumLBResults << " results";
4344 unsigned expectedNumUBResults = 0;
4345 for (APInt v : getUpperBoundsGroups()) {
4346 unsigned results = v.getZExtValue();
4347 if (results == 0)
4348 return emitOpError()
4349 << "expected upper bound map to have at least one result";
4350 expectedNumUBResults += results;
4351 }
4352 if (expectedNumUBResults != getUpperBoundsMap().getNumResults())
4353 return emitOpError() << "expected upper bounds map to have "
4354 << expectedNumUBResults << " results";
4355
4356 if (getReductions().size() != getNumResults())
4357 return emitOpError("a reduction must be specified for each output");
4358
4359 // Verify reduction ops are all valid and each result type matches reduction
4360 // ops
4361 for (auto it : llvm::enumerate((getReductions()))) {
4362 Attribute attr = it.value();
4363 auto intAttr = dyn_cast<IntegerAttr>(attr);
4364 if (!intAttr || !arith::symbolizeAtomicRMWKind(intAttr.getInt()))
4365 return emitOpError("invalid reduction attribute");
4366 auto kind = arith::symbolizeAtomicRMWKind(intAttr.getInt()).value();
4367 if (!isResultTypeMatchAtomicRMWKind(getResult(it.index()).getType(), kind))
4368 return emitOpError("result type cannot match reduction attribute");
4369 }
4370
4371 // Verify that the bound operands are valid dimension/symbols.
4372 /// Lower bounds.
4373 if (failed(verifyDimAndSymbolIdentifiers(*this, getLowerBoundsOperands(),
4374 getLowerBoundsMap().getNumDims())))
4375 return failure();
4376 /// Upper bounds.
4377 if (failed(verifyDimAndSymbolIdentifiers(*this, getUpperBoundsOperands(),
4378 getUpperBoundsMap().getNumDims())))
4379 return failure();
4380 return success();
4381}
4382
4384 SmallVector<Value, 4> newOperands{operands};
4385 auto newMap = getAffineMap();
4386 composeAffineMapAndOperands(&newMap, &newOperands);
4387 if (newMap == getAffineMap() && newOperands == operands)
4388 return failure();
4389 reset(newMap, newOperands);
4390 return success();
4391}
4392
4393/// Canonicalize the bounds of the given loop.
4394static LogicalResult canonicalizeLoopBounds(AffineParallelOp op) {
4395 AffineValueMap lb = op.getLowerBoundsValueMap();
4396 bool lbCanonicalized = succeeded(lb.canonicalize());
4397
4398 AffineValueMap ub = op.getUpperBoundsValueMap();
4399 bool ubCanonicalized = succeeded(ub.canonicalize());
4400
4401 // Any canonicalization change always leads to updated map(s).
4402 if (!lbCanonicalized && !ubCanonicalized)
4403 return failure();
4404
4405 if (lbCanonicalized)
4406 op.setLowerBounds(lb.getOperands(), lb.getAffineMap());
4407 if (ubCanonicalized)
4408 op.setUpperBounds(ub.getOperands(), ub.getAffineMap());
4409
4410 return success();
4411}
4412
4413LogicalResult AffineParallelOp::fold(FoldAdaptor adaptor,
4414 SmallVectorImpl<OpFoldResult> &results) {
4415 return canonicalizeLoopBounds(*this);
4416}
4417
4418/// Prints a lower(upper) bound of an affine parallel loop with max(min)
4419/// conditions in it. `mapAttr` is a flat list of affine expressions and `group`
4420/// identifies which of the those expressions form max/min groups. `operands`
4421/// are the SSA values of dimensions and symbols and `keyword` is either "min"
4422/// or "max".
4423static void printMinMaxBound(OpAsmPrinter &p, AffineMapAttr mapAttr,
4424 DenseIntElementsAttr group, ValueRange operands,
4425 StringRef keyword) {
4426 AffineMap map = mapAttr.getValue();
4427 unsigned numDims = map.getNumDims();
4428 ValueRange dimOperands = operands.take_front(numDims);
4429 ValueRange symOperands = operands.drop_front(numDims);
4430 unsigned start = 0;
4431 for (llvm::APInt groupSize : group) {
4432 if (start != 0)
4433 p << ", ";
4434
4435 unsigned size = groupSize.getZExtValue();
4436 if (size == 1) {
4437 p.printAffineExprOfSSAIds(map.getResult(start), dimOperands, symOperands);
4438 ++start;
4439 } else {
4440 p << keyword << '(';
4441 AffineMap submap = map.getSliceMap(start, size);
4442 p.printAffineMapOfSSAIds(AffineMapAttr::get(submap), operands);
4443 p << ')';
4444 start += size;
4445 }
4446 }
4447}
4448
4449void AffineParallelOp::print(OpAsmPrinter &p) {
4450 p << " (" << getBody()->getArguments() << ") = (";
4451 printMinMaxBound(p, getLowerBoundsMapAttr(), getLowerBoundsGroupsAttr(),
4452 getLowerBoundsOperands(), "max");
4453 p << ") to (";
4454 printMinMaxBound(p, getUpperBoundsMapAttr(), getUpperBoundsGroupsAttr(),
4455 getUpperBoundsOperands(), "min");
4456 p << ')';
4457 SmallVector<int64_t, 8> steps = getSteps();
4458 bool elideSteps = llvm::all_of(steps, [](int64_t step) { return step == 1; });
4459 if (!elideSteps) {
4460 p << " step (";
4461 llvm::interleaveComma(steps, p);
4462 p << ')';
4463 }
4464 if (getNumResults()) {
4465 p << " reduce (";
4466 llvm::interleaveComma(getReductions(), p, [&](auto &attr) {
4467 arith::AtomicRMWKind sym = *arith::symbolizeAtomicRMWKind(
4468 llvm::cast<IntegerAttr>(attr).getInt());
4469 p << "\"" << arith::stringifyAtomicRMWKind(sym) << "\"";
4470 });
4471 p << ") -> (" << getResultTypes() << ")";
4472 }
4473
4474 p << ' ';
4475 p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
4476 /*printBlockTerminators=*/getNumResults());
4478 (*this)->getAttrs(),
4479 /*elidedAttrs=*/{AffineParallelOp::getReductionsAttrStrName(),
4480 AffineParallelOp::getLowerBoundsMapAttrStrName(),
4481 AffineParallelOp::getLowerBoundsGroupsAttrStrName(),
4482 AffineParallelOp::getUpperBoundsMapAttrStrName(),
4483 AffineParallelOp::getUpperBoundsGroupsAttrStrName(),
4484 AffineParallelOp::getStepsAttrStrName()});
4485}
4486
4487/// Given a list of lists of parsed operands, populates `uniqueOperands` with
4488/// unique operands. Also populates `replacements with affine expressions of
4489/// `kind` that can be used to update affine maps previously accepting a
4490/// `operands` to accept `uniqueOperands` instead.
4491static ParseResult deduplicateAndResolveOperands(
4492 OpAsmParser &parser,
4493 ArrayRef<SmallVector<OpAsmParser::UnresolvedOperand>> operands,
4494 SmallVectorImpl<Value> &uniqueOperands,
4495 SmallVectorImpl<AffineExpr> &replacements, AffineExprKind kind) {
4496 assert((kind == AffineExprKind::DimId || kind == AffineExprKind::SymbolId) &&
4497 "expected operands to be dim or symbol expression");
4498
4499 Type indexType = parser.getBuilder().getIndexType();
4500 for (const auto &list : operands) {
4501 SmallVector<Value> valueOperands;
4502 if (parser.resolveOperands(list, indexType, valueOperands))
4503 return failure();
4504 for (Value operand : valueOperands) {
4505 unsigned pos = std::distance(uniqueOperands.begin(),
4506 llvm::find(uniqueOperands, operand));
4507 if (pos == uniqueOperands.size())
4508 uniqueOperands.push_back(operand);
4509 replacements.push_back(
4510 kind == AffineExprKind::DimId
4511 ? getAffineDimExpr(pos, parser.getContext())
4512 : getAffineSymbolExpr(pos, parser.getContext()));
4513 }
4514 }
4515 return success();
4516}
4517
4518namespace {
4519enum class MinMaxKind { Min, Max };
4520} // namespace
4521
4522/// Parses an affine map that can contain a min/max for groups of its results,
4523/// e.g., max(expr-1, expr-2), expr-3, max(expr-4, expr-5, expr-6). Populates
4524/// `result` attributes with the map (flat list of expressions) and the grouping
4525/// (list of integers that specify how many expressions to put into each
4526/// min/max) attributes. Deduplicates repeated operands.
4527///
4528/// parallel-bound ::= `(` parallel-group-list `)`
4529/// parallel-group-list ::= parallel-group (`,` parallel-group-list)?
4530/// parallel-group ::= simple-group | min-max-group
4531/// simple-group ::= expr-of-ssa-ids
4532/// min-max-group ::= ( `min` | `max` ) `(` expr-of-ssa-ids-list `)`
4533/// expr-of-ssa-ids-list ::= expr-of-ssa-ids (`,` expr-of-ssa-id-list)?
4534///
4535/// Examples:
4536/// (%0, min(%1 + %2, %3), %4, min(%5 floordiv 32, %6))
4537/// (%0, max(%1 - 2 * %2))
4538static ParseResult parseAffineMapWithMinMax(OpAsmParser &parser,
4539 OperationState &result,
4540 MinMaxKind kind) {
4541 // Using `const` not `constexpr` below to workaround a MSVC optimizer bug,
4542 // see: https://reviews.llvm.org/D134227#3821753
4543 const llvm::StringLiteral tmpAttrStrName = "__pseudo_bound_map";
4544
4545 StringRef mapName = kind == MinMaxKind::Min
4546 ? AffineParallelOp::getUpperBoundsMapAttrStrName()
4547 : AffineParallelOp::getLowerBoundsMapAttrStrName();
4548 StringRef groupsName =
4549 kind == MinMaxKind::Min
4550 ? AffineParallelOp::getUpperBoundsGroupsAttrStrName()
4551 : AffineParallelOp::getLowerBoundsGroupsAttrStrName();
4552
4553 if (failed(parser.parseLParen()))
4554 return failure();
4555
4556 if (succeeded(parser.parseOptionalRParen())) {
4557 result.addAttribute(
4558 mapName, AffineMapAttr::get(parser.getBuilder().getEmptyAffineMap()));
4559 result.addAttribute(groupsName, parser.getBuilder().getI32TensorAttr({}));
4560 return success();
4561 }
4562
4563 SmallVector<AffineExpr> flatExprs;
4564 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatDimOperands;
4565 SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatSymOperands;
4566 SmallVector<int32_t> numMapsPerGroup;
4567 SmallVector<OpAsmParser::UnresolvedOperand> mapOperands;
4568 auto parseOperands = [&]() {
4569 if (succeeded(parser.parseOptionalKeyword(
4570 kind == MinMaxKind::Min ? "min" : "max"))) {
4571 mapOperands.clear();
4572 AffineMapAttr map;
4573 if (failed(parser.parseAffineMapOfSSAIds(mapOperands, map, tmpAttrStrName,
4574 result.attributes,
4576 return failure();
4577 result.attributes.erase(tmpAttrStrName);
4578 llvm::append_range(flatExprs, map.getValue().getResults());
4579 auto operandsRef = llvm::ArrayRef(mapOperands);
4580 auto dimsRef = operandsRef.take_front(map.getValue().getNumDims());
4581 SmallVector<OpAsmParser::UnresolvedOperand> dims(dimsRef);
4582 auto symsRef = operandsRef.drop_front(map.getValue().getNumDims());
4583 SmallVector<OpAsmParser::UnresolvedOperand> syms(symsRef);
4584 flatDimOperands.append(map.getValue().getNumResults(), dims);
4585 flatSymOperands.append(map.getValue().getNumResults(), syms);
4586 numMapsPerGroup.push_back(map.getValue().getNumResults());
4587 } else {
4588 if (failed(parser.parseAffineExprOfSSAIds(flatDimOperands.emplace_back(),
4589 flatSymOperands.emplace_back(),
4590 flatExprs.emplace_back())))
4591 return failure();
4592 numMapsPerGroup.push_back(1);
4593 }
4594 return success();
4595 };
4596 if (parser.parseCommaSeparatedList(parseOperands) || parser.parseRParen())
4597 return failure();
4598
4599 unsigned totalNumDims = 0;
4600 unsigned totalNumSyms = 0;
4601 for (unsigned i = 0, e = flatExprs.size(); i < e; ++i) {
4602 unsigned numDims = flatDimOperands[i].size();
4603 unsigned numSyms = flatSymOperands[i].size();
4604 flatExprs[i] = flatExprs[i]
4605 .shiftDims(numDims, totalNumDims)
4606 .shiftSymbols(numSyms, totalNumSyms);
4607 totalNumDims += numDims;
4608 totalNumSyms += numSyms;
4609 }
4610
4611 // Deduplicate map operands.
4612 SmallVector<Value> dimOperands, symOperands;
4613 SmallVector<AffineExpr> dimRplacements, symRepacements;
4614 if (deduplicateAndResolveOperands(parser, flatDimOperands, dimOperands,
4615 dimRplacements, AffineExprKind::DimId) ||
4616 deduplicateAndResolveOperands(parser, flatSymOperands, symOperands,
4617 symRepacements, AffineExprKind::SymbolId))
4618 return failure();
4619
4620 result.operands.append(dimOperands.begin(), dimOperands.end());
4621 result.operands.append(symOperands.begin(), symOperands.end());
4622
4623 Builder &builder = parser.getBuilder();
4624 auto flatMap = AffineMap::get(totalNumDims, totalNumSyms, flatExprs,
4625 parser.getContext());
4626 flatMap = flatMap.replaceDimsAndSymbols(
4627 dimRplacements, symRepacements, dimOperands.size(), symOperands.size());
4628
4629 result.addAttribute(mapName, AffineMapAttr::get(flatMap));
4630 result.addAttribute(groupsName, builder.getI32TensorAttr(numMapsPerGroup));
4631 return success();
4632}
4633
4634//
4635// operation ::= `affine.parallel` `(` ssa-ids `)` `=` parallel-bound
4636// `to` parallel-bound steps? region attr-dict?
4637// steps ::= `steps` `(` integer-literals `)`
4638//
4639ParseResult AffineParallelOp::parse(OpAsmParser &parser,
4640 OperationState &result) {
4641 auto &builder = parser.getBuilder();
4642 auto indexType = builder.getIndexType();
4643 SmallVector<OpAsmParser::Argument, 4> ivs;
4645 parser.parseEqual() ||
4646 parseAffineMapWithMinMax(parser, result, MinMaxKind::Max) ||
4647 parser.parseKeyword("to") ||
4648 parseAffineMapWithMinMax(parser, result, MinMaxKind::Min))
4649 return failure();
4650
4651 AffineMapAttr stepsMapAttr;
4652 NamedAttrList stepsAttrs;
4653 SmallVector<OpAsmParser::UnresolvedOperand, 4> stepsMapOperands;
4654 if (failed(parser.parseOptionalKeyword("step"))) {
4655 SmallVector<int64_t, 4> steps(ivs.size(), 1);
4656 result.addAttribute(AffineParallelOp::getStepsAttrStrName(),
4657 builder.getI64ArrayAttr(steps));
4658 } else {
4659 if (parser.parseAffineMapOfSSAIds(stepsMapOperands, stepsMapAttr,
4660 AffineParallelOp::getStepsAttrStrName(),
4661 stepsAttrs,
4663 return failure();
4664
4665 // Convert steps from an AffineMap into an I64ArrayAttr.
4666 SmallVector<int64_t, 4> steps;
4667 auto stepsMap = stepsMapAttr.getValue();
4668 for (const auto &result : stepsMap.getResults()) {
4669 auto constExpr = dyn_cast<AffineConstantExpr>(result);
4670 if (!constExpr)
4671 return parser.emitError(parser.getNameLoc(),
4672 "steps must be constant integers");
4673 steps.push_back(constExpr.getValue());
4674 }
4675 result.addAttribute(AffineParallelOp::getStepsAttrStrName(),
4676 builder.getI64ArrayAttr(steps));
4677 }
4678
4679 // Parse optional clause of the form: `reduce ("addf", "maxf")`, where the
4680 // quoted strings are a member of the enum AtomicRMWKind.
4681 SmallVector<Attribute, 4> reductions;
4682 if (succeeded(parser.parseOptionalKeyword("reduce"))) {
4683 if (parser.parseLParen())
4684 return failure();
4685 auto parseAttributes = [&]() -> ParseResult {
4686 // Parse a single quoted string via the attribute parsing, and then
4687 // verify it is a member of the enum and convert to it's integer
4688 // representation.
4689 StringAttr attrVal;
4690 NamedAttrList attrStorage;
4691 auto loc = parser.getCurrentLocation();
4692 if (parser.parseAttribute(attrVal, builder.getNoneType(), "reduce",
4693 attrStorage))
4694 return failure();
4695 std::optional<arith::AtomicRMWKind> reduction =
4696 arith::symbolizeAtomicRMWKind(attrVal.getValue());
4697 if (!reduction)
4698 return parser.emitError(loc, "invalid reduction value: ") << attrVal;
4699 reductions.push_back(
4700 builder.getI64IntegerAttr(static_cast<int64_t>(reduction.value())));
4701 // While we keep getting commas, keep parsing.
4702 return success();
4703 };
4704 if (parser.parseCommaSeparatedList(parseAttributes) || parser.parseRParen())
4705 return failure();
4706 }
4707 result.addAttribute(AffineParallelOp::getReductionsAttrStrName(),
4708 builder.getArrayAttr(reductions));
4709
4710 // Parse return types of reductions (if any)
4711 if (parser.parseOptionalArrowTypeList(result.types))
4712 return failure();
4713
4714 // Now parse the body.
4715 Region *body = result.addRegion();
4716 for (auto &iv : ivs)
4717 iv.type = indexType;
4718 if (parser.parseRegion(*body, ivs) ||
4719 parser.parseOptionalAttrDict(result.attributes))
4720 return failure();
4721
4722 // Add a terminator if none was parsed.
4723 AffineParallelOp::ensureTerminator(*body, builder, result.location);
4724 return success();
4725}
4726
4727//===----------------------------------------------------------------------===//
4728// AffineYieldOp
4729//===----------------------------------------------------------------------===//
4730
4731LogicalResult AffineYieldOp::verify() {
4732 auto *parentOp = (*this)->getParentOp();
4733 auto results = parentOp->getResults();
4734 auto operands = getOperands();
4735
4736 if (!isa<AffineParallelOp, AffineIfOp, AffineForOp>(parentOp))
4737 return emitOpError() << "only terminates affine.if/for/parallel regions";
4738 if (parentOp->getNumResults() != getNumOperands())
4739 return emitOpError() << "parent of yield must have same number of "
4740 "results as the yield operands";
4741 for (auto it : llvm::zip(results, operands)) {
4742 if (std::get<0>(it).getType() != std::get<1>(it).getType())
4743 return emitOpError() << "types mismatch between yield op and its parent";
4744 }
4745
4746 return success();
4747}
4748
4749//===----------------------------------------------------------------------===//
4750// AffineVectorLoadOp
4751//===----------------------------------------------------------------------===//
4752
4753void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result,
4754 VectorType resultType, AffineMap map,
4755 ValueRange operands,
4756 llvm::MaybeAlign alignment) {
4757 assert(operands.size() == 1 + map.getNumInputs() && "inconsistent operands");
4758 result.addOperands(operands);
4759 if (map)
4760 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4761 addAlignmentAttr(builder, result, getAlignmentAttrName(result.name),
4762 alignment);
4763 result.types.push_back(resultType);
4764}
4765
4766void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result,
4767 VectorType resultType, Value memref,
4768 AffineMap map, ValueRange mapOperands,
4769 llvm::MaybeAlign alignment) {
4770 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
4771 result.addOperands(memref);
4772 result.addOperands(mapOperands);
4773 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4774 addAlignmentAttr(builder, result, getAlignmentAttrName(result.name),
4775 alignment);
4776 result.types.push_back(resultType);
4777}
4778
4779void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result,
4780 VectorType resultType, Value memref,
4781 ValueRange indices, llvm::MaybeAlign alignment) {
4782 auto memrefType = llvm::cast<MemRefType>(memref.getType());
4783 int64_t rank = memrefType.getRank();
4784 // Create identity map for memrefs with at least one dimension or () -> ()
4785 // for zero-dimensional memrefs.
4786 auto map =
4787 rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap();
4788 build(builder, result, resultType, memref, map, indices, alignment);
4789}
4790
4791void AffineVectorLoadOp::getCanonicalizationPatterns(RewritePatternSet &results,
4792 MLIRContext *context) {
4793 results.add<SimplifyAffineOp<AffineVectorLoadOp>>(context);
4794}
4795
4796ParseResult AffineVectorLoadOp::parse(OpAsmParser &parser,
4797 OperationState &result) {
4798 auto &builder = parser.getBuilder();
4799 auto indexTy = builder.getIndexType();
4800
4801 MemRefType memrefType;
4802 VectorType resultType;
4803 OpAsmParser::UnresolvedOperand memrefInfo;
4804 AffineMapAttr mapAttr;
4805 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
4806 return failure(
4807 parser.parseOperand(memrefInfo) ||
4808 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
4809 AffineVectorLoadOp::getMapAttrStrName(),
4810 result.attributes) ||
4811 parser.parseOptionalAttrDict(result.attributes) ||
4812 parser.parseColonType(memrefType) || parser.parseComma() ||
4813 parser.parseType(resultType) ||
4814 parser.resolveOperand(memrefInfo, memrefType, result.operands) ||
4815 parser.resolveOperands(mapOperands, indexTy, result.operands) ||
4816 parser.addTypeToList(resultType, result.types));
4817}
4818
4819void AffineVectorLoadOp::print(OpAsmPrinter &p) {
4820 p << " " << getMemRef() << '[';
4821 if (AffineMapAttr mapAttr =
4822 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
4823 p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
4824 p << ']';
4825 p.printOptionalAttrDict((*this)->getAttrs(),
4826 /*elidedAttrs=*/{getMapAttrStrName()});
4827 p << " : " << getMemRefType() << ", " << getType();
4828}
4829
4830/// Verify common invariants of affine.vector_load and affine.vector_store.
4831static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType,
4832 VectorType vectorType) {
4833 // Check that memref and vector element types match.
4834 if (memrefType.getElementType() != vectorType.getElementType())
4835 return op->emitOpError(
4836 "requires memref and vector types of the same elemental type");
4837 return success();
4838}
4839
4840LogicalResult AffineVectorLoadOp::verify() {
4841 MemRefType memrefType = getMemRefType();
4843 *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
4844 getMapOperands(), memrefType,
4845 /*numIndexOperands=*/getNumOperands() - 1)))
4846 return failure();
4847
4848 if (failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType())))
4849 return failure();
4850
4851 return success();
4852}
4853
4854//===----------------------------------------------------------------------===//
4855// AffineVectorStoreOp
4856//===----------------------------------------------------------------------===//
4857
4858void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &result,
4859 Value valueToStore, Value memref, AffineMap map,
4860 ValueRange mapOperands,
4861 llvm::MaybeAlign alignment) {
4862 assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
4863 result.addOperands(valueToStore);
4864 result.addOperands(memref);
4865 result.addOperands(mapOperands);
4866 result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
4867 addAlignmentAttr(builder, result, getAlignmentAttrName(result.name),
4868 alignment);
4869}
4870
4871// Use identity map.
4872void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &result,
4873 Value valueToStore, Value memref,
4875 llvm::MaybeAlign alignment) {
4876 auto memrefType = llvm::cast<MemRefType>(memref.getType());
4877 int64_t rank = memrefType.getRank();
4878 // Create identity map for memrefs with at least one dimension or () -> ()
4879 // for zero-dimensional memrefs.
4880 auto map =
4881 rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap();
4882 build(builder, result, valueToStore, memref, map, indices, alignment);
4883}
4884void AffineVectorStoreOp::getCanonicalizationPatterns(
4885 RewritePatternSet &results, MLIRContext *context) {
4886 results.add<SimplifyAffineOp<AffineVectorStoreOp>>(context);
4887}
4888
4889ParseResult AffineVectorStoreOp::parse(OpAsmParser &parser,
4890 OperationState &result) {
4891 auto indexTy = parser.getBuilder().getIndexType();
4892
4893 MemRefType memrefType;
4894 VectorType resultType;
4895 OpAsmParser::UnresolvedOperand storeValueInfo;
4896 OpAsmParser::UnresolvedOperand memrefInfo;
4897 AffineMapAttr mapAttr;
4898 SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
4899 return failure(
4900 parser.parseOperand(storeValueInfo) || parser.parseComma() ||
4901 parser.parseOperand(memrefInfo) ||
4902 parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
4903 AffineVectorStoreOp::getMapAttrStrName(),
4904 result.attributes) ||
4905 parser.parseOptionalAttrDict(result.attributes) ||
4906 parser.parseColonType(memrefType) || parser.parseComma() ||
4907 parser.parseType(resultType) ||
4908 parser.resolveOperand(storeValueInfo, resultType, result.operands) ||
4909 parser.resolveOperand(memrefInfo, memrefType, result.operands) ||
4910 parser.resolveOperands(mapOperands, indexTy, result.operands));
4911}
4912
4913void AffineVectorStoreOp::print(OpAsmPrinter &p) {
4914 p << " " << getValueToStore();
4915 p << ", " << getMemRef() << '[';
4916 if (AffineMapAttr mapAttr =
4917 (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
4918 p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
4919 p << ']';
4920 p.printOptionalAttrDict((*this)->getAttrs(),
4921 /*elidedAttrs=*/{getMapAttrStrName()});
4922 p << " : " << getMemRefType() << ", " << getValueToStore().getType();
4923}
4924
4925LogicalResult AffineVectorStoreOp::verify() {
4926 MemRefType memrefType = getMemRefType();
4928 *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
4929 getMapOperands(), memrefType,
4930 /*numIndexOperands=*/getNumOperands() - 2)))
4931 return failure();
4932
4933 if (failed(verifyVectorMemoryOp(*this, memrefType, getVectorType())))
4934 return failure();
4935
4936 return success();
4937}
4938
4939//===----------------------------------------------------------------------===//
4940// DelinearizeIndexOp
4941//===----------------------------------------------------------------------===//
4942
4943void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4944 OperationState &odsState,
4945 Value linearIndex, ValueRange dynamicBasis,
4946 ArrayRef<int64_t> staticBasis,
4947 bool hasOuterBound) {
4948 SmallVector<Type> returnTypes(hasOuterBound ? staticBasis.size()
4949 : staticBasis.size() + 1,
4950 linearIndex.getType());
4951 build(odsBuilder, odsState, returnTypes, linearIndex, dynamicBasis,
4952 staticBasis);
4953}
4954
4955void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4956 OperationState &odsState,
4957 Value linearIndex, ValueRange basis,
4958 bool hasOuterBound) {
4959 if (hasOuterBound && !basis.empty() && basis.front() == nullptr) {
4960 hasOuterBound = false;
4961 basis = basis.drop_front();
4962 }
4963 SmallVector<Value> dynamicBasis;
4964 SmallVector<int64_t> staticBasis;
4965 dispatchIndexOpFoldResults(getAsOpFoldResult(basis), dynamicBasis,
4966 staticBasis);
4967 build(odsBuilder, odsState, linearIndex, dynamicBasis, staticBasis,
4968 hasOuterBound);
4969}
4970
4971void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4972 OperationState &odsState,
4973 Value linearIndex,
4974 ArrayRef<OpFoldResult> basis,
4975 bool hasOuterBound) {
4976 if (hasOuterBound && !basis.empty() && basis.front() == OpFoldResult()) {
4977 hasOuterBound = false;
4978 basis = basis.drop_front();
4979 }
4980 SmallVector<Value> dynamicBasis;
4981 SmallVector<int64_t> staticBasis;
4982 dispatchIndexOpFoldResults(basis, dynamicBasis, staticBasis);
4983 build(odsBuilder, odsState, linearIndex, dynamicBasis, staticBasis,
4984 hasOuterBound);
4985}
4986
4987void AffineDelinearizeIndexOp::build(OpBuilder &odsBuilder,
4988 OperationState &odsState,
4989 Value linearIndex, ArrayRef<int64_t> basis,
4990 bool hasOuterBound) {
4991 build(odsBuilder, odsState, linearIndex, ValueRange{}, basis, hasOuterBound);
4992}
4993
4994LogicalResult AffineDelinearizeIndexOp::verify() {
4995 ArrayRef<int64_t> staticBasis = getStaticBasis();
4996 if (getNumResults() != staticBasis.size() &&
4997 getNumResults() != staticBasis.size() + 1)
4998 return emitOpError("should return an index for each basis element and up "
4999 "to one extra index");
5000
5001 auto dynamicMarkersCount = llvm::count_if(staticBasis, ShapedType::isDynamic);
5002 if (static_cast<size_t>(dynamicMarkersCount) != getDynamicBasis().size())
5003 return emitOpError(
5004 "mismatch between dynamic and static basis (kDynamic marker but no "
5005 "corresponding dynamic basis entry) -- this can only happen due to an "
5006 "incorrect fold/rewrite");
5007
5008 if (!llvm::all_of(staticBasis, [](int64_t v) {
5009 return v > 0 || ShapedType::isDynamic(v);
5010 }))
5011 return emitOpError("no basis element may be statically non-positive");
5012
5013 return success();
5014}
5015
5016/// Given mixed basis of affine.delinearize_index/linearize_index replace
5017/// constant SSA values with the constant integer value and return the new
5018/// static basis. In case no such candidate for replacement exists, this utility
5019/// returns std::nullopt.
5020static std::optional<SmallVector<int64_t>>
5022 MutableOperandRange mutableDynamicBasis,
5023 ArrayRef<Attribute> dynamicBasis) {
5024 uint64_t dynamicBasisIndex = 0;
5025 for (Attribute basis : dynamicBasis) {
5026 // Skip poison values: they don't have a concrete integer value, so erasing
5027 // them from the dynamic operands would create an inconsistency between
5028 // the static basis (which would still hold kDynamic) and the dynamic
5029 // operand list (which would be one element shorter).
5030 if (basis && isa<IntegerAttr>(basis)) {
5031 mutableDynamicBasis.erase(dynamicBasisIndex);
5032 } else {
5033 ++dynamicBasisIndex;
5034 }
5035 }
5036
5037 // No constant SSA value exists.
5038 if (dynamicBasisIndex == dynamicBasis.size())
5039 return std::nullopt;
5040
5041 SmallVector<int64_t> staticBasis;
5042 for (OpFoldResult basis : mixedBasis) {
5043 std::optional<int64_t> basisVal = getConstantIntValue(basis);
5044 if (!basisVal)
5045 staticBasis.push_back(ShapedType::kDynamic);
5046 else
5047 staticBasis.push_back(*basisVal);
5048 }
5049
5050 return staticBasis;
5051}
5052
5053LogicalResult
5054AffineDelinearizeIndexOp::fold(FoldAdaptor adaptor,
5055 SmallVectorImpl<OpFoldResult> &result) {
5056 std::optional<SmallVector<int64_t>> maybeStaticBasis =
5057 foldCstValueToCstAttrBasis(getMixedBasis(), getDynamicBasisMutable(),
5058 adaptor.getDynamicBasis());
5059 if (maybeStaticBasis) {
5060 setStaticBasis(*maybeStaticBasis);
5061 return success();
5062 }
5063 // If we won't be doing any division or modulo (no basis or the one basis
5064 // element is purely advisory), simply return the input value.
5065 if (getNumResults() == 1) {
5066 result.push_back(getLinearIndex());
5067 return success();
5068 }
5069
5070 if (adaptor.getLinearIndex() == nullptr)
5071 return failure();
5072
5073 if (!adaptor.getDynamicBasis().empty())
5074 return failure();
5075
5076 int64_t highPart = cast<IntegerAttr>(adaptor.getLinearIndex()).getInt();
5077 Type attrType = getLinearIndex().getType();
5078
5079 ArrayRef<int64_t> staticBasis = getStaticBasis();
5080 if (hasOuterBound())
5081 staticBasis = staticBasis.drop_front();
5082 for (int64_t modulus : llvm::reverse(staticBasis)) {
5083 result.push_back(IntegerAttr::get(attrType, llvm::mod(highPart, modulus)));
5084 highPart = llvm::divideFloorSigned(highPart, modulus);
5085 }
5086 result.push_back(IntegerAttr::get(attrType, highPart));
5087 std::reverse(result.begin(), result.end());
5088 return success();
5089}
5090
5091SmallVector<OpFoldResult> AffineDelinearizeIndexOp::getEffectiveBasis() {
5092 OpBuilder builder(getContext());
5093 if (hasOuterBound()) {
5094 if (getStaticBasis().front() == ::mlir::ShapedType::kDynamic)
5095 return getMixedValues(getStaticBasis().drop_front(),
5096 getDynamicBasis().drop_front(), builder);
5097
5098 return getMixedValues(getStaticBasis().drop_front(), getDynamicBasis(),
5099 builder);
5100 }
5101
5102 return getMixedValues(getStaticBasis(), getDynamicBasis(), builder);
5103}
5104
5105SmallVector<OpFoldResult> AffineDelinearizeIndexOp::getPaddedBasis() {
5106 SmallVector<OpFoldResult> ret = getMixedBasis();
5107 if (!hasOuterBound())
5108 ret.insert(ret.begin(), OpFoldResult());
5109 return ret;
5110}
5111
5112namespace {
5113
5114// Drops delinearization indices that correspond to unit-extent basis
5115struct DropUnitExtentBasis
5116 : public OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5118
5119 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5120 PatternRewriter &rewriter) const override {
5121 SmallVector<Value> replacements(delinearizeOp->getNumResults(), nullptr);
5122 std::optional<Value> zero = std::nullopt;
5123 Location loc = delinearizeOp->getLoc();
5124 Type indexType = delinearizeOp.getLinearIndex().getType();
5125 auto getZero = [&]() -> Value {
5126 if (!zero)
5127 zero = arith::ConstantOp::create(rewriter, loc,
5128 rewriter.getZeroAttr(indexType));
5129 return zero.value();
5130 };
5131
5132 // Replace all indices corresponding to unit-extent basis with 0.
5133 // Remaining basis can be used to get a new `affine.delinearize_index` op.
5134 SmallVector<OpFoldResult> newBasis;
5135 for (auto [index, basis] :
5136 llvm::enumerate(delinearizeOp.getPaddedBasis())) {
5137 std::optional<int64_t> basisVal =
5138 basis ? getConstantIntValue(basis) : std::nullopt;
5139 if (basisVal == 1)
5140 replacements[index] = getZero();
5141 else
5142 newBasis.push_back(basis);
5143 }
5144
5145 if (newBasis.size() == delinearizeOp.getNumResults())
5146 return rewriter.notifyMatchFailure(delinearizeOp,
5147 "no unit basis elements");
5148
5149 if (!newBasis.empty()) {
5150 // Will drop the leading nullptr from `basis` if there was no outer bound.
5151 auto newDelinearizeOp = affine::AffineDelinearizeIndexOp::create(
5152 rewriter, loc, delinearizeOp.getLinearIndex(), newBasis);
5153 int newIndex = 0;
5154 // Map back the new delinearized indices to the values they replace.
5155 for (auto &replacement : replacements) {
5156 if (replacement)
5157 continue;
5158 replacement = newDelinearizeOp->getResult(newIndex++);
5159 }
5160 }
5161
5162 rewriter.replaceOp(delinearizeOp, replacements);
5163 return success();
5164 }
5165};
5166
5167/// If a `affine.delinearize_index`'s input is a `affine.linearize_index
5168/// disjoint` and the two operations end with the same basis elements,
5169/// cancel those parts of the operations out because they are inverses
5170/// of each other.
5171///
5172/// If the operations have the same basis, cancel them entirely.
5173///
5174/// The `disjoint` flag is needed on the `affine.linearize_index` because
5175/// otherwise, there is no guarantee that the inputs to the linearization are
5176/// in-bounds the way the outputs of the delinearization would be.
5177struct CancelDelinearizeOfLinearizeDisjointExactTail
5178 : public OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5180
5181 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5182 PatternRewriter &rewriter) const override {
5183 auto linearizeOp = delinearizeOp.getLinearIndex()
5184 .getDefiningOp<affine::AffineLinearizeIndexOp>();
5185 if (!linearizeOp)
5186 return rewriter.notifyMatchFailure(delinearizeOp,
5187 "index doesn't come from linearize");
5188
5189 if (!linearizeOp.getDisjoint())
5190 return rewriter.notifyMatchFailure(linearizeOp, "not disjoint");
5191
5192 ValueRange linearizeIns = linearizeOp.getMultiIndex();
5193 // Note: we use the full basis so we don't lose outer bounds later.
5194 SmallVector<OpFoldResult> linearizeBasis = linearizeOp.getMixedBasis();
5195 SmallVector<OpFoldResult> delinearizeBasis = delinearizeOp.getMixedBasis();
5196 size_t numMatches = 0;
5197 for (auto [linSize, delinSize] : llvm::zip(
5198 llvm::reverse(linearizeBasis), llvm::reverse(delinearizeBasis))) {
5199 if (linSize != delinSize)
5200 break;
5201 ++numMatches;
5202 }
5203
5204 if (numMatches == 0)
5205 return rewriter.notifyMatchFailure(
5206 delinearizeOp, "final basis element doesn't match linearize");
5207
5208 // The easy case: everything lines up and the basis match sup completely.
5209 if (numMatches == linearizeBasis.size() &&
5210 numMatches == delinearizeBasis.size() &&
5211 linearizeIns.size() == delinearizeOp.getNumResults()) {
5212 rewriter.replaceOp(delinearizeOp, linearizeOp.getMultiIndex());
5213 return success();
5214 }
5215
5216 Value newLinearize = affine::AffineLinearizeIndexOp::create(
5217 rewriter, linearizeOp.getLoc(), linearizeIns.drop_back(numMatches),
5218 ArrayRef<OpFoldResult>{linearizeBasis}.drop_back(numMatches),
5219 linearizeOp.getDisjoint());
5220 auto newDelinearize = affine::AffineDelinearizeIndexOp::create(
5221 rewriter, delinearizeOp.getLoc(), newLinearize,
5222 ArrayRef<OpFoldResult>{delinearizeBasis}.drop_back(numMatches),
5223 delinearizeOp.hasOuterBound());
5224 SmallVector<Value> mergedResults(newDelinearize.getResults());
5225 mergedResults.append(linearizeIns.take_back(numMatches).begin(),
5226 linearizeIns.take_back(numMatches).end());
5227 rewriter.replaceOp(delinearizeOp, mergedResults);
5228 return success();
5229 }
5230};
5231
5232/// If the input to a delinearization is a disjoint linearization, and the
5233/// last k > 1 components of the delinearization basis multiply to the
5234/// last component of the linearization basis, break the linearization and
5235/// delinearization into two parts, peeling off the last input to linearization.
5236///
5237/// For example:
5238/// %0 = affine.linearize_index [%z, %y, %x] by (3, 2, 32) : index
5239/// %1:4 = affine.delinearize_index %0 by (2, 3, 8, 4) : index, ...
5240/// becomes
5241/// %0 = affine.linearize_index [%z, %y] by (3, 2) : index
5242/// %1:2 = affine.delinearize_index %0 by (2, 3) : index
5243/// %2:2 = affine.delinearize_index %x by (8, 4) : index
5244/// where the original %1:4 is replaced by %1:2 ++ %2:2
5245struct SplitDelinearizeSpanningLastLinearizeArg final
5246 : OpRewritePattern<affine::AffineDelinearizeIndexOp> {
5248
5249 LogicalResult matchAndRewrite(affine::AffineDelinearizeIndexOp delinearizeOp,
5250 PatternRewriter &rewriter) const override {
5251 auto linearizeOp = delinearizeOp.getLinearIndex()
5252 .getDefiningOp<affine::AffineLinearizeIndexOp>();
5253 if (!linearizeOp)
5254 return rewriter.notifyMatchFailure(delinearizeOp,
5255 "index doesn't come from linearize");
5256
5257 if (!linearizeOp.getDisjoint())
5258 return rewriter.notifyMatchFailure(linearizeOp,
5259 "linearize isn't disjoint");
5260
5261 // A linearize with no inputs has an empty basis and folds to a constant
5262 // zero; there is nothing to split, and reading its last basis element
5263 // below would be out of bounds.
5264 if (linearizeOp.getStaticBasis().empty())
5265 return rewriter.notifyMatchFailure(
5266 linearizeOp, "linearize has no basis elements (no inputs)");
5267
5268 int64_t target = linearizeOp.getStaticBasis().back();
5269 if (ShapedType::isDynamic(target))
5270 return rewriter.notifyMatchFailure(
5271 linearizeOp, "linearize ends with dynamic basis value");
5272
5273 int64_t sizeToSplit = 1;
5274 size_t elemsToSplit = 0;
5275 ArrayRef<int64_t> basis = delinearizeOp.getStaticBasis();
5276 for (int64_t basisElem : llvm::reverse(basis)) {
5277 if (ShapedType::isDynamic(basisElem))
5278 return rewriter.notifyMatchFailure(
5279 delinearizeOp, "dynamic basis element while scanning for split");
5280 sizeToSplit *= basisElem;
5281 elemsToSplit += 1;
5282
5283 if (sizeToSplit > target)
5284 return rewriter.notifyMatchFailure(delinearizeOp,
5285 "overshot last argument size");
5286 if (sizeToSplit == target)
5287 break;
5288 }
5289
5290 if (sizeToSplit < target)
5291 return rewriter.notifyMatchFailure(
5292 delinearizeOp, "product of known basis elements doesn't exceed last "
5293 "linearize argument");
5294
5295 if (elemsToSplit < 2)
5296 return rewriter.notifyMatchFailure(
5297 delinearizeOp,
5298 "need at least two elements to form the basis product");
5299
5300 Value linearizeWithoutBack = affine::AffineLinearizeIndexOp::create(
5301 rewriter, linearizeOp.getLoc(), linearizeOp.getLinearIndex().getType(),
5302 linearizeOp.getMultiIndex().drop_back(), linearizeOp.getDynamicBasis(),
5303 linearizeOp.getStaticBasis().drop_back(), linearizeOp.getDisjoint());
5304 auto delinearizeWithoutSplitPart = affine::AffineDelinearizeIndexOp::create(
5305 rewriter, delinearizeOp.getLoc(), linearizeWithoutBack,
5306 delinearizeOp.getDynamicBasis(), basis.drop_back(elemsToSplit),
5307 delinearizeOp.hasOuterBound());
5308 auto delinearizeBack = affine::AffineDelinearizeIndexOp::create(
5309 rewriter, delinearizeOp.getLoc(), linearizeOp.getMultiIndex().back(),
5310 basis.take_back(elemsToSplit), /*hasOuterBound=*/true);
5311 SmallVector<Value> results = llvm::to_vector(
5312 llvm::concat<Value>(delinearizeWithoutSplitPart.getResults(),
5313 delinearizeBack.getResults()));
5314 rewriter.replaceOp(delinearizeOp, results);
5315
5316 return success();
5317 }
5318};
5319} // namespace
5320
5321void affine::AffineDelinearizeIndexOp::getCanonicalizationPatterns(
5322 RewritePatternSet &patterns, MLIRContext *context) {
5323 patterns
5324 .insert<CancelDelinearizeOfLinearizeDisjointExactTail,
5325 DropUnitExtentBasis, SplitDelinearizeSpanningLastLinearizeArg>(
5326 context);
5327}
5328
5329//===----------------------------------------------------------------------===//
5330// LinearizeIndexOp
5331//===----------------------------------------------------------------------===//
5332
5333/// Infer the index type from a set of multi-index values. Returns the common
5334/// type (index or vector<...xindex>), or IndexType if the set is empty.
5335static Type inferIndexType(MLIRContext *ctx, ValueRange multiIndex) {
5336 if (multiIndex.empty())
5337 return IndexType::get(ctx);
5338 return multiIndex.front().getType();
5339}
5340
5341void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5342 OperationState &odsState,
5343 ValueRange multiIndex, ValueRange basis,
5344 bool disjoint) {
5345 if (!basis.empty() && basis.front() == Value())
5346 basis = basis.drop_front();
5347 SmallVector<Value> dynamicBasis;
5348 SmallVector<int64_t> staticBasis;
5349 dispatchIndexOpFoldResults(getAsOpFoldResult(basis), dynamicBasis,
5350 staticBasis);
5351 Type resultType = inferIndexType(odsBuilder.getContext(), multiIndex);
5352 build(odsBuilder, odsState, resultType, multiIndex, dynamicBasis, staticBasis,
5353 disjoint);
5354}
5355
5356void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5357 OperationState &odsState,
5358 ValueRange multiIndex,
5359 ArrayRef<OpFoldResult> basis,
5360 bool disjoint) {
5361 if (!basis.empty() && basis.front() == OpFoldResult())
5362 basis = basis.drop_front();
5363 SmallVector<Value> dynamicBasis;
5364 SmallVector<int64_t> staticBasis;
5365 dispatchIndexOpFoldResults(basis, dynamicBasis, staticBasis);
5366 Type resultType = inferIndexType(odsBuilder.getContext(), multiIndex);
5367 build(odsBuilder, odsState, resultType, multiIndex, dynamicBasis, staticBasis,
5368 disjoint);
5369}
5370
5371void AffineLinearizeIndexOp::build(OpBuilder &odsBuilder,
5372 OperationState &odsState,
5373 ValueRange multiIndex,
5374 ArrayRef<int64_t> basis, bool disjoint) {
5375 Type resultType = inferIndexType(odsBuilder.getContext(), multiIndex);
5376 build(odsBuilder, odsState, resultType, multiIndex, ValueRange{}, basis,
5377 disjoint);
5378}
5379
5380LogicalResult AffineLinearizeIndexOp::verify() {
5381 size_t numIndexes = getMultiIndex().size();
5382 size_t numBasisElems = getStaticBasis().size();
5383 if (numIndexes != numBasisElems && numIndexes != numBasisElems + 1)
5384 return emitOpError("should be passed a basis element for each index except "
5385 "possibly the first");
5386
5387 auto dynamicMarkersCount =
5388 llvm::count_if(getStaticBasis(), ShapedType::isDynamic);
5389 if (static_cast<size_t>(dynamicMarkersCount) != getDynamicBasis().size())
5390 return emitOpError(
5391 "mismatch between dynamic and static basis (kDynamic marker but no "
5392 "corresponding dynamic basis entry) -- this can only happen due to an "
5393 "incorrect fold/rewrite");
5394
5395 return success();
5396}
5397
5398OpFoldResult AffineLinearizeIndexOp::fold(FoldAdaptor adaptor) {
5399 std::optional<SmallVector<int64_t>> maybeStaticBasis =
5400 foldCstValueToCstAttrBasis(getMixedBasis(), getDynamicBasisMutable(),
5401 adaptor.getDynamicBasis());
5402 if (maybeStaticBasis) {
5403 setStaticBasis(*maybeStaticBasis);
5404 return getResult();
5405 }
5406 // No indices linearizes to zero.
5407 if (getMultiIndex().empty())
5408 return IntegerAttr::get(getResult().getType(), 0);
5409
5410 // One single index linearizes to itself.
5411 if (getMultiIndex().size() == 1)
5412 return getMultiIndex().front();
5413
5414 // Return nullptr if any multi-index attribute has not been folded to a
5415 // concrete integer (e.g. it is still a runtime value or has folded to a
5416 // non-integer attribute such as #ub.poison).
5417 if (llvm::any_of(adaptor.getMultiIndex(), [](Attribute a) {
5418 return !isa_and_nonnull<IntegerAttr>(a);
5419 }))
5420 return nullptr;
5421
5422 if (!adaptor.getDynamicBasis().empty())
5423 return nullptr;
5424
5425 int64_t result = 0;
5426 int64_t stride = 1;
5427 for (auto [length, indexAttr] :
5428 llvm::zip_first(llvm::reverse(getStaticBasis()),
5429 llvm::reverse(adaptor.getMultiIndex()))) {
5430 result = result + cast<IntegerAttr>(indexAttr).getInt() * stride;
5431 stride = stride * length;
5432 }
5433 // Handle the index element with no basis element.
5434 if (!hasOuterBound())
5435 result =
5436 result +
5437 cast<IntegerAttr>(adaptor.getMultiIndex().front()).getInt() * stride;
5438
5439 return IntegerAttr::get(getResult().getType(), result);
5440}
5441
5442SmallVector<OpFoldResult> AffineLinearizeIndexOp::getEffectiveBasis() {
5443 OpBuilder builder(getContext());
5444 if (hasOuterBound()) {
5445 if (getStaticBasis().front() == ::mlir::ShapedType::kDynamic)
5446 return getMixedValues(getStaticBasis().drop_front(),
5447 getDynamicBasis().drop_front(), builder);
5448
5449 return getMixedValues(getStaticBasis().drop_front(), getDynamicBasis(),
5450 builder);
5451 }
5452
5453 return getMixedValues(getStaticBasis(), getDynamicBasis(), builder);
5454}
5455
5456SmallVector<OpFoldResult> AffineLinearizeIndexOp::getPaddedBasis() {
5457 SmallVector<OpFoldResult> ret = getMixedBasis();
5458 if (!hasOuterBound())
5459 ret.insert(ret.begin(), OpFoldResult());
5460 return ret;
5461}
5462
5463namespace {
5464/// Rewrite `affine.linearize_index disjoint [%...a, %x, %...b] by (%...c, 1,
5465/// %...d)` to `affine.linearize_index disjoint [%...a, %...b] by (%...c,
5466/// %...d)`.
5467
5468/// Note that `disjoint` is required here, because, without it, we could have
5469/// `affine.linearize_index [%...a, %c64, %...b] by (%...c, 1, %...d)`
5470/// is a valid operation where the `%c64` cannot be trivially dropped.
5471///
5472/// Alternatively, if `%x` in the above is a known constant 0, remove it even if
5473/// the operation isn't asserted to be `disjoint`.
5474struct DropLinearizeUnitComponentsIfDisjointOrZero final
5475 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5477
5478 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp op,
5479 PatternRewriter &rewriter) const override {
5480 ValueRange multiIndex = op.getMultiIndex();
5481 size_t numIndices = multiIndex.size();
5482 SmallVector<Value> newIndices;
5483 newIndices.reserve(numIndices);
5484 SmallVector<OpFoldResult> newBasis;
5485 newBasis.reserve(numIndices);
5486
5487 if (!op.hasOuterBound()) {
5488 newIndices.push_back(multiIndex.front());
5489 multiIndex = multiIndex.drop_front();
5490 }
5491
5492 SmallVector<OpFoldResult> basis = op.getMixedBasis();
5493 for (auto [index, basisElem] : llvm::zip_equal(multiIndex, basis)) {
5494 std::optional<int64_t> basisEntry = getConstantIntValue(basisElem);
5495 if (!basisEntry || *basisEntry != 1) {
5496 newIndices.push_back(index);
5497 newBasis.push_back(basisElem);
5498 continue;
5499 }
5500
5501 std::optional<int64_t> indexValue = getConstantIntValue(index);
5502 if (!op.getDisjoint() && (!indexValue || *indexValue != 0)) {
5503 newIndices.push_back(index);
5504 newBasis.push_back(basisElem);
5505 continue;
5506 }
5507 }
5508 if (newIndices.size() == numIndices)
5509 return rewriter.notifyMatchFailure(op,
5510 "no unit basis entries to replace");
5511
5512 if (newIndices.empty()) {
5513 rewriter.replaceOpWithNewOp<arith::ConstantOp>(
5514 op, rewriter.getZeroAttr(op.getLinearIndex().getType()));
5515 return success();
5516 }
5517 rewriter.replaceOpWithNewOp<affine::AffineLinearizeIndexOp>(
5518 op, newIndices, newBasis, op.getDisjoint());
5519 return success();
5520 }
5521};
5522
5523OpFoldResult computeProduct(Location loc, OpBuilder &builder,
5524 ArrayRef<OpFoldResult> terms) {
5525 int64_t nDynamic = 0;
5526 SmallVector<Value> dynamicPart;
5527 AffineExpr result = builder.getAffineConstantExpr(1);
5528 for (OpFoldResult term : terms) {
5529 if (!term)
5530 return term;
5531 std::optional<int64_t> maybeConst = getConstantIntValue(term);
5532 if (maybeConst) {
5533 result = result * builder.getAffineConstantExpr(*maybeConst);
5534 } else {
5535 dynamicPart.push_back(cast<Value>(term));
5536 result = result * builder.getAffineSymbolExpr(nDynamic++);
5537 }
5538 }
5539 if (auto constant = dyn_cast<AffineConstantExpr>(result))
5540 return getAsIndexOpFoldResult(builder.getContext(), constant.getValue());
5541 return AffineApplyOp::create(builder, loc, result, dynamicPart).getResult();
5542}
5543
5544/// If conseceutive outputs of a delinearize_index are linearized with the same
5545/// bounds, canonicalize away the redundant arithmetic.
5546///
5547/// That is, if we have
5548/// ```
5549/// %s:N = affine.delinearize_index %x into (...a, B1, B2, ... BK, ...b)
5550/// %t = affine.linearize_index [...c, %s#I, %s#(I + 1), ... %s#(I+K-1), ...d]
5551/// by (...e, B1, B2, ..., BK, ...f)
5552/// ```
5553///
5554/// We can rewrite this to
5555/// ```
5556/// B = B1 * B2 ... BK
5557/// %sMerged:(N-K+1) affine.delinearize_index %x into (...a, B, ...b)
5558/// %t = affine.linearize_index [...c, %s#I, ...d] by (...e, B, ...f)
5559/// ```
5560/// where we replace all results of %s unaffected by the change with results
5561/// from %sMerged.
5562///
5563/// As a special case, if all results of the delinearize are merged in this way
5564/// we can replace those usages with %x, thus cancelling the delinearization
5565/// entirely, as in
5566/// ```
5567/// %s:3 = affine.delinearize_index %x into (2, 4, 8)
5568/// %t = affine.linearize_index [%s#0, %s#1, %s#2, %c0] by (2, 4, 8, 16)
5569/// ```
5570/// becoming `%t = affine.linearize_index [%x, %c0] by (64, 16)`
5571struct CancelLinearizeOfDelinearizePortion final
5572 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5574
5575private:
5576 // Struct representing a case where the cancellation pattern
5577 // applies. A `Match` means that `length` inputs to the linearize operation
5578 // starting at `linStart` can be cancelled with `length` outputs of
5579 // `delinearize`, starting from `delinStart`.
5580 struct Match {
5581 AffineDelinearizeIndexOp delinearize;
5582 unsigned linStart = 0;
5583 unsigned delinStart = 0;
5584 unsigned length = 0;
5585 };
5586
5587public:
5588 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp linearizeOp,
5589 PatternRewriter &rewriter) const override {
5590 SmallVector<Match> matches;
5591
5592 const SmallVector<OpFoldResult> linBasis = linearizeOp.getPaddedBasis();
5593 ArrayRef<OpFoldResult> linBasisRef = linBasis;
5594
5595 ValueRange multiIndex = linearizeOp.getMultiIndex();
5596 unsigned numLinArgs = multiIndex.size();
5597 unsigned linArgIdx = 0;
5598 // We only want to replace one run from the same delinearize op per
5599 // pattern invocation lest we run into invalidation issues.
5600 llvm::SmallPtrSet<Operation *, 2> alreadyMatchedDelinearize;
5601 while (linArgIdx < numLinArgs) {
5602 auto asResult = dyn_cast<OpResult>(multiIndex[linArgIdx]);
5603 if (!asResult) {
5604 linArgIdx++;
5605 continue;
5606 }
5607
5608 auto delinearizeOp =
5609 dyn_cast<AffineDelinearizeIndexOp>(asResult.getOwner());
5610 if (!delinearizeOp) {
5611 linArgIdx++;
5612 continue;
5613 }
5614
5615 /// Result 0 of the delinearize and argument 0 of the linearize can
5616 /// leave their maximum value unspecified. However, even if this happens
5617 /// we can still sometimes start the match process. Specifically, if
5618 /// - The argument we're matching is result 0 and argument 0 (so the
5619 /// bounds don't matter). For example,
5620 ///
5621 /// %0:2 = affine.delinearize_index %x into (8) : index, index
5622 /// %1 = affine.linearize_index [%s#0, %s#1, ...] (8, ...)
5623 /// allows cancellation
5624 /// - The delinearization doesn't specify a bound, but the linearization
5625 /// is `disjoint`, which asserts that the bound on the linearization is
5626 /// correct.
5627 unsigned delinArgIdx = asResult.getResultNumber();
5628 SmallVector<OpFoldResult> delinBasis = delinearizeOp.getPaddedBasis();
5629 OpFoldResult firstDelinBound = delinBasis[delinArgIdx];
5630 OpFoldResult firstLinBound = linBasis[linArgIdx];
5631 bool boundsMatch = firstDelinBound == firstLinBound;
5632 bool bothAtFront = linArgIdx == 0 && delinArgIdx == 0;
5633 bool knownByDisjoint =
5634 linearizeOp.getDisjoint() && delinArgIdx == 0 && !firstDelinBound;
5635 if (!boundsMatch && !bothAtFront && !knownByDisjoint) {
5636 linArgIdx++;
5637 continue;
5638 }
5639
5640 unsigned j = 1;
5641 unsigned numDelinOuts = delinearizeOp.getNumResults();
5642 for (; j + linArgIdx < numLinArgs && j + delinArgIdx < numDelinOuts;
5643 ++j) {
5644 if (multiIndex[linArgIdx + j] !=
5645 delinearizeOp.getResult(delinArgIdx + j))
5646 break;
5647 if (linBasis[linArgIdx + j] != delinBasis[delinArgIdx + j])
5648 break;
5649 }
5650 // If there're multiple matches against the same delinearize_index,
5651 // only rewrite the first one we find to prevent invalidations. The next
5652 // ones will be taken care of by subsequent pattern invocations.
5653 if (j <= 1 || !alreadyMatchedDelinearize.insert(delinearizeOp).second) {
5654 linArgIdx++;
5655 continue;
5656 }
5657 matches.push_back(Match{delinearizeOp, linArgIdx, delinArgIdx, j});
5658 linArgIdx += j;
5659 }
5660
5661 if (matches.empty())
5662 return rewriter.notifyMatchFailure(
5663 linearizeOp, "no run of delinearize outputs to deal with");
5664
5665 // Record all the delinearize replacements so we can do them after creating
5666 // the new linearization operation, since the new operation might use
5667 // outputs of something we're replacing.
5668 SmallVector<SmallVector<Value>> delinearizeReplacements;
5669
5670 SmallVector<Value> newIndex;
5671 newIndex.reserve(numLinArgs);
5672 SmallVector<OpFoldResult> newBasis;
5673 newBasis.reserve(numLinArgs);
5674 unsigned prevMatchEnd = 0;
5675 for (Match m : matches) {
5676 unsigned gap = m.linStart - prevMatchEnd;
5677 llvm::append_range(newIndex, multiIndex.slice(prevMatchEnd, gap));
5678 llvm::append_range(newBasis, linBasisRef.slice(prevMatchEnd, gap));
5679 // Update here so we don't forget this during early continues
5680 prevMatchEnd = m.linStart + m.length;
5681
5682 PatternRewriter::InsertionGuard g(rewriter);
5683 rewriter.setInsertionPoint(m.delinearize);
5684
5685 ArrayRef<OpFoldResult> basisToMerge =
5686 linBasisRef.slice(m.linStart, m.length);
5687 // We use the slice from the linearize's basis above because of the
5688 // "bounds inferred from `disjoint`" case above.
5689 OpFoldResult newSize =
5690 computeProduct(linearizeOp.getLoc(), rewriter, basisToMerge);
5691
5692 // Trivial case where we can just skip past the delinearize all together
5693 if (m.length == m.delinearize.getNumResults()) {
5694 newIndex.push_back(m.delinearize.getLinearIndex());
5695 newBasis.push_back(newSize);
5696 // Pad out set of replacements so we don't do anything with this one.
5697 delinearizeReplacements.push_back(SmallVector<Value>());
5698 continue;
5699 }
5700
5701 SmallVector<Value> newDelinResults;
5702 SmallVector<OpFoldResult> newDelinBasis = m.delinearize.getPaddedBasis();
5703 newDelinBasis.erase(newDelinBasis.begin() + m.delinStart,
5704 newDelinBasis.begin() + m.delinStart + m.length);
5705 newDelinBasis.insert(newDelinBasis.begin() + m.delinStart, newSize);
5706 auto newDelinearize = AffineDelinearizeIndexOp::create(
5707 rewriter, m.delinearize.getLoc(), m.delinearize.getLinearIndex(),
5708 newDelinBasis);
5709
5710 // Since there may be other uses of the indices we just merged together,
5711 // create a residual affine.delinearize_index that delinearizes the
5712 // merged output into its component parts.
5713 Value combinedElem = newDelinearize.getResult(m.delinStart);
5714 auto residualDelinearize = AffineDelinearizeIndexOp::create(
5715 rewriter, m.delinearize.getLoc(), combinedElem, basisToMerge);
5716
5717 // Swap all the uses of the unaffected delinearize outputs to the new
5718 // delinearization so that the old code can be removed if this
5719 // linearize_index is the only user of the merged results.
5720 llvm::append_range(newDelinResults,
5721 newDelinearize.getResults().take_front(m.delinStart));
5722 llvm::append_range(newDelinResults, residualDelinearize.getResults());
5723 llvm::append_range(
5724 newDelinResults,
5725 newDelinearize.getResults().drop_front(m.delinStart + 1));
5726
5727 delinearizeReplacements.push_back(newDelinResults);
5728 newIndex.push_back(combinedElem);
5729 newBasis.push_back(newSize);
5730 }
5731 llvm::append_range(newIndex, multiIndex.drop_front(prevMatchEnd));
5732 llvm::append_range(newBasis, linBasisRef.drop_front(prevMatchEnd));
5733 rewriter.replaceOpWithNewOp<AffineLinearizeIndexOp>(
5734 linearizeOp, newIndex, newBasis, linearizeOp.getDisjoint());
5735
5736 for (auto [m, newResults] :
5737 llvm::zip_equal(matches, delinearizeReplacements)) {
5738 if (newResults.empty())
5739 continue;
5740 rewriter.replaceOp(m.delinearize, newResults);
5741 }
5742
5743 return success();
5744 }
5745};
5746
5747/// Strip leading zero from affine.linearize_index.
5748///
5749/// `affine.linearize_index [%c0, ...a] by (%x, ...b)` can be rewritten
5750/// to `affine.linearize_index [...a] by (...b)` in all cases.
5751struct DropLinearizeLeadingZero final
5752 : OpRewritePattern<affine::AffineLinearizeIndexOp> {
5754
5755 LogicalResult matchAndRewrite(affine::AffineLinearizeIndexOp op,
5756 PatternRewriter &rewriter) const override {
5757 Value leadingIdx = op.getMultiIndex().front();
5758 if (!matchPattern(leadingIdx, m_Zero()))
5759 return failure();
5760
5761 if (op.getMultiIndex().size() == 1) {
5762 rewriter.replaceOp(op, leadingIdx);
5763 return success();
5764 }
5765
5766 SmallVector<OpFoldResult> mixedBasis = op.getMixedBasis();
5767 ArrayRef<OpFoldResult> newMixedBasis = mixedBasis;
5768 if (op.hasOuterBound())
5769 newMixedBasis = newMixedBasis.drop_front();
5770
5771 rewriter.replaceOpWithNewOp<affine::AffineLinearizeIndexOp>(
5772 op, op.getMultiIndex().drop_front(), newMixedBasis, op.getDisjoint());
5773 return success();
5774 }
5775};
5776} // namespace
5777
5778void affine::AffineLinearizeIndexOp::getCanonicalizationPatterns(
5779 RewritePatternSet &patterns, MLIRContext *context) {
5780 patterns.add<CancelLinearizeOfDelinearizePortion, DropLinearizeLeadingZero,
5781 DropLinearizeUnitComponentsIfDisjointOrZero>(context);
5782}
5783
5784//===----------------------------------------------------------------------===//
5785// TableGen'd op method definitions
5786//===----------------------------------------------------------------------===//
5787
5788#define GET_OP_CLASSES
5789#include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc"
return success()
static AffineForOp buildAffineLoopFromConstants(OpBuilder &builder, Location loc, int64_t lb, int64_t ub, int64_t step, AffineForOp::BodyBuilderFn bodyBuilderFn)
Creates an affine loop from the bounds known to be constants.
static bool hasTrivialZeroTripCount(AffineForOp op)
Returns true if the affine.for has zero iterations in trivial cases.
static Type inferIndexType(MLIRContext *ctx, ValueRange multiIndex)
Infer the index type from a set of multi-index values. Returns the common type (index or vector<....
static LogicalResult verifyMemoryOpIndexing(AffineMemOpTy op, AffineMapAttr mapAttr, Operation::operand_range mapOperands, MemRefType memrefType, unsigned numIndexOperands)
Verify common indexing invariants of affine.load, affine.store, affine.vector_load and affine....
static void printAffineMinMaxOp(OpAsmPrinter &p, T op)
static bool isResultTypeMatchAtomicRMWKind(Type resultType, arith::AtomicRMWKind op)
static bool remainsLegalAfterInline(Value value, Region *src, Region *dest, const IRMapping &mapping, function_ref< bool(Value, Region *)> legalityCheck)
Checks if value known to be a legal affine dimension or symbol in src region remains legal if the ope...
Definition AffineOps.cpp:62
static void printMinMaxBound(OpAsmPrinter &p, AffineMapAttr mapAttr, DenseIntElementsAttr group, ValueRange operands, StringRef keyword)
Prints a lower(upper) bound of an affine parallel loop with max(min) conditions in it.
static OpFoldResult foldMinMaxOp(T op, ArrayRef< Attribute > operands)
Fold an affine min or max operation with the given operands.
static bool isTopLevelValueOrAbove(Value value, Region *region)
A utility function to check if a value is defined at the top level of region or is an argument of reg...
static LogicalResult canonicalizeLoopBounds(AffineForOp forOp)
Canonicalize the bounds of the given loop.
static void simplifyExprAndOperands(AffineExpr &expr, unsigned numDims, unsigned numSymbols, ArrayRef< Value > operands)
Simplify expr while exploiting information from the values in operands.
static bool isValidAffineIndexOperand(Value value, Region *region)
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static void canonicalizeMapOrSetAndOperands(MapOrSet *mapOrSet, SmallVectorImpl< Value > *operands)
static ParseResult parseBound(bool isLower, OperationState &result, OpAsmParser &p)
Parse a for operation loop bounds.
static std::optional< SmallVector< int64_t > > foldCstValueToCstAttrBasis(ArrayRef< OpFoldResult > mixedBasis, MutableOperandRange mutableDynamicBasis, ArrayRef< Attribute > dynamicBasis)
Given mixed basis of affine.delinearize_index/linearize_index replace constant SSA values with the co...
static void canonicalizePromotedSymbols(MapOrSet *mapOrSet, SmallVectorImpl< Value > *operands)
static void simplifyMinOrMaxExprWithOperands(AffineMap &map, ArrayRef< Value > operands, bool isMax)
Simplify the expressions in map while making use of lower or upper bounds of its operands.
static ParseResult parseAffineMinMaxOp(OpAsmParser &parser, OperationState &result)
static LogicalResult replaceAffineDelinearizeIndexInverseExpression(AffineDelinearizeIndexOp delinOp, Value resultToReplace, AffineMap *map, SmallVectorImpl< Value > &dims, SmallVectorImpl< Value > &syms)
If this map contains of the expression x_1 + x_1 * C_1 + ... x_n * C_N + / ... (not necessarily in or...
static void composeSetAndOperands(IntegerSet &set, SmallVectorImpl< Value > &operands, bool composeAffineMin=false)
Compose any affine.apply ops feeding into operands of the integer set set by composing the maps of su...
static bool isMemRefSizeValidSymbol(AnyMemRefDefOp memrefDefOp, unsigned index, Region *region)
Returns true if the 'index' dimension of the memref defined by memrefDefOp is a statically shaped one...
static bool isNonNegativeBoundedBy(AffineExpr e, ArrayRef< Value > operands, int64_t k)
Check if e is known to be: 0 <= e < k.
static AffineForOp buildAffineLoopFromValues(OpBuilder &builder, Location loc, Value lb, Value ub, int64_t step, AffineForOp::BodyBuilderFn bodyBuilderFn)
Creates an affine loop from the bounds that may or may not be constants.
static void simplifyMapWithOperands(AffineMap &map, ArrayRef< Value > operands)
Simplify the map while exploiting information on the values in operands.
static void printDimAndSymbolList(Operation::operand_iterator begin, Operation::operand_iterator end, unsigned numDims, OpAsmPrinter &printer)
Prints dimension and symbol list.
static int64_t getLargestKnownDivisor(AffineExpr e, ArrayRef< Value > operands)
Returns the largest known divisor of e.
static void composeAffineMapAndOperands(AffineMap *map, SmallVectorImpl< Value > *operands, bool composeAffineMin=false)
Iterate over operands and fold away all those produced by an AffineApplyOp iteratively.
static void legalizeDemotedDims(MapOrSet &mapOrSet, SmallVectorImpl< Value > &operands)
A valid affine dimension may appear as a symbol in affine.apply operations.
static OpTy makeComposedMinMax(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
static std::optional< int64_t > getUpperBound(Value iv)
Gets the constant upper bound on an affine.for iv.
static void buildAffineLoopNestImpl(OpBuilder &builder, Location loc, BoundListTy lbs, BoundListTy ubs, ArrayRef< int64_t > steps, function_ref< void(OpBuilder &, Location, ValueRange)> bodyBuilderFn, LoopCreatorTy &&loopCreatorFn)
Builds an affine loop nest, using "loopCreatorFn" to create individual loop operations.
static LogicalResult foldLoopBounds(AffineForOp forOp)
Fold the constant bounds of a loop.
return success()
static LogicalResult replaceAffineMinBoundingBoxExpression(AffineMinOp minOp, AffineExpr dimOrSym, AffineMap *map, ValueRange dims, ValueRange syms)
Assuming dimOrSym is a quantity in the apply op map map and defined by minOp = affine_min(x_1,...
static void addAlignmentAttr(OpBuilder &builder, OperationState &result, StringAttr attrName, llvm::MaybeAlign alignment)
Adds the optional alignment attribute to result, if one is given.
static SmallVector< OpFoldResult > AffineForEmptyLoopFolder(AffineForOp forOp)
Fold the empty loop.
static LogicalResult verifyDimAndSymbolIdentifiers(OpTy &op, Operation::operand_range operands, unsigned numDims)
Utility function to verify that a set of operands are valid dimension and symbol identifiers.
static OpFoldResult makeComposedFoldedMinMax(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
static bool isDimOpValidSymbol(ShapedDimOpInterface dimOp, Region *region)
Returns true if the result of the dim op is a valid symbol for region.
static bool isQTimesDPlusR(AffineExpr e, ArrayRef< Value > operands, int64_t &div, AffineExpr &quotientTimesDiv, AffineExpr &rem)
Check if expression e is of the form d*e_1 + e_2 where 0 <= e_2 < d.
static LogicalResult replaceDimOrSym(AffineMap *map, unsigned dimOrSymbolPosition, SmallVectorImpl< Value > &dims, SmallVectorImpl< Value > &syms, bool replaceAffineMin)
Replace all occurrences of AffineExpr at position pos in map by the defining AffineApplyOp expression...
static std::optional< int64_t > getLowerBound(Value iv)
Gets the constant lower bound on an iv.
static std::optional< uint64_t > getTrivialConstantTripCount(AffineForOp forOp)
Returns constant trip count in trivial cases.
static LogicalResult verifyAffineMinMaxOp(T op)
static void printBound(AffineMapAttr boundMap, Operation::operand_range boundOperands, const char *prefix, OpAsmPrinter &p)
static void shortenAddChainsContainingAll(AffineExpr e, const llvm::SmallDenseSet< AffineExpr, 4 > &exprsToRemove, AffineExpr newVal, DenseMap< AffineExpr, AffineExpr > &replacementsMap)
Recursively traverse e.
static void composeMultiResultAffineMap(AffineMap &map, SmallVectorImpl< Value > &operands, bool composeAffineMin=false)
Composes the given affine map with the given list of operands, pulling in the maps from any affine....
static LogicalResult canonicalizeMapExprAndTermOrder(AffineMap &map)
Canonicalize the result expression order of an affine map and return success if the order changed.
static Value getZero(OpBuilder &b, Location loc, Type elementType)
Get zero value for an element type.
static Value getMemRef(Operation *memOp)
Returns the memref being read/written by a memref/affine load/store op.
Definition Utils.cpp:247
lhs
static bool isLegalToInline(InlinerInterface &interface, Region *src, Region *insertRegion, bool shouldCloneInlinedRegion, IRMapping &valueMapping)
Utility to check that all of the operations within 'src' can be inlined.
static int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getI64ArrayAttr(paddingDimensions)
b getContext())
auto load
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static Operation::operand_range getLowerBoundOperands(AffineForOp forOp)
Definition SCFToGPU.cpp:75
static Operation::operand_range getUpperBoundOperands(AffineForOp forOp)
Definition SCFToGPU.cpp:80
static VectorType getVectorType(Type scalarTy, const VectorizationStrategy *strategy)
Returns the vector type resulting from applying the provided vectorization strategy on the scalar typ...
#define div(a, b)
#define rem(a, b)
RetTy walkPostOrder(AffineExpr expr)
Base type for affine expression.
Definition AffineExpr.h:68
AffineExpr shiftDims(unsigned numDims, unsigned shift, unsigned offset=0) const
Replace dims[offset ... numDims) by dims[offset + shift ... shift + numDims).
AffineExpr shiftSymbols(unsigned numSymbols, unsigned shift, unsigned offset=0) const
Replace symbols[offset ... numSymbols) by symbols[offset + shift ... shift + numSymbols).
AffineExpr floorDiv(uint64_t v) const
AffineExprKind getKind() const
Return the classification for this type.
int64_t getLargestKnownDivisor() const
Returns the greatest known integral divisor of this affine expression.
MLIRContext * getContext() const
AffineExpr replace(AffineExpr expr, AffineExpr replacement) const
Sparse replace method.
AffineExpr ceilDiv(uint64_t v) 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.
MLIRContext * getContext() const
bool isFunctionOfDim(unsigned position) const
Return true if any affine expression involves AffineDimExpr position.
Definition AffineMap.h:221
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
AffineMap shiftDims(unsigned shift, unsigned offset=0) const
Replace dims[offset ... numDims) by dims[offset + shift ... shift + numDims).
Definition AffineMap.h:267
unsigned getNumSymbols() const
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
bool isFunctionOfSymbol(unsigned position) const
Return true if any affine expression involves AffineSymbolExpr position.
Definition AffineMap.h:228
unsigned getNumResults() const
static SmallVector< AffineMap, 4 > inferFromExprList(ArrayRef< ArrayRef< AffineExpr > > exprsList, MLIRContext *context)
Returns a vector of AffineMaps; each with as many results as exprs.size(), as many dims as the larges...
AffineMap replaceDimsAndSymbols(ArrayRef< AffineExpr > dimReplacements, ArrayRef< AffineExpr > symReplacements, unsigned numResultDims, unsigned numResultSyms) const
This method substitutes any uses of dimensions and symbols (e.g.
unsigned getNumInputs() const
AffineMap shiftSymbols(unsigned shift, unsigned offset=0) const
Replace symbols[offset ... numSymbols) by symbols[offset + shift ... shift + numSymbols).
Definition AffineMap.h:280
AffineExpr getResult(unsigned idx) const
AffineMap replace(AffineExpr expr, AffineExpr replacement, unsigned numResultDims, unsigned numResultSyms) const
Sparse replace method.
static AffineMap getConstantMap(int64_t val, MLIRContext *context)
Returns a single constant result affine map.
AffineMap getSubMap(ArrayRef< unsigned > resultPos) const
Returns the map consisting of the resultPos subset.
LogicalResult constantFold(ArrayRef< Attribute > operandConstants, SmallVectorImpl< Attribute > &results, bool *hasPoison=nullptr) const
Folds the results of the application of an affine map on the provided operands to a constant if possi...
@ Paren
Parens surrounding zero or more operands.
@ OptionalSquare
Square brackets supporting zero or more ops, or nothing.
virtual ParseResult parseColonTypeList(SmallVectorImpl< Type > &result)=0
Parse a colon followed by a type list, which must have at least one type.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
ParseResult addTypeToList(Type type, SmallVectorImpl< Type > &result)
Add the specified type to the end of the specified type list and return success.
virtual ParseResult parseOptionalRParen()=0
Parse a ) token if present.
virtual ParseResult parseLess()=0
Parse a '<' token.
virtual ParseResult parseEqual()=0
Parse a = token.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseGreater()=0
Parse a '>' token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseOptionalArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional arrow followed by a type list.
virtual ParseResult parseArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an arrow followed by a type list.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
void printOptionalArrowTypeList(TypeRange &&types)
Print an optional arrow followed by a type list.
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class represents an argument of a Block.
Definition Value.h:306
Block represents an ordered list of Operations.
Definition Block.h:33
Operation & front()
Definition Block.h:177
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
Definition Block.cpp:158
BlockArgListType getArguments()
Definition Block.h:111
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
Definition Builders.cpp:171
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
AffineMap getDimIdentityMap()
Definition Builders.cpp:392
AffineMap getMultiDimIdentityMap(unsigned rank)
Definition Builders.cpp:396
AffineExpr getAffineSymbolExpr(unsigned position)
Definition Builders.cpp:377
AffineExpr getAffineConstantExpr(int64_t constant)
Definition Builders.cpp:381
DenseIntElementsAttr getI32TensorAttr(ArrayRef< int32_t > values)
Tensor-typed DenseIntElementsAttr getters.
Definition Builders.cpp:187
IntegerAttr getI64IntegerAttr(int64_t value)
Definition Builders.cpp:120
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
NoneType getNoneType()
Definition Builders.cpp:96
BoolAttr getBoolAttr(bool value)
Definition Builders.cpp:108
AffineMap getEmptyAffineMap()
Returns a zero result affine map with no dimensions or symbols: () -> ().
Definition Builders.cpp:385
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
AffineMap getConstantAffineMap(int64_t val)
Returns a single constant result affine map with 0 dimensions and 0 symbols.
Definition Builders.cpp:387
AffineMap getSymbolIdentityMap()
Definition Builders.cpp:405
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
MLIRContext * getContext() const
Definition Builders.h:56
ArrayAttr getI64ArrayAttr(ArrayRef< int64_t > values)
Definition Builders.cpp:290
IndexType getIndexType()
Definition Builders.cpp:59
An attribute that represents a reference to a dense integer vector or tensor object.
This is a utility class for mapping one set of IR entities to another.
Definition IRMapping.h:26
auto lookup(T from) const
Lookup a mapped value within the map.
Definition IRMapping.h:72
An integer set representing a conjunction of one or more affine equalities and inequalities.
Definition IntegerSet.h:44
unsigned getNumDims() const
static IntegerSet get(unsigned dimCount, unsigned symbolCount, ArrayRef< AffineExpr > constraints, ArrayRef< bool > eqFlags)
MLIRContext * getContext() const
unsigned getNumInputs() const
ArrayRef< AffineExpr > getConstraints() const
ArrayRef< bool > getEqFlags() const
Returns the equality bits, which specify whether each of the constraints is an equality or inequality...
unsigned getNumSymbols() const
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class provides a mutable adaptor for a range of operands.
Definition ValueRange.h:119
void erase(unsigned subStart, unsigned subLen=1)
Erase the operands within the given sub-range.
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region &region, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult parseArgument(Argument &result, bool allowType=false, bool allowAttrs=false)=0
Parse a single argument with the following syntax:
ParseResult parseTrailingOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None)
Parse zero or more trailing SSA comma-separated trailing operand references with a specified surround...
virtual ParseResult parseArgumentList(SmallVectorImpl< Argument > &result, Delimiter delimiter=Delimiter::None, bool allowType=false, bool allowAttrs=false)=0
Parse zero or more arguments with a specified surrounding delimiter.
virtual ParseResult parseAffineMapOfSSAIds(SmallVectorImpl< UnresolvedOperand > &operands, Attribute &map, StringRef attrName, NamedAttrList &attrs, Delimiter delimiter=Delimiter::Square)=0
Parses an affine map attribute where dims and symbols are SSA operands.
ParseResult parseAssignmentList(SmallVectorImpl< Argument > &lhs, SmallVectorImpl< UnresolvedOperand > &rhs)
Parse a list of assignments of the form (x1 = y1, x2 = y2, ...)
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseAffineExprOfSSAIds(SmallVectorImpl< UnresolvedOperand > &dimOperands, SmallVectorImpl< UnresolvedOperand > &symbOperands, AffineExpr &expr)=0
Parses an affine expression where dims and symbols are SSA operands.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printAffineExprOfSSAIds(AffineExpr expr, ValueRange dimOperands, ValueRange symOperands)=0
Prints an affine expression of SSA ids with SSA id names used instead of dims and symbols.
virtual void printAffineMapOfSSAIds(AffineMapAttr mapAttr, ValueRange operands)=0
Prints an affine map of SSA ids, where SSA id names are used in place of dims/symbols.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
virtual void printRegionArgument(BlockArgument arg, ArrayRef< NamedAttribute > argAttrs={}, bool omitType=false)=0
Print a block argument in the usual format of: ssaName : type {attr1=42} loc("here") where location p...
virtual void printOperand(Value value)=0
Print implementations for various things an operation contains.
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
Definition Builders.cpp:439
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
This class represents a single result from folding an operation.
A trait of region holding operations that defines a new scope for polyhedral optimization purposes.
This class provides the API for ops that are known to be isolated from above.
This class implements the operand iterators for the Operation class.
Definition ValueRange.h:44
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:774
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
OperandRange operand_range
Definition Operation.h:396
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
Region * getParentRegion()
Returns the region to which the instruction belongs.
Definition Operation.h:247
operand_range::iterator operand_iterator
Definition Operation.h:397
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This class represents a point being branched from in the methods of the RegionBranchOpInterface.
bool isParent() const
Returns true if branching from the parent op.
RegionBranchTerminatorOpInterface getTerminatorPredecessorOrNull() const
Returns the terminator if branching from a region.
This class represents a successor of a region.
Region * getSuccessor() const
Return the given region successor.
bool isOperation() const
Return true if the successor is an operation.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
Block & front()
Definition Region.h:65
bool empty()
Definition Region.h:60
Operation * getParentOp()
Return the parent operation this region is attached to.
Definition Region.h:213
bool hasOneBlock()
Return true if this region has exactly one block.
Definition Region.h:68
RewritePatternSet & insert(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
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 eraseBlock(Block *block)
This method erases all operations in a block.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void finalizeOpModification(Operation *op)
This method is used to signal the end of an in-place modification of the given operation.
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
virtual void replaceUsesWithIf(Value from, Value to, function_ref< bool(OpOperand &)> functor, bool *allUsesReplaced=nullptr)
Find uses of from and replace them with to if the functor returns true.
virtual void inlineBlockBefore(Block *source, Block *dest, Block::iterator before, ValueRange argValues={})
Inline the operations of block 'source' into block 'dest' before the given position.
void mergeBlocks(Block *source, Block *dest, ValueRange argValues={})
Inline the operations of block 'source' into the end of block 'dest'.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
virtual void startOpModification(Operation *op)
This method is used to notify the rewriter that an in-place operation modification is about to happen...
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 represents a specific instance of an effect.
std::vector< SmallVector< int64_t, 8 > > operandExprStack
static Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isIndex() const
Definition Types.cpp:56
A variable that can be added to the constraint set as a "column".
static bool compare(const Variable &lhs, ComparisonOperator cmp, const Variable &rhs)
Return "true" if "lhs cmp rhs" was proven to hold.
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getType() const
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
Region * getParentRegion()
Return the Region in which this Value is defined.
Definition Value.cpp:39
AffineBound represents a lower or upper bound in the for operation.
Definition AffineOps.h:218
An AffineValueMap is an affine map plus its ML value operands and results for analysis purposes.
LogicalResult canonicalize()
Attempts to canonicalize the map and operands.
ArrayRef< Value > getOperands() const
AffineExpr getResult(unsigned i)
void reset(AffineMap map, ValueRange operands, ValueRange results={})
static void difference(const AffineValueMap &a, const AffineValueMap &b, AffineValueMap *res)
Return the value map that is the difference of value maps 'a' and 'b', represented as an affine map a...
Operation * getOwner() const
Return the owner of this operand.
Definition UseDefLists.h:38
constexpr auto RecursivelySpeculatable
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto NotSpeculatable
void buildAffineLoopNest(OpBuilder &builder, Location loc, ArrayRef< int64_t > lbs, ArrayRef< int64_t > ubs, ArrayRef< int64_t > steps, function_ref< void(OpBuilder &, Location, ValueRange)> bodyBuilderFn=nullptr)
Builds a perfect nest of affine.for loops, i.e., each loop except the innermost one contains only ano...
AffineApplyOp makeComposedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Returns a composed AffineApplyOp by composing map and operands with other AffineApplyOps supplying th...
void extractForInductionVars(ArrayRef< AffineForOp > forInsts, SmallVectorImpl< Value > *ivs)
Extracts the induction variables from a list of AffineForOps and places them in the output argument i...
bool isValidDim(Value value)
Returns true if the given Value can be used as a dimension id in the region of the closest surroundin...
bool isAffineInductionVar(Value val)
Returns true if the provided value is the induction variable of an AffineForOp or AffineParallelOp.
SmallVector< OpFoldResult > makeComposedFoldedMultiResultAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Variant of makeComposedFoldedAffineApply suitable for multi-result maps.
OpFoldResult computeProduct(Location loc, OpBuilder &builder, ArrayRef< OpFoldResult > terms)
Return the product of terms, creating an affine.apply if any of them are non-constant values.
AffineForOp getForInductionVarOwner(Value val)
Returns the loop parent of an induction variable.
void canonicalizeMapAndOperands(AffineMap *map, SmallVectorImpl< Value > *operands)
Modifies both map and operands in-place so as to:
OpFoldResult makeComposedFoldedAffineMax(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
Constructs an AffineMinOp that computes a maximum across the results of applying map to operands,...
bool isAffineForInductionVar(Value val)
Returns true if the provided value is the induction variable of an AffineForOp.
OpFoldResult makeComposedFoldedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Constructs an AffineApplyOp that applies map to operands after composing the map with the maps of any...
OpFoldResult makeComposedFoldedAffineMin(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
Constructs an AffineMinOp that computes a minimum across the results of applying map to operands,...
bool isTopLevelValue(Value value)
A utility function to check if a value is defined at the top level of an op with trait AffineScope or...
Region * getAffineAnalysisScope(Operation *op)
Returns the closest region enclosing op that is held by a non-affine operation; nullptr if there is n...
void fullyComposeAffineMapAndOperands(AffineMap *map, SmallVectorImpl< Value > *operands, bool composeAffineMin=false)
Given an affine map map and its input operands, this method composes into map, maps of AffineApplyOps...
void canonicalizeSetAndOperands(IntegerSet *set, SmallVectorImpl< Value > *operands)
Canonicalizes an integer set the same way canonicalizeMapAndOperands does for affine maps.
void extractInductionVars(ArrayRef< Operation * > affineOps, SmallVectorImpl< Value > &ivs)
Extracts the induction variables from a list of either AffineForOp or AffineParallelOp and places the...
bool isValidSymbol(Value value)
Returns true if the given value can be used as a symbol in the region of the closest surrounding op t...
AffineParallelOp getAffineParallelInductionVarOwner(Value val)
Returns true if the provided value is among the induction variables of an AffineParallelOp.
Region * getAffineScope(Operation *op)
Returns the closest region enclosing op that is held by an operation with trait AffineScope; nullptr ...
ParseResult parseDimAndSymbolList(OpAsmParser &parser, SmallVectorImpl< Value > &operands, unsigned &numDims)
Parses dimension and symbol list.
bool isAffineParallelInductionVar(Value val)
Returns true if val is the induction variable of an AffineParallelOp.
AffineMinOp makeComposedAffineMin(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands)
Returns an AffineMinOp obtained by composing map and operands with AffineApplyOps supplying those ope...
LogicalResult foldMemRefCast(Operation *op, Value inner=nullptr)
This is a common utility used for patterns of the form "someop(memref.cast) -> someop".
Definition MemRefOps.cpp:47
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:717
MemRefType getMemRefType(T &&t)
Convenience method to abbreviate casting getType().
Include the generated interface declarations.
AffineMap simplifyAffineMap(AffineMap map)
Simplifies an affine map by simplifying its underlying AffineExpr results.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
OpFoldResult getAsIndexOpFoldResult(MLIRContext *ctx, int64_t val)
Convert int64_t to integer attributes of index type and return them as OpFoldResult.
AffineMap removeDuplicateExprs(AffineMap map)
Returns a map with the same dimension and symbol count as map, but whose results are the unique affin...
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
std::function< SmallVector< Value >( OpBuilder &b, Location loc, ArrayRef< BlockArgument > newBbArgs)> NewYieldValuesFn
A function that returns the additional yielded values during replaceWithAdditionalYields.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307
std::optional< int64_t > getBoundForAffineExpr(AffineExpr expr, unsigned numDims, unsigned numSymbols, ArrayRef< std::optional< int64_t > > constLowerBounds, ArrayRef< std::optional< int64_t > > constUpperBounds, bool isUpper)
Get a lower or upper (depending on isUpper) bound for expr while using the constant lower and upper b...
SmallVector< int64_t > delinearize(int64_t linearIndex, ArrayRef< int64_t > strides)
Given the strides together with a linear index in the dimension space, return the vector-space offset...
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
bool isPure(Operation *op)
Returns true if the given operation is pure, i.e., is speculatable that does not touch memory.
AffineExprKind
Definition AffineExpr.h:40
@ CeilDiv
RHS of ceildiv is always a constant or a symbolic expression.
Definition AffineExpr.h:50
@ Mod
RHS of mod is always a constant or a symbolic expression with a positive value.
Definition AffineExpr.h:46
@ DimId
Dimensional identifier.
Definition AffineExpr.h:59
@ FloorDiv
RHS of floordiv is always a constant or a symbolic expression.
Definition AffineExpr.h:48
@ SymbolId
Symbolic identifier.
Definition AffineExpr.h:61
AffineExpr getAffineBinaryOpExpr(AffineExprKind kind, AffineExpr lhs, AffineExpr rhs)
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
Definition Matchers.h:442
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
AffineMap foldAttributesIntoMap(Builder &b, AffineMap map, ArrayRef< OpFoldResult > operands, SmallVector< Value > &remainingValues)
Fold all attributes among the given operands into the affine map.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
AffineExpr getAffineSymbolExpr(unsigned position, MLIRContext *context)
Canonicalize the affine map result expression order of an affine min/max operation.
LogicalResult matchAndRewrite(T affineOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(T affineOp, PatternRewriter &rewriter) const override
Remove duplicated expressions in affine min/max ops.
LogicalResult matchAndRewrite(T affineOp, PatternRewriter &rewriter) const override
Merge an affine min/max op to its consumers if its consumer is also an affine min/max op.
LogicalResult matchAndRewrite(T affineOp, PatternRewriter &rewriter) const override
This is the representation of an operand reference.
This class represents a listener that may be used to hook into various actions within an OpBuilder.
Definition Builders.h:288
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.