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 
9 #include "mlir/Dialect/Affine/IR/AffineOps.h"
10 #include "mlir/Dialect/Affine/IR/AffineValueMap.h"
11 #include "mlir/Dialect/MemRef/IR/MemRef.h"
12 #include "mlir/Dialect/Tensor/IR/Tensor.h"
13 #include "mlir/IR/AffineExprVisitor.h"
14 #include "mlir/IR/BlockAndValueMapping.h"
15 #include "mlir/IR/IntegerSet.h"
16 #include "mlir/IR/Matchers.h"
17 #include "mlir/IR/PatternMatch.h"
18 #include "mlir/Transforms/InliningUtils.h"
19 #include "llvm/ADT/SmallBitVector.h"
20 #include "llvm/ADT/TypeSwitch.h"
21 #include "llvm/Support/Debug.h"
22 
23 using namespace mlir;
24 
25 #define DEBUG_TYPE "affine-analysis"
26 
27 #include "mlir/Dialect/Affine/IR/AffineOpsDialect.cpp.inc"
28 
29 /// A utility function to check if a value is defined at the top level of
30 /// `region` or is an argument of `region`. A value of index type defined at the
31 /// top level of a `AffineScope` region is always a valid symbol for all
32 /// uses in that region.
33 bool mlir::isTopLevelValue(Value value, Region *region) {
34   if (auto arg = value.dyn_cast<BlockArgument>())
35     return arg.getParentRegion() == region;
36   return value.getDefiningOp()->getParentRegion() == region;
37 }
38 
39 /// Checks if `value` known to be a legal affine dimension or symbol in `src`
40 /// region remains legal if the operation that uses it is inlined into `dest`
41 /// with the given value mapping. `legalityCheck` is either `isValidDim` or
42 /// `isValidSymbol`, depending on the value being required to remain a valid
43 /// dimension or symbol.
44 static bool
45 remainsLegalAfterInline(Value value, Region *src, Region *dest,
46                         const BlockAndValueMapping &mapping,
47                         function_ref<bool(Value, Region *)> legalityCheck) {
48   // If the value is a valid dimension for any other reason than being
49   // a top-level value, it will remain valid: constants get inlined
50   // with the function, transitive affine applies also get inlined and
51   // will be checked themselves, etc.
52   if (!isTopLevelValue(value, src))
53     return true;
54 
55   // If it's a top-level value because it's a block operand, i.e. a
56   // function argument, check whether the value replacing it after
57   // inlining is a valid dimension in the new region.
58   if (value.isa<BlockArgument>())
59     return legalityCheck(mapping.lookup(value), dest);
60 
61   // If it's a top-level value because it's defined in the region,
62   // it can only be inlined if the defining op is a constant or a
63   // `dim`, which can appear anywhere and be valid, since the defining
64   // op won't be top-level anymore after inlining.
65   Attribute operandCst;
66   return matchPattern(value.getDefiningOp(), m_Constant(&operandCst)) ||
67          value.getDefiningOp<memref::DimOp>() ||
68          value.getDefiningOp<tensor::DimOp>();
69 }
70 
71 /// Checks if all values known to be legal affine dimensions or symbols in `src`
72 /// remain so if their respective users are inlined into `dest`.
73 static bool
74 remainsLegalAfterInline(ValueRange values, Region *src, Region *dest,
75                         const BlockAndValueMapping &mapping,
76                         function_ref<bool(Value, Region *)> legalityCheck) {
77   return llvm::all_of(values, [&](Value v) {
78     return remainsLegalAfterInline(v, src, dest, mapping, legalityCheck);
79   });
80 }
81 
82 /// Checks if an affine read or write operation remains legal after inlining
83 /// from `src` to `dest`.
84 template <typename OpTy>
85 static bool remainsLegalAfterInline(OpTy op, Region *src, Region *dest,
86                                     const BlockAndValueMapping &mapping) {
87   static_assert(llvm::is_one_of<OpTy, AffineReadOpInterface,
88                                 AffineWriteOpInterface>::value,
89                 "only ops with affine read/write interface are supported");
90 
91   AffineMap map = op.getAffineMap();
92   ValueRange dimOperands = op.getMapOperands().take_front(map.getNumDims());
93   ValueRange symbolOperands =
94       op.getMapOperands().take_back(map.getNumSymbols());
95   if (!remainsLegalAfterInline(
96           dimOperands, src, dest, mapping,
97           static_cast<bool (*)(Value, Region *)>(isValidDim)))
98     return false;
99   if (!remainsLegalAfterInline(
100           symbolOperands, src, dest, mapping,
101           static_cast<bool (*)(Value, Region *)>(isValidSymbol)))
102     return false;
103   return true;
104 }
105 
106 /// Checks if an affine apply operation remains legal after inlining from `src`
107 /// to `dest`.
108 //  Use "unused attribute" marker to silence clang-tidy warning stemming from
109 //  the inability to see through "llvm::TypeSwitch".
110 template <>
111 bool LLVM_ATTRIBUTE_UNUSED
112 remainsLegalAfterInline(AffineApplyOp op, Region *src, Region *dest,
113                         const BlockAndValueMapping &mapping) {
114   // If it's a valid dimension, we need to check that it remains so.
115   if (isValidDim(op.getResult(), src))
116     return remainsLegalAfterInline(
117         op.getMapOperands(), src, dest, mapping,
118         static_cast<bool (*)(Value, Region *)>(isValidDim));
119 
120   // Otherwise it must be a valid symbol, check that it remains so.
121   return remainsLegalAfterInline(
122       op.getMapOperands(), src, dest, mapping,
123       static_cast<bool (*)(Value, Region *)>(isValidSymbol));
124 }
125 
126 //===----------------------------------------------------------------------===//
127 // AffineDialect Interfaces
128 //===----------------------------------------------------------------------===//
129 
130 namespace {
131 /// This class defines the interface for handling inlining with affine
132 /// operations.
133 struct AffineInlinerInterface : public DialectInlinerInterface {
134   using DialectInlinerInterface::DialectInlinerInterface;
135 
136   //===--------------------------------------------------------------------===//
137   // Analysis Hooks
138   //===--------------------------------------------------------------------===//
139 
140   /// Returns true if the given region 'src' can be inlined into the region
141   /// 'dest' that is attached to an operation registered to the current dialect.
142   /// 'wouldBeCloned' is set if the region is cloned into its new location
143   /// rather than moved, indicating there may be other users.
144   bool isLegalToInline(Region *dest, Region *src, bool wouldBeCloned,
145                        BlockAndValueMapping &valueMapping) const final {
146     // We can inline into affine loops and conditionals if this doesn't break
147     // affine value categorization rules.
148     Operation *destOp = dest->getParentOp();
149     if (!isa<AffineParallelOp, AffineForOp, AffineIfOp>(destOp))
150       return false;
151 
152     // Multi-block regions cannot be inlined into affine constructs, all of
153     // which require single-block regions.
154     if (!llvm::hasSingleElement(*src))
155       return false;
156 
157     // Side-effecting operations that the affine dialect cannot understand
158     // should not be inlined.
159     Block &srcBlock = src->front();
160     for (Operation &op : srcBlock) {
161       // Ops with no side effects are fine,
162       if (auto iface = dyn_cast<MemoryEffectOpInterface>(op)) {
163         if (iface.hasNoEffect())
164           continue;
165       }
166 
167       // Assuming the inlined region is valid, we only need to check if the
168       // inlining would change it.
169       bool remainsValid =
170           llvm::TypeSwitch<Operation *, bool>(&op)
171               .Case<AffineApplyOp, AffineReadOpInterface,
172                     AffineWriteOpInterface>([&](auto op) {
173                 return remainsLegalAfterInline(op, src, dest, valueMapping);
174               })
175               .Default([](Operation *) {
176                 // Conservatively disallow inlining ops we cannot reason about.
177                 return false;
178               });
179 
180       if (!remainsValid)
181         return false;
182     }
183 
184     return true;
185   }
186 
187   /// Returns true if the given operation 'op', that is registered to this
188   /// dialect, can be inlined into the given region, false otherwise.
189   bool isLegalToInline(Operation *op, Region *region, bool wouldBeCloned,
190                        BlockAndValueMapping &valueMapping) const final {
191     // Always allow inlining affine operations into a region that is marked as
192     // affine scope, or into affine loops and conditionals. There are some edge
193     // cases when inlining *into* affine structures, but that is handled in the
194     // other 'isLegalToInline' hook above.
195     Operation *parentOp = region->getParentOp();
196     return parentOp->hasTrait<OpTrait::AffineScope>() ||
197            isa<AffineForOp, AffineParallelOp, AffineIfOp>(parentOp);
198   }
199 
200   /// Affine regions should be analyzed recursively.
201   bool shouldAnalyzeRecursively(Operation *op) const final { return true; }
202 };
203 } // namespace
204 
205 //===----------------------------------------------------------------------===//
206 // AffineDialect
207 //===----------------------------------------------------------------------===//
208 
209 void AffineDialect::initialize() {
210   addOperations<AffineDmaStartOp, AffineDmaWaitOp,
211 #define GET_OP_LIST
212 #include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc"
213                 >();
214   addInterfaces<AffineInlinerInterface>();
215 }
216 
217 /// Materialize a single constant operation from a given attribute value with
218 /// the desired resultant type.
219 Operation *AffineDialect::materializeConstant(OpBuilder &builder,
220                                               Attribute value, Type type,
221                                               Location loc) {
222   return builder.create<arith::ConstantOp>(loc, type, value);
223 }
224 
225 /// A utility function to check if a value is defined at the top level of an
226 /// op with trait `AffineScope`. If the value is defined in an unlinked region,
227 /// conservatively assume it is not top-level. A value of index type defined at
228 /// the top level is always a valid symbol.
229 bool mlir::isTopLevelValue(Value value) {
230   if (auto arg = value.dyn_cast<BlockArgument>()) {
231     // The block owning the argument may be unlinked, e.g. when the surrounding
232     // region has not yet been attached to an Op, at which point the parent Op
233     // is null.
234     Operation *parentOp = arg.getOwner()->getParentOp();
235     return parentOp && parentOp->hasTrait<OpTrait::AffineScope>();
236   }
237   // The defining Op may live in an unlinked block so its parent Op may be null.
238   Operation *parentOp = value.getDefiningOp()->getParentOp();
239   return parentOp && parentOp->hasTrait<OpTrait::AffineScope>();
240 }
241 
242 /// Returns the closest region enclosing `op` that is held by an operation with
243 /// trait `AffineScope`; `nullptr` if there is no such region.
244 Region *mlir::getAffineScope(Operation *op) {
245   auto *curOp = op;
246   while (auto *parentOp = curOp->getParentOp()) {
247     if (parentOp->hasTrait<OpTrait::AffineScope>())
248       return curOp->getParentRegion();
249     curOp = parentOp;
250   }
251   return nullptr;
252 }
253 
254 // A Value can be used as a dimension id iff it meets one of the following
255 // conditions:
256 // *) It is valid as a symbol.
257 // *) It is an induction variable.
258 // *) It is the result of affine apply operation with dimension id arguments.
259 bool mlir::isValidDim(Value value) {
260   // The value must be an index type.
261   if (!value.getType().isIndex())
262     return false;
263 
264   if (auto *defOp = value.getDefiningOp())
265     return isValidDim(value, getAffineScope(defOp));
266 
267   // This value has to be a block argument for an op that has the
268   // `AffineScope` trait or for an affine.for or affine.parallel.
269   auto *parentOp = value.cast<BlockArgument>().getOwner()->getParentOp();
270   return parentOp && (parentOp->hasTrait<OpTrait::AffineScope>() ||
271                       isa<AffineForOp, AffineParallelOp>(parentOp));
272 }
273 
274 // Value can be used as a dimension id iff it meets one of the following
275 // conditions:
276 // *) It is valid as a symbol.
277 // *) It is an induction variable.
278 // *) It is the result of an affine apply operation with dimension id operands.
279 bool mlir::isValidDim(Value value, Region *region) {
280   // The value must be an index type.
281   if (!value.getType().isIndex())
282     return false;
283 
284   // All valid symbols are okay.
285   if (isValidSymbol(value, region))
286     return true;
287 
288   auto *op = value.getDefiningOp();
289   if (!op) {
290     // This value has to be a block argument for an affine.for or an
291     // affine.parallel.
292     auto *parentOp = value.cast<BlockArgument>().getOwner()->getParentOp();
293     return isa<AffineForOp, AffineParallelOp>(parentOp);
294   }
295 
296   // Affine apply operation is ok if all of its operands are ok.
297   if (auto applyOp = dyn_cast<AffineApplyOp>(op))
298     return applyOp.isValidDim(region);
299   // The dim op is okay if its operand memref/tensor is defined at the top
300   // level.
301   if (auto dimOp = dyn_cast<memref::DimOp>(op))
302     return isTopLevelValue(dimOp.source());
303   if (auto dimOp = dyn_cast<tensor::DimOp>(op))
304     return isTopLevelValue(dimOp.getSource());
305   return false;
306 }
307 
308 /// Returns true if the 'index' dimension of the `memref` defined by
309 /// `memrefDefOp` is a statically  shaped one or defined using a valid symbol
310 /// for `region`.
311 template <typename AnyMemRefDefOp>
312 static bool isMemRefSizeValidSymbol(AnyMemRefDefOp memrefDefOp, unsigned index,
313                                     Region *region) {
314   auto memRefType = memrefDefOp.getType();
315   // Statically shaped.
316   if (!memRefType.isDynamicDim(index))
317     return true;
318   // Get the position of the dimension among dynamic dimensions;
319   unsigned dynamicDimPos = memRefType.getDynamicDimIndex(index);
320   return isValidSymbol(*(memrefDefOp.getDynamicSizes().begin() + dynamicDimPos),
321                        region);
322 }
323 
324 /// Returns true if the result of the dim op is a valid symbol for `region`.
325 template <typename OpTy>
326 static bool isDimOpValidSymbol(OpTy dimOp, Region *region) {
327   // The dim op is okay if its source is defined at the top level.
328   if (isTopLevelValue(dimOp.getSource()))
329     return true;
330 
331   // Conservatively handle remaining BlockArguments as non-valid symbols.
332   // E.g. scf.for iterArgs.
333   if (dimOp.getSource().template isa<BlockArgument>())
334     return false;
335 
336   // The dim op is also okay if its operand memref is a view/subview whose
337   // corresponding size is a valid symbol.
338   Optional<int64_t> index = dimOp.getConstantIndex();
339   assert(index.hasValue() &&
340          "expect only `dim` operations with a constant index");
341   int64_t i = index.getValue();
342   return TypeSwitch<Operation *, bool>(dimOp.getSource().getDefiningOp())
343       .Case<memref::ViewOp, memref::SubViewOp, memref::AllocOp>(
344           [&](auto op) { return isMemRefSizeValidSymbol(op, i, region); })
345       .Default([](Operation *) { return false; });
346 }
347 
348 // A value can be used as a symbol (at all its use sites) iff it meets one of
349 // the following conditions:
350 // *) It is a constant.
351 // *) Its defining op or block arg appearance is immediately enclosed by an op
352 //    with `AffineScope` trait.
353 // *) It is the result of an affine.apply operation with symbol operands.
354 // *) It is a result of the dim op on a memref whose corresponding size is a
355 //    valid symbol.
356 bool mlir::isValidSymbol(Value value) {
357   if (!value)
358     return false;
359 
360   // The value must be an index type.
361   if (!value.getType().isIndex())
362     return false;
363 
364   // Check that the value is a top level value.
365   if (isTopLevelValue(value))
366     return true;
367 
368   if (auto *defOp = value.getDefiningOp())
369     return isValidSymbol(value, getAffineScope(defOp));
370 
371   return false;
372 }
373 
374 /// A value can be used as a symbol for `region` iff it meets one of the
375 /// following conditions:
376 /// *) It is a constant.
377 /// *) It is the result of an affine apply operation with symbol arguments.
378 /// *) It is a result of the dim op on a memref whose corresponding size is
379 ///    a valid symbol.
380 /// *) It is defined at the top level of 'region' or is its argument.
381 /// *) It dominates `region`'s parent op.
382 /// If `region` is null, conservatively assume the symbol definition scope does
383 /// not exist and only accept the values that would be symbols regardless of
384 /// the surrounding region structure, i.e. the first three cases above.
385 bool mlir::isValidSymbol(Value value, Region *region) {
386   // The value must be an index type.
387   if (!value.getType().isIndex())
388     return false;
389 
390   // A top-level value is a valid symbol.
391   if (region && ::isTopLevelValue(value, region))
392     return true;
393 
394   auto *defOp = value.getDefiningOp();
395   if (!defOp) {
396     // A block argument that is not a top-level value is a valid symbol if it
397     // dominates region's parent op.
398     Operation *regionOp = region ? region->getParentOp() : nullptr;
399     if (regionOp && !regionOp->hasTrait<OpTrait::IsIsolatedFromAbove>())
400       if (auto *parentOpRegion = region->getParentOp()->getParentRegion())
401         return isValidSymbol(value, parentOpRegion);
402     return false;
403   }
404 
405   // Constant operation is ok.
406   Attribute operandCst;
407   if (matchPattern(defOp, m_Constant(&operandCst)))
408     return true;
409 
410   // Affine apply operation is ok if all of its operands are ok.
411   if (auto applyOp = dyn_cast<AffineApplyOp>(defOp))
412     return applyOp.isValidSymbol(region);
413 
414   // Dim op results could be valid symbols at any level.
415   if (auto dimOp = dyn_cast<memref::DimOp>(defOp))
416     return isDimOpValidSymbol(dimOp, region);
417   if (auto dimOp = dyn_cast<tensor::DimOp>(defOp))
418     return isDimOpValidSymbol(dimOp, region);
419 
420   // Check for values dominating `region`'s parent op.
421   Operation *regionOp = region ? region->getParentOp() : nullptr;
422   if (regionOp && !regionOp->hasTrait<OpTrait::IsIsolatedFromAbove>())
423     if (auto *parentRegion = region->getParentOp()->getParentRegion())
424       return isValidSymbol(value, parentRegion);
425 
426   return false;
427 }
428 
429 // Returns true if 'value' is a valid index to an affine operation (e.g.
430 // affine.load, affine.store, affine.dma_start, affine.dma_wait) where
431 // `region` provides the polyhedral symbol scope. Returns false otherwise.
432 static bool isValidAffineIndexOperand(Value value, Region *region) {
433   return isValidDim(value, region) || isValidSymbol(value, region);
434 }
435 
436 /// Prints dimension and symbol list.
437 static void printDimAndSymbolList(Operation::operand_iterator begin,
438                                   Operation::operand_iterator end,
439                                   unsigned numDims, OpAsmPrinter &printer) {
440   OperandRange operands(begin, end);
441   printer << '(' << operands.take_front(numDims) << ')';
442   if (operands.size() > numDims)
443     printer << '[' << operands.drop_front(numDims) << ']';
444 }
445 
446 /// Parses dimension and symbol list and returns true if parsing failed.
447 ParseResult mlir::parseDimAndSymbolList(OpAsmParser &parser,
448                                         SmallVectorImpl<Value> &operands,
449                                         unsigned &numDims) {
450   SmallVector<OpAsmParser::UnresolvedOperand, 8> opInfos;
451   if (parser.parseOperandList(opInfos, OpAsmParser::Delimiter::Paren))
452     return failure();
453   // Store number of dimensions for validation by caller.
454   numDims = opInfos.size();
455 
456   // Parse the optional symbol operands.
457   auto indexTy = parser.getBuilder().getIndexType();
458   return failure(parser.parseOperandList(
459                      opInfos, OpAsmParser::Delimiter::OptionalSquare) ||
460                  parser.resolveOperands(opInfos, indexTy, operands));
461 }
462 
463 /// Utility function to verify that a set of operands are valid dimension and
464 /// symbol identifiers. The operands should be laid out such that the dimension
465 /// operands are before the symbol operands. This function returns failure if
466 /// there was an invalid operand. An operation is provided to emit any necessary
467 /// errors.
468 template <typename OpTy>
469 static LogicalResult
470 verifyDimAndSymbolIdentifiers(OpTy &op, Operation::operand_range operands,
471                               unsigned numDims) {
472   unsigned opIt = 0;
473   for (auto operand : operands) {
474     if (opIt++ < numDims) {
475       if (!isValidDim(operand, getAffineScope(op)))
476         return op.emitOpError("operand cannot be used as a dimension id");
477     } else if (!isValidSymbol(operand, getAffineScope(op))) {
478       return op.emitOpError("operand cannot be used as a symbol");
479     }
480   }
481   return success();
482 }
483 
484 //===----------------------------------------------------------------------===//
485 // AffineApplyOp
486 //===----------------------------------------------------------------------===//
487 
488 AffineValueMap AffineApplyOp::getAffineValueMap() {
489   return AffineValueMap(getAffineMap(), getOperands(), getResult());
490 }
491 
492 ParseResult AffineApplyOp::parse(OpAsmParser &parser, OperationState &result) {
493   auto &builder = parser.getBuilder();
494   auto indexTy = builder.getIndexType();
495 
496   AffineMapAttr mapAttr;
497   unsigned numDims;
498   if (parser.parseAttribute(mapAttr, "map", result.attributes) ||
499       parseDimAndSymbolList(parser, result.operands, numDims) ||
500       parser.parseOptionalAttrDict(result.attributes))
501     return failure();
502   auto map = mapAttr.getValue();
503 
504   if (map.getNumDims() != numDims ||
505       numDims + map.getNumSymbols() != result.operands.size()) {
506     return parser.emitError(parser.getNameLoc(),
507                             "dimension or symbol index mismatch");
508   }
509 
510   result.types.append(map.getNumResults(), indexTy);
511   return success();
512 }
513 
514 void AffineApplyOp::print(OpAsmPrinter &p) {
515   p << " " << getMapAttr();
516   printDimAndSymbolList(operand_begin(), operand_end(),
517                         getAffineMap().getNumDims(), p);
518   p.printOptionalAttrDict((*this)->getAttrs(), /*elidedAttrs=*/{"map"});
519 }
520 
521 LogicalResult AffineApplyOp::verify() {
522   // Check input and output dimensions match.
523   AffineMap affineMap = getMap();
524 
525   // Verify that operand count matches affine map dimension and symbol count.
526   if (getNumOperands() != affineMap.getNumDims() + affineMap.getNumSymbols())
527     return emitOpError(
528         "operand count and affine map dimension and symbol count must match");
529 
530   // Verify that the map only produces one result.
531   if (affineMap.getNumResults() != 1)
532     return emitOpError("mapping must produce one value");
533 
534   return success();
535 }
536 
537 // The result of the affine apply operation can be used as a dimension id if all
538 // its operands are valid dimension ids.
539 bool AffineApplyOp::isValidDim() {
540   return llvm::all_of(getOperands(),
541                       [](Value op) { return mlir::isValidDim(op); });
542 }
543 
544 // The result of the affine apply operation can be used as a dimension id if all
545 // its operands are valid dimension ids with the parent operation of `region`
546 // defining the polyhedral scope for symbols.
547 bool AffineApplyOp::isValidDim(Region *region) {
548   return llvm::all_of(getOperands(),
549                       [&](Value op) { return ::isValidDim(op, region); });
550 }
551 
552 // The result of the affine apply operation can be used as a symbol if all its
553 // operands are symbols.
554 bool AffineApplyOp::isValidSymbol() {
555   return llvm::all_of(getOperands(),
556                       [](Value op) { return mlir::isValidSymbol(op); });
557 }
558 
559 // The result of the affine apply operation can be used as a symbol in `region`
560 // if all its operands are symbols in `region`.
561 bool AffineApplyOp::isValidSymbol(Region *region) {
562   return llvm::all_of(getOperands(), [&](Value operand) {
563     return mlir::isValidSymbol(operand, region);
564   });
565 }
566 
567 OpFoldResult AffineApplyOp::fold(ArrayRef<Attribute> operands) {
568   auto map = getAffineMap();
569 
570   // Fold dims and symbols to existing values.
571   auto expr = map.getResult(0);
572   if (auto dim = expr.dyn_cast<AffineDimExpr>())
573     return getOperand(dim.getPosition());
574   if (auto sym = expr.dyn_cast<AffineSymbolExpr>())
575     return getOperand(map.getNumDims() + sym.getPosition());
576 
577   // Otherwise, default to folding the map.
578   SmallVector<Attribute, 1> result;
579   if (failed(map.constantFold(operands, result)))
580     return {};
581   return result[0];
582 }
583 
584 /// Replace all occurrences of AffineExpr at position `pos` in `map` by the
585 /// defining AffineApplyOp expression and operands.
586 /// When `dimOrSymbolPosition < dims.size()`, AffineDimExpr@[pos] is replaced.
587 /// When `dimOrSymbolPosition >= dims.size()`,
588 /// AffineSymbolExpr@[pos - dims.size()] is replaced.
589 /// Mutate `map`,`dims` and `syms` in place as follows:
590 ///   1. `dims` and `syms` are only appended to.
591 ///   2. `map` dim and symbols are gradually shifted to higer positions.
592 ///   3. Old `dim` and `sym` entries are replaced by nullptr
593 /// This avoids the need for any bookkeeping.
594 static LogicalResult replaceDimOrSym(AffineMap *map,
595                                      unsigned dimOrSymbolPosition,
596                                      SmallVectorImpl<Value> &dims,
597                                      SmallVectorImpl<Value> &syms) {
598   bool isDimReplacement = (dimOrSymbolPosition < dims.size());
599   unsigned pos = isDimReplacement ? dimOrSymbolPosition
600                                   : dimOrSymbolPosition - dims.size();
601   Value &v = isDimReplacement ? dims[pos] : syms[pos];
602   if (!v)
603     return failure();
604 
605   auto affineApply = v.getDefiningOp<AffineApplyOp>();
606   if (!affineApply)
607     return failure();
608 
609   // At this point we will perform a replacement of `v`, set the entry in `dim`
610   // or `sym` to nullptr immediately.
611   v = nullptr;
612 
613   // Compute the map, dims and symbols coming from the AffineApplyOp.
614   AffineMap composeMap = affineApply.getAffineMap();
615   assert(composeMap.getNumResults() == 1 && "affine.apply with >1 results");
616   AffineExpr composeExpr =
617       composeMap.shiftDims(dims.size()).shiftSymbols(syms.size()).getResult(0);
618   ValueRange composeDims =
619       affineApply.getMapOperands().take_front(composeMap.getNumDims());
620   ValueRange composeSyms =
621       affineApply.getMapOperands().take_back(composeMap.getNumSymbols());
622 
623   // Append the dims and symbols where relevant and perform the replacement.
624   MLIRContext *ctx = map->getContext();
625   AffineExpr toReplace = isDimReplacement ? getAffineDimExpr(pos, ctx)
626                                           : getAffineSymbolExpr(pos, ctx);
627   dims.append(composeDims.begin(), composeDims.end());
628   syms.append(composeSyms.begin(), composeSyms.end());
629   *map = map->replace(toReplace, composeExpr, dims.size(), syms.size());
630 
631   return success();
632 }
633 
634 /// Iterate over `operands` and fold away all those produced by an AffineApplyOp
635 /// iteratively. Perform canonicalization of map and operands as well as
636 /// AffineMap simplification. `map` and `operands` are mutated in place.
637 static void composeAffineMapAndOperands(AffineMap *map,
638                                         SmallVectorImpl<Value> *operands) {
639   if (map->getNumResults() == 0) {
640     canonicalizeMapAndOperands(map, operands);
641     *map = simplifyAffineMap(*map);
642     return;
643   }
644 
645   MLIRContext *ctx = map->getContext();
646   SmallVector<Value, 4> dims(operands->begin(),
647                              operands->begin() + map->getNumDims());
648   SmallVector<Value, 4> syms(operands->begin() + map->getNumDims(),
649                              operands->end());
650 
651   // Iterate over dims and symbols coming from AffineApplyOp and replace until
652   // exhaustion. This iteratively mutates `map`, `dims` and `syms`. Both `dims`
653   // and `syms` can only increase by construction.
654   // The implementation uses a `while` loop to support the case of symbols
655   // that may be constructed from dims ;this may be overkill.
656   while (true) {
657     bool changed = false;
658     for (unsigned pos = 0; pos != dims.size() + syms.size(); ++pos)
659       if ((changed |= succeeded(replaceDimOrSym(map, pos, dims, syms))))
660         break;
661     if (!changed)
662       break;
663   }
664 
665   // Clear operands so we can fill them anew.
666   operands->clear();
667 
668   // At this point we may have introduced null operands, prune them out before
669   // canonicalizing map and operands.
670   unsigned nDims = 0, nSyms = 0;
671   SmallVector<AffineExpr, 4> dimReplacements, symReplacements;
672   dimReplacements.reserve(dims.size());
673   symReplacements.reserve(syms.size());
674   for (auto *container : {&dims, &syms}) {
675     bool isDim = (container == &dims);
676     auto &repls = isDim ? dimReplacements : symReplacements;
677     for (const auto &en : llvm::enumerate(*container)) {
678       Value v = en.value();
679       if (!v) {
680         assert(isDim ? !map->isFunctionOfDim(en.index())
681                      : !map->isFunctionOfSymbol(en.index()) &&
682                            "map is function of unexpected expr@pos");
683         repls.push_back(getAffineConstantExpr(0, ctx));
684         continue;
685       }
686       repls.push_back(isDim ? getAffineDimExpr(nDims++, ctx)
687                             : getAffineSymbolExpr(nSyms++, ctx));
688       operands->push_back(v);
689     }
690   }
691   *map = map->replaceDimsAndSymbols(dimReplacements, symReplacements, nDims,
692                                     nSyms);
693 
694   // Canonicalize and simplify before returning.
695   canonicalizeMapAndOperands(map, operands);
696   *map = simplifyAffineMap(*map);
697 }
698 
699 void mlir::fullyComposeAffineMapAndOperands(AffineMap *map,
700                                             SmallVectorImpl<Value> *operands) {
701   while (llvm::any_of(*operands, [](Value v) {
702     return isa_and_nonnull<AffineApplyOp>(v.getDefiningOp());
703   })) {
704     composeAffineMapAndOperands(map, operands);
705   }
706 }
707 
708 AffineApplyOp mlir::makeComposedAffineApply(OpBuilder &b, Location loc,
709                                             AffineMap map,
710                                             ValueRange operands) {
711   AffineMap normalizedMap = map;
712   SmallVector<Value, 8> normalizedOperands(operands.begin(), operands.end());
713   composeAffineMapAndOperands(&normalizedMap, &normalizedOperands);
714   assert(normalizedMap);
715   return b.create<AffineApplyOp>(loc, normalizedMap, normalizedOperands);
716 }
717 
718 AffineApplyOp mlir::makeComposedAffineApply(OpBuilder &b, Location loc,
719                                             AffineExpr e, ValueRange values) {
720   return makeComposedAffineApply(
721       b, loc, AffineMap::inferFromExprList(ArrayRef<AffineExpr>{e}).front(),
722       values);
723 }
724 
725 /// Fully compose map with operands and canonicalize the result.
726 /// Return the `createOrFold`'ed AffineApply op.
727 static Value createFoldedComposedAffineApply(OpBuilder &b, Location loc,
728                                              AffineMap map,
729                                              ValueRange operandsRef) {
730   SmallVector<Value, 4> operands(operandsRef.begin(), operandsRef.end());
731   fullyComposeAffineMapAndOperands(&map, &operands);
732   canonicalizeMapAndOperands(&map, &operands);
733   return b.createOrFold<AffineApplyOp>(loc, map, operands);
734 }
735 
736 SmallVector<Value, 4> mlir::applyMapToValues(OpBuilder &b, Location loc,
737                                              AffineMap map, ValueRange values) {
738   SmallVector<Value, 4> res;
739   res.reserve(map.getNumResults());
740   unsigned numDims = map.getNumDims(), numSym = map.getNumSymbols();
741   // For each `expr` in `map`, applies the `expr` to the values extracted from
742   // ranges. If the resulting application can be folded into a Value, the
743   // folding occurs eagerly.
744   for (auto expr : map.getResults()) {
745     AffineMap map = AffineMap::get(numDims, numSym, expr);
746     res.push_back(createFoldedComposedAffineApply(b, loc, map, values));
747   }
748   return res;
749 }
750 
751 SmallVector<OpFoldResult>
752 mlir::applyMapToValues(IRRewriter &b, Location loc, AffineMap map,
753                        ArrayRef<OpFoldResult> values) {
754   // Materialize constants and keep track of produced operations so we can clean
755   // them up later.
756   SmallVector<Operation *> constants;
757   SmallVector<Value> actualValues;
758   actualValues.reserve(values.size());
759   auto *dialect = b.getContext()->getLoadedDialect<AffineDialect>();
760   for (OpFoldResult ofr : values) {
761     if (auto value = ofr.dyn_cast<Value>()) {
762       actualValues.push_back(value);
763       continue;
764     }
765     constants.push_back(dialect->materializeConstant(b, ofr.get<Attribute>(),
766                                                      b.getIndexType(), loc));
767     actualValues.push_back(constants.back()->getResult(0));
768   }
769 
770   // Compose, fold and construct maps for each result independently because they
771   // may simplify more effectively.
772   SmallVector<OpFoldResult> results;
773   results.reserve(map.getNumResults());
774   bool foldedAll = true;
775   for (auto i : llvm::seq<unsigned>(0, map.getNumResults())) {
776     AffineMap submap = map.getSubMap({i});
777     SmallVector<Value> operands = actualValues;
778     fullyComposeAffineMapAndOperands(&submap, &operands);
779     canonicalizeMapAndOperands(&submap, &operands);
780 
781     // Identify the constant operands and extract their values as attributes.
782     // Note that we cannot use the original values directly because the list of
783     // operands may have changed due to canonicalization and composition.
784     SmallVector<Attribute> constantOperands;
785     constantOperands.reserve(operands.size());
786     for (Value operand : operands) {
787       IntegerAttr attr;
788       if (matchPattern(operand, m_Constant(&attr)))
789         constantOperands.push_back(attr);
790       else
791         constantOperands.push_back(nullptr);
792     }
793 
794     // Create an apply operation and immediately attempt to fold it. On sucess,
795     // delete the operation and prepare the (unmaterialized) value for being
796     // returned. On failure, return the function result.
797     // TODO: arguably, the main folder (createOrFold) API should support this
798     // use case instead of indiscriminately materializing constants.
799     auto apply = b.create<AffineApplyOp>(loc, submap, operands);
800     SmallVector<OpFoldResult, 1> foldResult;
801     if (succeeded(apply->fold(constantOperands, foldResult))) {
802       assert(foldResult.size() == 1 && "expected single-result map");
803       b.eraseOp(apply);
804       results.push_back(foldResult.front());
805     } else {
806       results.push_back(apply.getResult());
807       foldedAll = false;
808     }
809   }
810 
811   // If the entire map could be folded, remove the constants that were used in
812   // the initial ops.
813   if (foldedAll) {
814     for (Operation *constant : constants)
815       b.eraseOp(constant);
816   }
817 
818   return results;
819 }
820 
821 // A symbol may appear as a dim in affine.apply operations. This function
822 // canonicalizes dims that are valid symbols into actual symbols.
823 template <class MapOrSet>
824 static void canonicalizePromotedSymbols(MapOrSet *mapOrSet,
825                                         SmallVectorImpl<Value> *operands) {
826   if (!mapOrSet || operands->empty())
827     return;
828 
829   assert(mapOrSet->getNumInputs() == operands->size() &&
830          "map/set inputs must match number of operands");
831 
832   auto *context = mapOrSet->getContext();
833   SmallVector<Value, 8> resultOperands;
834   resultOperands.reserve(operands->size());
835   SmallVector<Value, 8> remappedSymbols;
836   remappedSymbols.reserve(operands->size());
837   unsigned nextDim = 0;
838   unsigned nextSym = 0;
839   unsigned oldNumSyms = mapOrSet->getNumSymbols();
840   SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims());
841   for (unsigned i = 0, e = mapOrSet->getNumInputs(); i != e; ++i) {
842     if (i < mapOrSet->getNumDims()) {
843       if (isValidSymbol((*operands)[i])) {
844         // This is a valid symbol that appears as a dim, canonicalize it.
845         dimRemapping[i] = getAffineSymbolExpr(oldNumSyms + nextSym++, context);
846         remappedSymbols.push_back((*operands)[i]);
847       } else {
848         dimRemapping[i] = getAffineDimExpr(nextDim++, context);
849         resultOperands.push_back((*operands)[i]);
850       }
851     } else {
852       resultOperands.push_back((*operands)[i]);
853     }
854   }
855 
856   resultOperands.append(remappedSymbols.begin(), remappedSymbols.end());
857   *operands = resultOperands;
858   *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, {}, nextDim,
859                                               oldNumSyms + nextSym);
860 
861   assert(mapOrSet->getNumInputs() == operands->size() &&
862          "map/set inputs must match number of operands");
863 }
864 
865 // Works for either an affine map or an integer set.
866 template <class MapOrSet>
867 static void canonicalizeMapOrSetAndOperands(MapOrSet *mapOrSet,
868                                             SmallVectorImpl<Value> *operands) {
869   static_assert(llvm::is_one_of<MapOrSet, AffineMap, IntegerSet>::value,
870                 "Argument must be either of AffineMap or IntegerSet type");
871 
872   if (!mapOrSet || operands->empty())
873     return;
874 
875   assert(mapOrSet->getNumInputs() == operands->size() &&
876          "map/set inputs must match number of operands");
877 
878   canonicalizePromotedSymbols<MapOrSet>(mapOrSet, operands);
879 
880   // Check to see what dims are used.
881   llvm::SmallBitVector usedDims(mapOrSet->getNumDims());
882   llvm::SmallBitVector usedSyms(mapOrSet->getNumSymbols());
883   mapOrSet->walkExprs([&](AffineExpr expr) {
884     if (auto dimExpr = expr.dyn_cast<AffineDimExpr>())
885       usedDims[dimExpr.getPosition()] = true;
886     else if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>())
887       usedSyms[symExpr.getPosition()] = true;
888   });
889 
890   auto *context = mapOrSet->getContext();
891 
892   SmallVector<Value, 8> resultOperands;
893   resultOperands.reserve(operands->size());
894 
895   llvm::SmallDenseMap<Value, AffineExpr, 8> seenDims;
896   SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims());
897   unsigned nextDim = 0;
898   for (unsigned i = 0, e = mapOrSet->getNumDims(); i != e; ++i) {
899     if (usedDims[i]) {
900       // Remap dim positions for duplicate operands.
901       auto it = seenDims.find((*operands)[i]);
902       if (it == seenDims.end()) {
903         dimRemapping[i] = getAffineDimExpr(nextDim++, context);
904         resultOperands.push_back((*operands)[i]);
905         seenDims.insert(std::make_pair((*operands)[i], dimRemapping[i]));
906       } else {
907         dimRemapping[i] = it->second;
908       }
909     }
910   }
911   llvm::SmallDenseMap<Value, AffineExpr, 8> seenSymbols;
912   SmallVector<AffineExpr, 8> symRemapping(mapOrSet->getNumSymbols());
913   unsigned nextSym = 0;
914   for (unsigned i = 0, e = mapOrSet->getNumSymbols(); i != e; ++i) {
915     if (!usedSyms[i])
916       continue;
917     // Handle constant operands (only needed for symbolic operands since
918     // constant operands in dimensional positions would have already been
919     // promoted to symbolic positions above).
920     IntegerAttr operandCst;
921     if (matchPattern((*operands)[i + mapOrSet->getNumDims()],
922                      m_Constant(&operandCst))) {
923       symRemapping[i] =
924           getAffineConstantExpr(operandCst.getValue().getSExtValue(), context);
925       continue;
926     }
927     // Remap symbol positions for duplicate operands.
928     auto it = seenSymbols.find((*operands)[i + mapOrSet->getNumDims()]);
929     if (it == seenSymbols.end()) {
930       symRemapping[i] = getAffineSymbolExpr(nextSym++, context);
931       resultOperands.push_back((*operands)[i + mapOrSet->getNumDims()]);
932       seenSymbols.insert(std::make_pair((*operands)[i + mapOrSet->getNumDims()],
933                                         symRemapping[i]));
934     } else {
935       symRemapping[i] = it->second;
936     }
937   }
938   *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, symRemapping,
939                                               nextDim, nextSym);
940   *operands = resultOperands;
941 }
942 
943 void mlir::canonicalizeMapAndOperands(AffineMap *map,
944                                       SmallVectorImpl<Value> *operands) {
945   canonicalizeMapOrSetAndOperands<AffineMap>(map, operands);
946 }
947 
948 void mlir::canonicalizeSetAndOperands(IntegerSet *set,
949                                       SmallVectorImpl<Value> *operands) {
950   canonicalizeMapOrSetAndOperands<IntegerSet>(set, operands);
951 }
952 
953 namespace {
954 /// Simplify AffineApply, AffineLoad, and AffineStore operations by composing
955 /// maps that supply results into them.
956 ///
957 template <typename AffineOpTy>
958 struct SimplifyAffineOp : public OpRewritePattern<AffineOpTy> {
959   using OpRewritePattern<AffineOpTy>::OpRewritePattern;
960 
961   /// Replace the affine op with another instance of it with the supplied
962   /// map and mapOperands.
963   void replaceAffineOp(PatternRewriter &rewriter, AffineOpTy affineOp,
964                        AffineMap map, ArrayRef<Value> mapOperands) const;
965 
966   LogicalResult matchAndRewrite(AffineOpTy affineOp,
967                                 PatternRewriter &rewriter) const override {
968     static_assert(
969         llvm::is_one_of<AffineOpTy, AffineLoadOp, AffinePrefetchOp,
970                         AffineStoreOp, AffineApplyOp, AffineMinOp, AffineMaxOp,
971                         AffineVectorStoreOp, AffineVectorLoadOp>::value,
972         "affine load/store/vectorstore/vectorload/apply/prefetch/min/max op "
973         "expected");
974     auto map = affineOp.getAffineMap();
975     AffineMap oldMap = map;
976     auto oldOperands = affineOp.getMapOperands();
977     SmallVector<Value, 8> resultOperands(oldOperands);
978     composeAffineMapAndOperands(&map, &resultOperands);
979     canonicalizeMapAndOperands(&map, &resultOperands);
980     if (map == oldMap && std::equal(oldOperands.begin(), oldOperands.end(),
981                                     resultOperands.begin()))
982       return failure();
983 
984     replaceAffineOp(rewriter, affineOp, map, resultOperands);
985     return success();
986   }
987 };
988 
989 // Specialize the template to account for the different build signatures for
990 // affine load, store, and apply ops.
991 template <>
992 void SimplifyAffineOp<AffineLoadOp>::replaceAffineOp(
993     PatternRewriter &rewriter, AffineLoadOp load, AffineMap map,
994     ArrayRef<Value> mapOperands) const {
995   rewriter.replaceOpWithNewOp<AffineLoadOp>(load, load.getMemRef(), map,
996                                             mapOperands);
997 }
998 template <>
999 void SimplifyAffineOp<AffinePrefetchOp>::replaceAffineOp(
1000     PatternRewriter &rewriter, AffinePrefetchOp prefetch, AffineMap map,
1001     ArrayRef<Value> mapOperands) const {
1002   rewriter.replaceOpWithNewOp<AffinePrefetchOp>(
1003       prefetch, prefetch.getMemref(), map, mapOperands,
1004       prefetch.getLocalityHint(), prefetch.getIsWrite(),
1005       prefetch.getIsDataCache());
1006 }
1007 template <>
1008 void SimplifyAffineOp<AffineStoreOp>::replaceAffineOp(
1009     PatternRewriter &rewriter, AffineStoreOp store, AffineMap map,
1010     ArrayRef<Value> mapOperands) const {
1011   rewriter.replaceOpWithNewOp<AffineStoreOp>(
1012       store, store.getValueToStore(), store.getMemRef(), map, mapOperands);
1013 }
1014 template <>
1015 void SimplifyAffineOp<AffineVectorLoadOp>::replaceAffineOp(
1016     PatternRewriter &rewriter, AffineVectorLoadOp vectorload, AffineMap map,
1017     ArrayRef<Value> mapOperands) const {
1018   rewriter.replaceOpWithNewOp<AffineVectorLoadOp>(
1019       vectorload, vectorload.getVectorType(), vectorload.getMemRef(), map,
1020       mapOperands);
1021 }
1022 template <>
1023 void SimplifyAffineOp<AffineVectorStoreOp>::replaceAffineOp(
1024     PatternRewriter &rewriter, AffineVectorStoreOp vectorstore, AffineMap map,
1025     ArrayRef<Value> mapOperands) const {
1026   rewriter.replaceOpWithNewOp<AffineVectorStoreOp>(
1027       vectorstore, vectorstore.getValueToStore(), vectorstore.getMemRef(), map,
1028       mapOperands);
1029 }
1030 
1031 // Generic version for ops that don't have extra operands.
1032 template <typename AffineOpTy>
1033 void SimplifyAffineOp<AffineOpTy>::replaceAffineOp(
1034     PatternRewriter &rewriter, AffineOpTy op, AffineMap map,
1035     ArrayRef<Value> mapOperands) const {
1036   rewriter.replaceOpWithNewOp<AffineOpTy>(op, map, mapOperands);
1037 }
1038 } // namespace
1039 
1040 void AffineApplyOp::getCanonicalizationPatterns(RewritePatternSet &results,
1041                                                 MLIRContext *context) {
1042   results.add<SimplifyAffineOp<AffineApplyOp>>(context);
1043 }
1044 
1045 //===----------------------------------------------------------------------===//
1046 // Common canonicalization pattern support logic
1047 //===----------------------------------------------------------------------===//
1048 
1049 /// This is a common class used for patterns of the form
1050 /// "someop(memrefcast) -> someop".  It folds the source of any memref.cast
1051 /// into the root operation directly.
1052 static LogicalResult foldMemRefCast(Operation *op, Value ignore = nullptr) {
1053   bool folded = false;
1054   for (OpOperand &operand : op->getOpOperands()) {
1055     auto cast = operand.get().getDefiningOp<memref::CastOp>();
1056     if (cast && operand.get() != ignore &&
1057         !cast.getOperand().getType().isa<UnrankedMemRefType>()) {
1058       operand.set(cast.getOperand());
1059       folded = true;
1060     }
1061   }
1062   return success(folded);
1063 }
1064 
1065 //===----------------------------------------------------------------------===//
1066 // AffineDmaStartOp
1067 //===----------------------------------------------------------------------===//
1068 
1069 // TODO: Check that map operands are loop IVs or symbols.
1070 void AffineDmaStartOp::build(OpBuilder &builder, OperationState &result,
1071                              Value srcMemRef, AffineMap srcMap,
1072                              ValueRange srcIndices, Value destMemRef,
1073                              AffineMap dstMap, ValueRange destIndices,
1074                              Value tagMemRef, AffineMap tagMap,
1075                              ValueRange tagIndices, Value numElements,
1076                              Value stride, Value elementsPerStride) {
1077   result.addOperands(srcMemRef);
1078   result.addAttribute(getSrcMapAttrStrName(), AffineMapAttr::get(srcMap));
1079   result.addOperands(srcIndices);
1080   result.addOperands(destMemRef);
1081   result.addAttribute(getDstMapAttrStrName(), AffineMapAttr::get(dstMap));
1082   result.addOperands(destIndices);
1083   result.addOperands(tagMemRef);
1084   result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap));
1085   result.addOperands(tagIndices);
1086   result.addOperands(numElements);
1087   if (stride) {
1088     result.addOperands({stride, elementsPerStride});
1089   }
1090 }
1091 
1092 void AffineDmaStartOp::print(OpAsmPrinter &p) {
1093   p << " " << getSrcMemRef() << '[';
1094   p.printAffineMapOfSSAIds(getSrcMapAttr(), getSrcIndices());
1095   p << "], " << getDstMemRef() << '[';
1096   p.printAffineMapOfSSAIds(getDstMapAttr(), getDstIndices());
1097   p << "], " << getTagMemRef() << '[';
1098   p.printAffineMapOfSSAIds(getTagMapAttr(), getTagIndices());
1099   p << "], " << getNumElements();
1100   if (isStrided()) {
1101     p << ", " << getStride();
1102     p << ", " << getNumElementsPerStride();
1103   }
1104   p << " : " << getSrcMemRefType() << ", " << getDstMemRefType() << ", "
1105     << getTagMemRefType();
1106 }
1107 
1108 // Parse AffineDmaStartOp.
1109 // Ex:
1110 //   affine.dma_start %src[%i, %j], %dst[%k, %l], %tag[%index], %size,
1111 //     %stride, %num_elt_per_stride
1112 //       : memref<3076 x f32, 0>, memref<1024 x f32, 2>, memref<1 x i32>
1113 //
1114 ParseResult AffineDmaStartOp::parse(OpAsmParser &parser,
1115                                     OperationState &result) {
1116   OpAsmParser::UnresolvedOperand srcMemRefInfo;
1117   AffineMapAttr srcMapAttr;
1118   SmallVector<OpAsmParser::UnresolvedOperand, 4> srcMapOperands;
1119   OpAsmParser::UnresolvedOperand dstMemRefInfo;
1120   AffineMapAttr dstMapAttr;
1121   SmallVector<OpAsmParser::UnresolvedOperand, 4> dstMapOperands;
1122   OpAsmParser::UnresolvedOperand tagMemRefInfo;
1123   AffineMapAttr tagMapAttr;
1124   SmallVector<OpAsmParser::UnresolvedOperand, 4> tagMapOperands;
1125   OpAsmParser::UnresolvedOperand numElementsInfo;
1126   SmallVector<OpAsmParser::UnresolvedOperand, 2> strideInfo;
1127 
1128   SmallVector<Type, 3> types;
1129   auto indexType = parser.getBuilder().getIndexType();
1130 
1131   // Parse and resolve the following list of operands:
1132   // *) dst memref followed by its affine maps operands (in square brackets).
1133   // *) src memref followed by its affine map operands (in square brackets).
1134   // *) tag memref followed by its affine map operands (in square brackets).
1135   // *) number of elements transferred by DMA operation.
1136   if (parser.parseOperand(srcMemRefInfo) ||
1137       parser.parseAffineMapOfSSAIds(srcMapOperands, srcMapAttr,
1138                                     getSrcMapAttrStrName(),
1139                                     result.attributes) ||
1140       parser.parseComma() || parser.parseOperand(dstMemRefInfo) ||
1141       parser.parseAffineMapOfSSAIds(dstMapOperands, dstMapAttr,
1142                                     getDstMapAttrStrName(),
1143                                     result.attributes) ||
1144       parser.parseComma() || parser.parseOperand(tagMemRefInfo) ||
1145       parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr,
1146                                     getTagMapAttrStrName(),
1147                                     result.attributes) ||
1148       parser.parseComma() || parser.parseOperand(numElementsInfo))
1149     return failure();
1150 
1151   // Parse optional stride and elements per stride.
1152   if (parser.parseTrailingOperandList(strideInfo))
1153     return failure();
1154 
1155   if (!strideInfo.empty() && strideInfo.size() != 2) {
1156     return parser.emitError(parser.getNameLoc(),
1157                             "expected two stride related operands");
1158   }
1159   bool isStrided = strideInfo.size() == 2;
1160 
1161   if (parser.parseColonTypeList(types))
1162     return failure();
1163 
1164   if (types.size() != 3)
1165     return parser.emitError(parser.getNameLoc(), "expected three types");
1166 
1167   if (parser.resolveOperand(srcMemRefInfo, types[0], result.operands) ||
1168       parser.resolveOperands(srcMapOperands, indexType, result.operands) ||
1169       parser.resolveOperand(dstMemRefInfo, types[1], result.operands) ||
1170       parser.resolveOperands(dstMapOperands, indexType, result.operands) ||
1171       parser.resolveOperand(tagMemRefInfo, types[2], result.operands) ||
1172       parser.resolveOperands(tagMapOperands, indexType, result.operands) ||
1173       parser.resolveOperand(numElementsInfo, indexType, result.operands))
1174     return failure();
1175 
1176   if (isStrided) {
1177     if (parser.resolveOperands(strideInfo, indexType, result.operands))
1178       return failure();
1179   }
1180 
1181   // Check that src/dst/tag operand counts match their map.numInputs.
1182   if (srcMapOperands.size() != srcMapAttr.getValue().getNumInputs() ||
1183       dstMapOperands.size() != dstMapAttr.getValue().getNumInputs() ||
1184       tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
1185     return parser.emitError(parser.getNameLoc(),
1186                             "memref operand count not equal to map.numInputs");
1187   return success();
1188 }
1189 
1190 LogicalResult AffineDmaStartOp::verifyInvariantsImpl() {
1191   if (!getOperand(getSrcMemRefOperandIndex()).getType().isa<MemRefType>())
1192     return emitOpError("expected DMA source to be of memref type");
1193   if (!getOperand(getDstMemRefOperandIndex()).getType().isa<MemRefType>())
1194     return emitOpError("expected DMA destination to be of memref type");
1195   if (!getOperand(getTagMemRefOperandIndex()).getType().isa<MemRefType>())
1196     return emitOpError("expected DMA tag to be of memref type");
1197 
1198   unsigned numInputsAllMaps = getSrcMap().getNumInputs() +
1199                               getDstMap().getNumInputs() +
1200                               getTagMap().getNumInputs();
1201   if (getNumOperands() != numInputsAllMaps + 3 + 1 &&
1202       getNumOperands() != numInputsAllMaps + 3 + 1 + 2) {
1203     return emitOpError("incorrect number of operands");
1204   }
1205 
1206   Region *scope = getAffineScope(*this);
1207   for (auto idx : getSrcIndices()) {
1208     if (!idx.getType().isIndex())
1209       return emitOpError("src index to dma_start must have 'index' type");
1210     if (!isValidAffineIndexOperand(idx, scope))
1211       return emitOpError("src index must be a dimension or symbol identifier");
1212   }
1213   for (auto idx : getDstIndices()) {
1214     if (!idx.getType().isIndex())
1215       return emitOpError("dst index to dma_start must have 'index' type");
1216     if (!isValidAffineIndexOperand(idx, scope))
1217       return emitOpError("dst index must be a dimension or symbol identifier");
1218   }
1219   for (auto idx : getTagIndices()) {
1220     if (!idx.getType().isIndex())
1221       return emitOpError("tag index to dma_start must have 'index' type");
1222     if (!isValidAffineIndexOperand(idx, scope))
1223       return emitOpError("tag index must be a dimension or symbol identifier");
1224   }
1225   return success();
1226 }
1227 
1228 LogicalResult AffineDmaStartOp::fold(ArrayRef<Attribute> cstOperands,
1229                                      SmallVectorImpl<OpFoldResult> &results) {
1230   /// dma_start(memrefcast) -> dma_start
1231   return foldMemRefCast(*this);
1232 }
1233 
1234 //===----------------------------------------------------------------------===//
1235 // AffineDmaWaitOp
1236 //===----------------------------------------------------------------------===//
1237 
1238 // TODO: Check that map operands are loop IVs or symbols.
1239 void AffineDmaWaitOp::build(OpBuilder &builder, OperationState &result,
1240                             Value tagMemRef, AffineMap tagMap,
1241                             ValueRange tagIndices, Value numElements) {
1242   result.addOperands(tagMemRef);
1243   result.addAttribute(getTagMapAttrStrName(), AffineMapAttr::get(tagMap));
1244   result.addOperands(tagIndices);
1245   result.addOperands(numElements);
1246 }
1247 
1248 void AffineDmaWaitOp::print(OpAsmPrinter &p) {
1249   p << " " << getTagMemRef() << '[';
1250   SmallVector<Value, 2> operands(getTagIndices());
1251   p.printAffineMapOfSSAIds(getTagMapAttr(), operands);
1252   p << "], ";
1253   p.printOperand(getNumElements());
1254   p << " : " << getTagMemRef().getType();
1255 }
1256 
1257 // Parse AffineDmaWaitOp.
1258 // Eg:
1259 //   affine.dma_wait %tag[%index], %num_elements
1260 //     : memref<1 x i32, (d0) -> (d0), 4>
1261 //
1262 ParseResult AffineDmaWaitOp::parse(OpAsmParser &parser,
1263                                    OperationState &result) {
1264   OpAsmParser::UnresolvedOperand tagMemRefInfo;
1265   AffineMapAttr tagMapAttr;
1266   SmallVector<OpAsmParser::UnresolvedOperand, 2> tagMapOperands;
1267   Type type;
1268   auto indexType = parser.getBuilder().getIndexType();
1269   OpAsmParser::UnresolvedOperand numElementsInfo;
1270 
1271   // Parse tag memref, its map operands, and dma size.
1272   if (parser.parseOperand(tagMemRefInfo) ||
1273       parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr,
1274                                     getTagMapAttrStrName(),
1275                                     result.attributes) ||
1276       parser.parseComma() || parser.parseOperand(numElementsInfo) ||
1277       parser.parseColonType(type) ||
1278       parser.resolveOperand(tagMemRefInfo, type, result.operands) ||
1279       parser.resolveOperands(tagMapOperands, indexType, result.operands) ||
1280       parser.resolveOperand(numElementsInfo, indexType, result.operands))
1281     return failure();
1282 
1283   if (!type.isa<MemRefType>())
1284     return parser.emitError(parser.getNameLoc(),
1285                             "expected tag to be of memref type");
1286 
1287   if (tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
1288     return parser.emitError(parser.getNameLoc(),
1289                             "tag memref operand count != to map.numInputs");
1290   return success();
1291 }
1292 
1293 LogicalResult AffineDmaWaitOp::verifyInvariantsImpl() {
1294   if (!getOperand(0).getType().isa<MemRefType>())
1295     return emitOpError("expected DMA tag to be of memref type");
1296   Region *scope = getAffineScope(*this);
1297   for (auto idx : getTagIndices()) {
1298     if (!idx.getType().isIndex())
1299       return emitOpError("index to dma_wait must have 'index' type");
1300     if (!isValidAffineIndexOperand(idx, scope))
1301       return emitOpError("index must be a dimension or symbol identifier");
1302   }
1303   return success();
1304 }
1305 
1306 LogicalResult AffineDmaWaitOp::fold(ArrayRef<Attribute> cstOperands,
1307                                     SmallVectorImpl<OpFoldResult> &results) {
1308   /// dma_wait(memrefcast) -> dma_wait
1309   return foldMemRefCast(*this);
1310 }
1311 
1312 //===----------------------------------------------------------------------===//
1313 // AffineForOp
1314 //===----------------------------------------------------------------------===//
1315 
1316 /// 'bodyBuilder' is used to build the body of affine.for. If iterArgs and
1317 /// bodyBuilder are empty/null, we include default terminator op.
1318 void AffineForOp::build(OpBuilder &builder, OperationState &result,
1319                         ValueRange lbOperands, AffineMap lbMap,
1320                         ValueRange ubOperands, AffineMap ubMap, int64_t step,
1321                         ValueRange iterArgs, BodyBuilderFn bodyBuilder) {
1322   assert(((!lbMap && lbOperands.empty()) ||
1323           lbOperands.size() == lbMap.getNumInputs()) &&
1324          "lower bound operand count does not match the affine map");
1325   assert(((!ubMap && ubOperands.empty()) ||
1326           ubOperands.size() == ubMap.getNumInputs()) &&
1327          "upper bound operand count does not match the affine map");
1328   assert(step > 0 && "step has to be a positive integer constant");
1329 
1330   for (Value val : iterArgs)
1331     result.addTypes(val.getType());
1332 
1333   // Add an attribute for the step.
1334   result.addAttribute(getStepAttrStrName(),
1335                       builder.getIntegerAttr(builder.getIndexType(), step));
1336 
1337   // Add the lower bound.
1338   result.addAttribute(getLowerBoundAttrStrName(), AffineMapAttr::get(lbMap));
1339   result.addOperands(lbOperands);
1340 
1341   // Add the upper bound.
1342   result.addAttribute(getUpperBoundAttrStrName(), AffineMapAttr::get(ubMap));
1343   result.addOperands(ubOperands);
1344 
1345   result.addOperands(iterArgs);
1346   // Create a region and a block for the body.  The argument of the region is
1347   // the loop induction variable.
1348   Region *bodyRegion = result.addRegion();
1349   bodyRegion->push_back(new Block);
1350   Block &bodyBlock = bodyRegion->front();
1351   Value inductionVar =
1352       bodyBlock.addArgument(builder.getIndexType(), result.location);
1353   for (Value val : iterArgs)
1354     bodyBlock.addArgument(val.getType(), val.getLoc());
1355 
1356   // Create the default terminator if the builder is not provided and if the
1357   // iteration arguments are not provided. Otherwise, leave this to the caller
1358   // because we don't know which values to return from the loop.
1359   if (iterArgs.empty() && !bodyBuilder) {
1360     ensureTerminator(*bodyRegion, builder, result.location);
1361   } else if (bodyBuilder) {
1362     OpBuilder::InsertionGuard guard(builder);
1363     builder.setInsertionPointToStart(&bodyBlock);
1364     bodyBuilder(builder, result.location, inductionVar,
1365                 bodyBlock.getArguments().drop_front());
1366   }
1367 }
1368 
1369 void AffineForOp::build(OpBuilder &builder, OperationState &result, int64_t lb,
1370                         int64_t ub, int64_t step, ValueRange iterArgs,
1371                         BodyBuilderFn bodyBuilder) {
1372   auto lbMap = AffineMap::getConstantMap(lb, builder.getContext());
1373   auto ubMap = AffineMap::getConstantMap(ub, builder.getContext());
1374   return build(builder, result, {}, lbMap, {}, ubMap, step, iterArgs,
1375                bodyBuilder);
1376 }
1377 
1378 LogicalResult AffineForOp::verifyRegions() {
1379   // Check that the body defines as single block argument for the induction
1380   // variable.
1381   auto *body = getBody();
1382   if (body->getNumArguments() == 0 || !body->getArgument(0).getType().isIndex())
1383     return emitOpError("expected body to have a single index argument for the "
1384                        "induction variable");
1385 
1386   // Verify that the bound operands are valid dimension/symbols.
1387   /// Lower bound.
1388   if (getLowerBoundMap().getNumInputs() > 0)
1389     if (failed(verifyDimAndSymbolIdentifiers(*this, getLowerBoundOperands(),
1390                                              getLowerBoundMap().getNumDims())))
1391       return failure();
1392   /// Upper bound.
1393   if (getUpperBoundMap().getNumInputs() > 0)
1394     if (failed(verifyDimAndSymbolIdentifiers(*this, getUpperBoundOperands(),
1395                                              getUpperBoundMap().getNumDims())))
1396       return failure();
1397 
1398   unsigned opNumResults = getNumResults();
1399   if (opNumResults == 0)
1400     return success();
1401 
1402   // If ForOp defines values, check that the number and types of the defined
1403   // values match ForOp initial iter operands and backedge basic block
1404   // arguments.
1405   if (getNumIterOperands() != opNumResults)
1406     return emitOpError(
1407         "mismatch between the number of loop-carried values and results");
1408   if (getNumRegionIterArgs() != opNumResults)
1409     return emitOpError(
1410         "mismatch between the number of basic block args and results");
1411 
1412   return success();
1413 }
1414 
1415 /// Parse a for operation loop bounds.
1416 static ParseResult parseBound(bool isLower, OperationState &result,
1417                               OpAsmParser &p) {
1418   // 'min' / 'max' prefixes are generally syntactic sugar, but are required if
1419   // the map has multiple results.
1420   bool failedToParsedMinMax =
1421       failed(p.parseOptionalKeyword(isLower ? "max" : "min"));
1422 
1423   auto &builder = p.getBuilder();
1424   auto boundAttrStrName = isLower ? AffineForOp::getLowerBoundAttrStrName()
1425                                   : AffineForOp::getUpperBoundAttrStrName();
1426 
1427   // Parse ssa-id as identity map.
1428   SmallVector<OpAsmParser::UnresolvedOperand, 1> boundOpInfos;
1429   if (p.parseOperandList(boundOpInfos))
1430     return failure();
1431 
1432   if (!boundOpInfos.empty()) {
1433     // Check that only one operand was parsed.
1434     if (boundOpInfos.size() > 1)
1435       return p.emitError(p.getNameLoc(),
1436                          "expected only one loop bound operand");
1437 
1438     // TODO: improve error message when SSA value is not of index type.
1439     // Currently it is 'use of value ... expects different type than prior uses'
1440     if (p.resolveOperand(boundOpInfos.front(), builder.getIndexType(),
1441                          result.operands))
1442       return failure();
1443 
1444     // Create an identity map using symbol id. This representation is optimized
1445     // for storage. Analysis passes may expand it into a multi-dimensional map
1446     // if desired.
1447     AffineMap map = builder.getSymbolIdentityMap();
1448     result.addAttribute(boundAttrStrName, AffineMapAttr::get(map));
1449     return success();
1450   }
1451 
1452   // Get the attribute location.
1453   SMLoc attrLoc = p.getCurrentLocation();
1454 
1455   Attribute boundAttr;
1456   if (p.parseAttribute(boundAttr, builder.getIndexType(), boundAttrStrName,
1457                        result.attributes))
1458     return failure();
1459 
1460   // Parse full form - affine map followed by dim and symbol list.
1461   if (auto affineMapAttr = boundAttr.dyn_cast<AffineMapAttr>()) {
1462     unsigned currentNumOperands = result.operands.size();
1463     unsigned numDims;
1464     if (parseDimAndSymbolList(p, result.operands, numDims))
1465       return failure();
1466 
1467     auto map = affineMapAttr.getValue();
1468     if (map.getNumDims() != numDims)
1469       return p.emitError(
1470           p.getNameLoc(),
1471           "dim operand count and affine map dim count must match");
1472 
1473     unsigned numDimAndSymbolOperands =
1474         result.operands.size() - currentNumOperands;
1475     if (numDims + map.getNumSymbols() != numDimAndSymbolOperands)
1476       return p.emitError(
1477           p.getNameLoc(),
1478           "symbol operand count and affine map symbol count must match");
1479 
1480     // If the map has multiple results, make sure that we parsed the min/max
1481     // prefix.
1482     if (map.getNumResults() > 1 && failedToParsedMinMax) {
1483       if (isLower) {
1484         return p.emitError(attrLoc, "lower loop bound affine map with "
1485                                     "multiple results requires 'max' prefix");
1486       }
1487       return p.emitError(attrLoc, "upper loop bound affine map with multiple "
1488                                   "results requires 'min' prefix");
1489     }
1490     return success();
1491   }
1492 
1493   // Parse custom assembly form.
1494   if (auto integerAttr = boundAttr.dyn_cast<IntegerAttr>()) {
1495     result.attributes.pop_back();
1496     result.addAttribute(
1497         boundAttrStrName,
1498         AffineMapAttr::get(builder.getConstantAffineMap(integerAttr.getInt())));
1499     return success();
1500   }
1501 
1502   return p.emitError(
1503       p.getNameLoc(),
1504       "expected valid affine map representation for loop bounds");
1505 }
1506 
1507 ParseResult AffineForOp::parse(OpAsmParser &parser, OperationState &result) {
1508   auto &builder = parser.getBuilder();
1509   OpAsmParser::Argument inductionVariable;
1510   inductionVariable.type = builder.getIndexType();
1511   // Parse the induction variable followed by '='.
1512   if (parser.parseArgument(inductionVariable) || parser.parseEqual())
1513     return failure();
1514 
1515   // Parse loop bounds.
1516   if (parseBound(/*isLower=*/true, result, parser) ||
1517       parser.parseKeyword("to", " between bounds") ||
1518       parseBound(/*isLower=*/false, result, parser))
1519     return failure();
1520 
1521   // Parse the optional loop step, we default to 1 if one is not present.
1522   if (parser.parseOptionalKeyword("step")) {
1523     result.addAttribute(
1524         AffineForOp::getStepAttrStrName(),
1525         builder.getIntegerAttr(builder.getIndexType(), /*value=*/1));
1526   } else {
1527     SMLoc stepLoc = parser.getCurrentLocation();
1528     IntegerAttr stepAttr;
1529     if (parser.parseAttribute(stepAttr, builder.getIndexType(),
1530                               AffineForOp::getStepAttrStrName().data(),
1531                               result.attributes))
1532       return failure();
1533 
1534     if (stepAttr.getValue().getSExtValue() < 0)
1535       return parser.emitError(
1536           stepLoc,
1537           "expected step to be representable as a positive signed integer");
1538   }
1539 
1540   // Parse the optional initial iteration arguments.
1541   SmallVector<OpAsmParser::Argument, 4> regionArgs;
1542   SmallVector<OpAsmParser::UnresolvedOperand, 4> operands;
1543 
1544   // Induction variable.
1545   regionArgs.push_back(inductionVariable);
1546 
1547   if (succeeded(parser.parseOptionalKeyword("iter_args"))) {
1548     // Parse assignment list and results type list.
1549     if (parser.parseAssignmentList(regionArgs, operands) ||
1550         parser.parseArrowTypeList(result.types))
1551       return failure();
1552     // Resolve input operands.
1553     for (auto argOperandType :
1554          llvm::zip(llvm::drop_begin(regionArgs), operands, result.types)) {
1555       Type type = std::get<2>(argOperandType);
1556       std::get<0>(argOperandType).type = type;
1557       if (parser.resolveOperand(std::get<1>(argOperandType), type,
1558                                 result.operands))
1559         return failure();
1560     }
1561   }
1562 
1563   // Parse the body region.
1564   Region *body = result.addRegion();
1565   if (regionArgs.size() != result.types.size() + 1)
1566     return parser.emitError(
1567         parser.getNameLoc(),
1568         "mismatch between the number of loop-carried values and results");
1569   if (parser.parseRegion(*body, regionArgs))
1570     return failure();
1571 
1572   AffineForOp::ensureTerminator(*body, builder, result.location);
1573 
1574   // Parse the optional attribute list.
1575   return parser.parseOptionalAttrDict(result.attributes);
1576 }
1577 
1578 static void printBound(AffineMapAttr boundMap,
1579                        Operation::operand_range boundOperands,
1580                        const char *prefix, OpAsmPrinter &p) {
1581   AffineMap map = boundMap.getValue();
1582 
1583   // Check if this bound should be printed using custom assembly form.
1584   // The decision to restrict printing custom assembly form to trivial cases
1585   // comes from the will to roundtrip MLIR binary -> text -> binary in a
1586   // lossless way.
1587   // Therefore, custom assembly form parsing and printing is only supported for
1588   // zero-operand constant maps and single symbol operand identity maps.
1589   if (map.getNumResults() == 1) {
1590     AffineExpr expr = map.getResult(0);
1591 
1592     // Print constant bound.
1593     if (map.getNumDims() == 0 && map.getNumSymbols() == 0) {
1594       if (auto constExpr = expr.dyn_cast<AffineConstantExpr>()) {
1595         p << constExpr.getValue();
1596         return;
1597       }
1598     }
1599 
1600     // Print bound that consists of a single SSA symbol if the map is over a
1601     // single symbol.
1602     if (map.getNumDims() == 0 && map.getNumSymbols() == 1) {
1603       if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>()) {
1604         p.printOperand(*boundOperands.begin());
1605         return;
1606       }
1607     }
1608   } else {
1609     // Map has multiple results. Print 'min' or 'max' prefix.
1610     p << prefix << ' ';
1611   }
1612 
1613   // Print the map and its operands.
1614   p << boundMap;
1615   printDimAndSymbolList(boundOperands.begin(), boundOperands.end(),
1616                         map.getNumDims(), p);
1617 }
1618 
1619 unsigned AffineForOp::getNumIterOperands() {
1620   AffineMap lbMap = getLowerBoundMapAttr().getValue();
1621   AffineMap ubMap = getUpperBoundMapAttr().getValue();
1622 
1623   return getNumOperands() - lbMap.getNumInputs() - ubMap.getNumInputs();
1624 }
1625 
1626 void AffineForOp::print(OpAsmPrinter &p) {
1627   p << ' ';
1628   p.printRegionArgument(getBody()->getArgument(0), /*argAttrs=*/{},
1629                         /*omitType=*/true);
1630   p << " = ";
1631   printBound(getLowerBoundMapAttr(), getLowerBoundOperands(), "max", p);
1632   p << " to ";
1633   printBound(getUpperBoundMapAttr(), getUpperBoundOperands(), "min", p);
1634 
1635   if (getStep() != 1)
1636     p << " step " << getStep();
1637 
1638   bool printBlockTerminators = false;
1639   if (getNumIterOperands() > 0) {
1640     p << " iter_args(";
1641     auto regionArgs = getRegionIterArgs();
1642     auto operands = getIterOperands();
1643 
1644     llvm::interleaveComma(llvm::zip(regionArgs, operands), p, [&](auto it) {
1645       p << std::get<0>(it) << " = " << std::get<1>(it);
1646     });
1647     p << ") -> (" << getResultTypes() << ")";
1648     printBlockTerminators = true;
1649   }
1650 
1651   p << ' ';
1652   p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
1653                 printBlockTerminators);
1654   p.printOptionalAttrDict((*this)->getAttrs(),
1655                           /*elidedAttrs=*/{getLowerBoundAttrStrName(),
1656                                            getUpperBoundAttrStrName(),
1657                                            getStepAttrStrName()});
1658 }
1659 
1660 /// Fold the constant bounds of a loop.
1661 static LogicalResult foldLoopBounds(AffineForOp forOp) {
1662   auto foldLowerOrUpperBound = [&forOp](bool lower) {
1663     // Check to see if each of the operands is the result of a constant.  If
1664     // so, get the value.  If not, ignore it.
1665     SmallVector<Attribute, 8> operandConstants;
1666     auto boundOperands =
1667         lower ? forOp.getLowerBoundOperands() : forOp.getUpperBoundOperands();
1668     for (auto operand : boundOperands) {
1669       Attribute operandCst;
1670       matchPattern(operand, m_Constant(&operandCst));
1671       operandConstants.push_back(operandCst);
1672     }
1673 
1674     AffineMap boundMap =
1675         lower ? forOp.getLowerBoundMap() : forOp.getUpperBoundMap();
1676     assert(boundMap.getNumResults() >= 1 &&
1677            "bound maps should have at least one result");
1678     SmallVector<Attribute, 4> foldedResults;
1679     if (failed(boundMap.constantFold(operandConstants, foldedResults)))
1680       return failure();
1681 
1682     // Compute the max or min as applicable over the results.
1683     assert(!foldedResults.empty() && "bounds should have at least one result");
1684     auto maxOrMin = foldedResults[0].cast<IntegerAttr>().getValue();
1685     for (unsigned i = 1, e = foldedResults.size(); i < e; i++) {
1686       auto foldedResult = foldedResults[i].cast<IntegerAttr>().getValue();
1687       maxOrMin = lower ? llvm::APIntOps::smax(maxOrMin, foldedResult)
1688                        : llvm::APIntOps::smin(maxOrMin, foldedResult);
1689     }
1690     lower ? forOp.setConstantLowerBound(maxOrMin.getSExtValue())
1691           : forOp.setConstantUpperBound(maxOrMin.getSExtValue());
1692     return success();
1693   };
1694 
1695   // Try to fold the lower bound.
1696   bool folded = false;
1697   if (!forOp.hasConstantLowerBound())
1698     folded |= succeeded(foldLowerOrUpperBound(/*lower=*/true));
1699 
1700   // Try to fold the upper bound.
1701   if (!forOp.hasConstantUpperBound())
1702     folded |= succeeded(foldLowerOrUpperBound(/*lower=*/false));
1703   return success(folded);
1704 }
1705 
1706 /// Canonicalize the bounds of the given loop.
1707 static LogicalResult canonicalizeLoopBounds(AffineForOp forOp) {
1708   SmallVector<Value, 4> lbOperands(forOp.getLowerBoundOperands());
1709   SmallVector<Value, 4> ubOperands(forOp.getUpperBoundOperands());
1710 
1711   auto lbMap = forOp.getLowerBoundMap();
1712   auto ubMap = forOp.getUpperBoundMap();
1713   auto prevLbMap = lbMap;
1714   auto prevUbMap = ubMap;
1715 
1716   composeAffineMapAndOperands(&lbMap, &lbOperands);
1717   canonicalizeMapAndOperands(&lbMap, &lbOperands);
1718   lbMap = removeDuplicateExprs(lbMap);
1719 
1720   composeAffineMapAndOperands(&ubMap, &ubOperands);
1721   canonicalizeMapAndOperands(&ubMap, &ubOperands);
1722   ubMap = removeDuplicateExprs(ubMap);
1723 
1724   // Any canonicalization change always leads to updated map(s).
1725   if (lbMap == prevLbMap && ubMap == prevUbMap)
1726     return failure();
1727 
1728   if (lbMap != prevLbMap)
1729     forOp.setLowerBound(lbOperands, lbMap);
1730   if (ubMap != prevUbMap)
1731     forOp.setUpperBound(ubOperands, ubMap);
1732   return success();
1733 }
1734 
1735 namespace {
1736 /// Returns constant trip count in trivial cases.
1737 static Optional<uint64_t> getTrivialConstantTripCount(AffineForOp forOp) {
1738   int64_t step = forOp.getStep();
1739   if (!forOp.hasConstantBounds() || step <= 0)
1740     return None;
1741   int64_t lb = forOp.getConstantLowerBound();
1742   int64_t ub = forOp.getConstantUpperBound();
1743   return ub - lb <= 0 ? 0 : (ub - lb + step - 1) / step;
1744 }
1745 
1746 /// This is a pattern to fold trivially empty loop bodies.
1747 /// TODO: This should be moved into the folding hook.
1748 struct AffineForEmptyLoopFolder : public OpRewritePattern<AffineForOp> {
1749   using OpRewritePattern<AffineForOp>::OpRewritePattern;
1750 
1751   LogicalResult matchAndRewrite(AffineForOp forOp,
1752                                 PatternRewriter &rewriter) const override {
1753     // Check that the body only contains a yield.
1754     if (!llvm::hasSingleElement(*forOp.getBody()))
1755       return failure();
1756     if (forOp.getNumResults() == 0)
1757       return success();
1758     Optional<uint64_t> tripCount = getTrivialConstantTripCount(forOp);
1759     if (tripCount && *tripCount == 0) {
1760       // The initial values of the iteration arguments would be the op's
1761       // results.
1762       rewriter.replaceOp(forOp, forOp.getIterOperands());
1763       return success();
1764     }
1765     SmallVector<Value, 4> replacements;
1766     auto yieldOp = cast<AffineYieldOp>(forOp.getBody()->getTerminator());
1767     auto iterArgs = forOp.getRegionIterArgs();
1768     bool hasValDefinedOutsideLoop = false;
1769     bool iterArgsNotInOrder = false;
1770     for (unsigned i = 0, e = yieldOp->getNumOperands(); i < e; ++i) {
1771       Value val = yieldOp.getOperand(i);
1772       auto *iterArgIt = llvm::find(iterArgs, val);
1773       if (iterArgIt == iterArgs.end()) {
1774         // `val` is defined outside of the loop.
1775         assert(forOp.isDefinedOutsideOfLoop(val) &&
1776                "must be defined outside of the loop");
1777         hasValDefinedOutsideLoop = true;
1778         replacements.push_back(val);
1779       } else {
1780         unsigned pos = std::distance(iterArgs.begin(), iterArgIt);
1781         if (pos != i)
1782           iterArgsNotInOrder = true;
1783         replacements.push_back(forOp.getIterOperands()[pos]);
1784       }
1785     }
1786     // Bail out when the trip count is unknown and the loop returns any value
1787     // defined outside of the loop or any iterArg out of order.
1788     if (!tripCount.hasValue() &&
1789         (hasValDefinedOutsideLoop || iterArgsNotInOrder))
1790       return failure();
1791     // Bail out when the loop iterates more than once and it returns any iterArg
1792     // out of order.
1793     if (tripCount.hasValue() && tripCount.getValue() >= 2 && iterArgsNotInOrder)
1794       return failure();
1795     rewriter.replaceOp(forOp, replacements);
1796     return success();
1797   }
1798 };
1799 } // namespace
1800 
1801 void AffineForOp::getCanonicalizationPatterns(RewritePatternSet &results,
1802                                               MLIRContext *context) {
1803   results.add<AffineForEmptyLoopFolder>(context);
1804 }
1805 
1806 /// Return operands used when entering the region at 'index'. These operands
1807 /// correspond to the loop iterator operands, i.e., those excluding the
1808 /// induction variable. AffineForOp only has one region, so zero is the only
1809 /// valid value for `index`.
1810 OperandRange AffineForOp::getSuccessorEntryOperands(Optional<unsigned> index) {
1811   assert((!index || *index == 0) && "invalid region index");
1812 
1813   // The initial operands map to the loop arguments after the induction
1814   // variable or are forwarded to the results when the trip count is zero.
1815   return getIterOperands();
1816 }
1817 
1818 /// Given the region at `index`, or the parent operation if `index` is None,
1819 /// return the successor regions. These are the regions that may be selected
1820 /// during the flow of control. `operands` is a set of optional attributes that
1821 /// correspond to a constant value for each operand, or null if that operand is
1822 /// not a constant.
1823 void AffineForOp::getSuccessorRegions(
1824     Optional<unsigned> index, ArrayRef<Attribute> operands,
1825     SmallVectorImpl<RegionSuccessor> &regions) {
1826   assert((!index.hasValue() || index.getValue() == 0) &&
1827          "expected loop region");
1828   // The loop may typically branch back to its body or to the parent operation.
1829   // If the predecessor is the parent op and the trip count is known to be at
1830   // least one, branch into the body using the iterator arguments. And in cases
1831   // we know the trip count is zero, it can only branch back to its parent.
1832   Optional<uint64_t> tripCount = getTrivialConstantTripCount(*this);
1833   if (!index.hasValue() && tripCount.hasValue()) {
1834     if (tripCount.getValue() > 0) {
1835       regions.push_back(RegionSuccessor(&getLoopBody(), getRegionIterArgs()));
1836       return;
1837     }
1838     if (tripCount.getValue() == 0) {
1839       regions.push_back(RegionSuccessor(getResults()));
1840       return;
1841     }
1842   }
1843 
1844   // From the loop body, if the trip count is one, we can only branch back to
1845   // the parent.
1846   if (index && tripCount && *tripCount == 1) {
1847     regions.push_back(RegionSuccessor(getResults()));
1848     return;
1849   }
1850 
1851   // In all other cases, the loop may branch back to itself or the parent
1852   // operation.
1853   regions.push_back(RegionSuccessor(&getLoopBody(), getRegionIterArgs()));
1854   regions.push_back(RegionSuccessor(getResults()));
1855 }
1856 
1857 /// Returns true if the affine.for has zero iterations in trivial cases.
1858 static bool hasTrivialZeroTripCount(AffineForOp op) {
1859   Optional<uint64_t> tripCount = getTrivialConstantTripCount(op);
1860   return tripCount && *tripCount == 0;
1861 }
1862 
1863 LogicalResult AffineForOp::fold(ArrayRef<Attribute> operands,
1864                                 SmallVectorImpl<OpFoldResult> &results) {
1865   bool folded = succeeded(foldLoopBounds(*this));
1866   folded |= succeeded(canonicalizeLoopBounds(*this));
1867   if (hasTrivialZeroTripCount(*this)) {
1868     // The initial values of the loop-carried variables (iter_args) are the
1869     // results of the op.
1870     results.assign(getIterOperands().begin(), getIterOperands().end());
1871     folded = true;
1872   }
1873   return success(folded);
1874 }
1875 
1876 AffineBound AffineForOp::getLowerBound() {
1877   auto lbMap = getLowerBoundMap();
1878   return AffineBound(AffineForOp(*this), 0, lbMap.getNumInputs(), lbMap);
1879 }
1880 
1881 AffineBound AffineForOp::getUpperBound() {
1882   auto lbMap = getLowerBoundMap();
1883   auto ubMap = getUpperBoundMap();
1884   return AffineBound(AffineForOp(*this), lbMap.getNumInputs(),
1885                      lbMap.getNumInputs() + ubMap.getNumInputs(), ubMap);
1886 }
1887 
1888 void AffineForOp::setLowerBound(ValueRange lbOperands, AffineMap map) {
1889   assert(lbOperands.size() == map.getNumInputs());
1890   assert(map.getNumResults() >= 1 && "bound map has at least one result");
1891 
1892   SmallVector<Value, 4> newOperands(lbOperands.begin(), lbOperands.end());
1893 
1894   auto ubOperands = getUpperBoundOperands();
1895   newOperands.append(ubOperands.begin(), ubOperands.end());
1896   auto iterOperands = getIterOperands();
1897   newOperands.append(iterOperands.begin(), iterOperands.end());
1898   (*this)->setOperands(newOperands);
1899 
1900   (*this)->setAttr(getLowerBoundAttrStrName(), AffineMapAttr::get(map));
1901 }
1902 
1903 void AffineForOp::setUpperBound(ValueRange ubOperands, AffineMap map) {
1904   assert(ubOperands.size() == map.getNumInputs());
1905   assert(map.getNumResults() >= 1 && "bound map has at least one result");
1906 
1907   SmallVector<Value, 4> newOperands(getLowerBoundOperands());
1908   newOperands.append(ubOperands.begin(), ubOperands.end());
1909   auto iterOperands = getIterOperands();
1910   newOperands.append(iterOperands.begin(), iterOperands.end());
1911   (*this)->setOperands(newOperands);
1912 
1913   (*this)->setAttr(getUpperBoundAttrStrName(), AffineMapAttr::get(map));
1914 }
1915 
1916 void AffineForOp::setLowerBoundMap(AffineMap map) {
1917   auto lbMap = getLowerBoundMap();
1918   assert(lbMap.getNumDims() == map.getNumDims() &&
1919          lbMap.getNumSymbols() == map.getNumSymbols());
1920   assert(map.getNumResults() >= 1 && "bound map has at least one result");
1921   (void)lbMap;
1922   (*this)->setAttr(getLowerBoundAttrStrName(), AffineMapAttr::get(map));
1923 }
1924 
1925 void AffineForOp::setUpperBoundMap(AffineMap map) {
1926   auto ubMap = getUpperBoundMap();
1927   assert(ubMap.getNumDims() == map.getNumDims() &&
1928          ubMap.getNumSymbols() == map.getNumSymbols());
1929   assert(map.getNumResults() >= 1 && "bound map has at least one result");
1930   (void)ubMap;
1931   (*this)->setAttr(getUpperBoundAttrStrName(), AffineMapAttr::get(map));
1932 }
1933 
1934 bool AffineForOp::hasConstantLowerBound() {
1935   return getLowerBoundMap().isSingleConstant();
1936 }
1937 
1938 bool AffineForOp::hasConstantUpperBound() {
1939   return getUpperBoundMap().isSingleConstant();
1940 }
1941 
1942 int64_t AffineForOp::getConstantLowerBound() {
1943   return getLowerBoundMap().getSingleConstantResult();
1944 }
1945 
1946 int64_t AffineForOp::getConstantUpperBound() {
1947   return getUpperBoundMap().getSingleConstantResult();
1948 }
1949 
1950 void AffineForOp::setConstantLowerBound(int64_t value) {
1951   setLowerBound({}, AffineMap::getConstantMap(value, getContext()));
1952 }
1953 
1954 void AffineForOp::setConstantUpperBound(int64_t value) {
1955   setUpperBound({}, AffineMap::getConstantMap(value, getContext()));
1956 }
1957 
1958 AffineForOp::operand_range AffineForOp::getLowerBoundOperands() {
1959   return {operand_begin(), operand_begin() + getLowerBoundMap().getNumInputs()};
1960 }
1961 
1962 AffineForOp::operand_range AffineForOp::getUpperBoundOperands() {
1963   return {operand_begin() + getLowerBoundMap().getNumInputs(),
1964           operand_begin() + getLowerBoundMap().getNumInputs() +
1965               getUpperBoundMap().getNumInputs()};
1966 }
1967 
1968 AffineForOp::operand_range AffineForOp::getControlOperands() {
1969   return {operand_begin(), operand_begin() + getLowerBoundMap().getNumInputs() +
1970                                getUpperBoundMap().getNumInputs()};
1971 }
1972 
1973 bool AffineForOp::matchingBoundOperandList() {
1974   auto lbMap = getLowerBoundMap();
1975   auto ubMap = getUpperBoundMap();
1976   if (lbMap.getNumDims() != ubMap.getNumDims() ||
1977       lbMap.getNumSymbols() != ubMap.getNumSymbols())
1978     return false;
1979 
1980   unsigned numOperands = lbMap.getNumInputs();
1981   for (unsigned i = 0, e = lbMap.getNumInputs(); i < e; i++) {
1982     // Compare Value 's.
1983     if (getOperand(i) != getOperand(numOperands + i))
1984       return false;
1985   }
1986   return true;
1987 }
1988 
1989 Region &AffineForOp::getLoopBody() { return getRegion(); }
1990 
1991 Optional<Value> AffineForOp::getSingleInductionVar() {
1992   return getInductionVar();
1993 }
1994 
1995 Optional<OpFoldResult> AffineForOp::getSingleLowerBound() {
1996   if (!hasConstantLowerBound())
1997     return llvm::None;
1998   OpBuilder b(getContext());
1999   return OpFoldResult(b.getI64IntegerAttr(getConstantLowerBound()));
2000 }
2001 
2002 Optional<OpFoldResult> AffineForOp::getSingleStep() {
2003   OpBuilder b(getContext());
2004   return OpFoldResult(b.getI64IntegerAttr(getStep()));
2005 }
2006 
2007 Optional<OpFoldResult> AffineForOp::getSingleUpperBound() {
2008   if (!hasConstantUpperBound())
2009     return llvm::None;
2010   OpBuilder b(getContext());
2011   return OpFoldResult(b.getI64IntegerAttr(getConstantUpperBound()));
2012 }
2013 
2014 /// Returns true if the provided value is the induction variable of a
2015 /// AffineForOp.
2016 bool mlir::isForInductionVar(Value val) {
2017   return getForInductionVarOwner(val) != AffineForOp();
2018 }
2019 
2020 /// Returns the loop parent of an induction variable. If the provided value is
2021 /// not an induction variable, then return nullptr.
2022 AffineForOp mlir::getForInductionVarOwner(Value val) {
2023   auto ivArg = val.dyn_cast<BlockArgument>();
2024   if (!ivArg || !ivArg.getOwner())
2025     return AffineForOp();
2026   auto *containingInst = ivArg.getOwner()->getParent()->getParentOp();
2027   if (auto forOp = dyn_cast<AffineForOp>(containingInst))
2028     // Check to make sure `val` is the induction variable, not an iter_arg.
2029     return forOp.getInductionVar() == val ? forOp : AffineForOp();
2030   return AffineForOp();
2031 }
2032 
2033 /// Extracts the induction variables from a list of AffineForOps and returns
2034 /// them.
2035 void mlir::extractForInductionVars(ArrayRef<AffineForOp> forInsts,
2036                                    SmallVectorImpl<Value> *ivs) {
2037   ivs->reserve(forInsts.size());
2038   for (auto forInst : forInsts)
2039     ivs->push_back(forInst.getInductionVar());
2040 }
2041 
2042 /// Builds an affine loop nest, using "loopCreatorFn" to create individual loop
2043 /// operations.
2044 template <typename BoundListTy, typename LoopCreatorTy>
2045 static void buildAffineLoopNestImpl(
2046     OpBuilder &builder, Location loc, BoundListTy lbs, BoundListTy ubs,
2047     ArrayRef<int64_t> steps,
2048     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn,
2049     LoopCreatorTy &&loopCreatorFn) {
2050   assert(lbs.size() == ubs.size() && "Mismatch in number of arguments");
2051   assert(lbs.size() == steps.size() && "Mismatch in number of arguments");
2052 
2053   // If there are no loops to be constructed, construct the body anyway.
2054   OpBuilder::InsertionGuard guard(builder);
2055   if (lbs.empty()) {
2056     if (bodyBuilderFn)
2057       bodyBuilderFn(builder, loc, ValueRange());
2058     return;
2059   }
2060 
2061   // Create the loops iteratively and store the induction variables.
2062   SmallVector<Value, 4> ivs;
2063   ivs.reserve(lbs.size());
2064   for (unsigned i = 0, e = lbs.size(); i < e; ++i) {
2065     // Callback for creating the loop body, always creates the terminator.
2066     auto loopBody = [&](OpBuilder &nestedBuilder, Location nestedLoc, Value iv,
2067                         ValueRange iterArgs) {
2068       ivs.push_back(iv);
2069       // In the innermost loop, call the body builder.
2070       if (i == e - 1 && bodyBuilderFn) {
2071         OpBuilder::InsertionGuard nestedGuard(nestedBuilder);
2072         bodyBuilderFn(nestedBuilder, nestedLoc, ivs);
2073       }
2074       nestedBuilder.create<AffineYieldOp>(nestedLoc);
2075     };
2076 
2077     // Delegate actual loop creation to the callback in order to dispatch
2078     // between constant- and variable-bound loops.
2079     auto loop = loopCreatorFn(builder, loc, lbs[i], ubs[i], steps[i], loopBody);
2080     builder.setInsertionPointToStart(loop.getBody());
2081   }
2082 }
2083 
2084 /// Creates an affine loop from the bounds known to be constants.
2085 static AffineForOp
2086 buildAffineLoopFromConstants(OpBuilder &builder, Location loc, int64_t lb,
2087                              int64_t ub, int64_t step,
2088                              AffineForOp::BodyBuilderFn bodyBuilderFn) {
2089   return builder.create<AffineForOp>(loc, lb, ub, step, /*iterArgs=*/llvm::None,
2090                                      bodyBuilderFn);
2091 }
2092 
2093 /// Creates an affine loop from the bounds that may or may not be constants.
2094 static AffineForOp
2095 buildAffineLoopFromValues(OpBuilder &builder, Location loc, Value lb, Value ub,
2096                           int64_t step,
2097                           AffineForOp::BodyBuilderFn bodyBuilderFn) {
2098   auto lbConst = lb.getDefiningOp<arith::ConstantIndexOp>();
2099   auto ubConst = ub.getDefiningOp<arith::ConstantIndexOp>();
2100   if (lbConst && ubConst)
2101     return buildAffineLoopFromConstants(builder, loc, lbConst.value(),
2102                                         ubConst.value(), step, bodyBuilderFn);
2103   return builder.create<AffineForOp>(loc, lb, builder.getDimIdentityMap(), ub,
2104                                      builder.getDimIdentityMap(), step,
2105                                      /*iterArgs=*/llvm::None, bodyBuilderFn);
2106 }
2107 
2108 void mlir::buildAffineLoopNest(
2109     OpBuilder &builder, Location loc, ArrayRef<int64_t> lbs,
2110     ArrayRef<int64_t> ubs, ArrayRef<int64_t> steps,
2111     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn) {
2112   buildAffineLoopNestImpl(builder, loc, lbs, ubs, steps, bodyBuilderFn,
2113                           buildAffineLoopFromConstants);
2114 }
2115 
2116 void mlir::buildAffineLoopNest(
2117     OpBuilder &builder, Location loc, ValueRange lbs, ValueRange ubs,
2118     ArrayRef<int64_t> steps,
2119     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn) {
2120   buildAffineLoopNestImpl(builder, loc, lbs, ubs, steps, bodyBuilderFn,
2121                           buildAffineLoopFromValues);
2122 }
2123 
2124 AffineForOp mlir::replaceForOpWithNewYields(OpBuilder &b, AffineForOp loop,
2125                                             ValueRange newIterOperands,
2126                                             ValueRange newYieldedValues,
2127                                             ValueRange newIterArgs,
2128                                             bool replaceLoopResults) {
2129   assert(newIterOperands.size() == newYieldedValues.size() &&
2130          "newIterOperands must be of the same size as newYieldedValues");
2131   // Create a new loop before the existing one, with the extra operands.
2132   OpBuilder::InsertionGuard g(b);
2133   b.setInsertionPoint(loop);
2134   auto operands = llvm::to_vector<4>(loop.getIterOperands());
2135   operands.append(newIterOperands.begin(), newIterOperands.end());
2136   SmallVector<Value, 4> lbOperands(loop.getLowerBoundOperands());
2137   SmallVector<Value, 4> ubOperands(loop.getUpperBoundOperands());
2138   SmallVector<Value, 4> steps(loop.getStep());
2139   auto lbMap = loop.getLowerBoundMap();
2140   auto ubMap = loop.getUpperBoundMap();
2141   AffineForOp newLoop =
2142       b.create<AffineForOp>(loop.getLoc(), lbOperands, lbMap, ubOperands, ubMap,
2143                             loop.getStep(), operands);
2144   // Take the body of the original parent loop.
2145   newLoop.getLoopBody().takeBody(loop.getLoopBody());
2146   for (Value val : newIterArgs)
2147     newLoop.getLoopBody().addArgument(val.getType(), val.getLoc());
2148 
2149   // Update yield operation with new values to be added.
2150   if (!newYieldedValues.empty()) {
2151     auto yield = cast<AffineYieldOp>(newLoop.getBody()->getTerminator());
2152     b.setInsertionPoint(yield);
2153     auto yieldOperands = llvm::to_vector<4>(yield.getOperands());
2154     yieldOperands.append(newYieldedValues.begin(), newYieldedValues.end());
2155     b.create<AffineYieldOp>(yield.getLoc(), yieldOperands);
2156     yield.erase();
2157   }
2158   if (replaceLoopResults) {
2159     for (auto it : llvm::zip(loop.getResults(), newLoop.getResults().take_front(
2160                                                     loop.getNumResults()))) {
2161       std::get<0>(it).replaceAllUsesWith(std::get<1>(it));
2162     }
2163   }
2164   return newLoop;
2165 }
2166 
2167 //===----------------------------------------------------------------------===//
2168 // AffineIfOp
2169 //===----------------------------------------------------------------------===//
2170 
2171 namespace {
2172 /// Remove else blocks that have nothing other than a zero value yield.
2173 struct SimplifyDeadElse : public OpRewritePattern<AffineIfOp> {
2174   using OpRewritePattern<AffineIfOp>::OpRewritePattern;
2175 
2176   LogicalResult matchAndRewrite(AffineIfOp ifOp,
2177                                 PatternRewriter &rewriter) const override {
2178     if (ifOp.getElseRegion().empty() ||
2179         !llvm::hasSingleElement(*ifOp.getElseBlock()) || ifOp.getNumResults())
2180       return failure();
2181 
2182     rewriter.startRootUpdate(ifOp);
2183     rewriter.eraseBlock(ifOp.getElseBlock());
2184     rewriter.finalizeRootUpdate(ifOp);
2185     return success();
2186   }
2187 };
2188 
2189 /// Removes affine.if cond if the condition is always true or false in certain
2190 /// trivial cases. Promotes the then/else block in the parent operation block.
2191 struct AlwaysTrueOrFalseIf : public OpRewritePattern<AffineIfOp> {
2192   using OpRewritePattern<AffineIfOp>::OpRewritePattern;
2193 
2194   LogicalResult matchAndRewrite(AffineIfOp op,
2195                                 PatternRewriter &rewriter) const override {
2196 
2197     auto isTriviallyFalse = [](IntegerSet iSet) {
2198       return iSet.isEmptyIntegerSet();
2199     };
2200 
2201     auto isTriviallyTrue = [](IntegerSet iSet) {
2202       return (iSet.getNumEqualities() == 1 && iSet.getNumInequalities() == 0 &&
2203               iSet.getConstraint(0) == 0);
2204     };
2205 
2206     IntegerSet affineIfConditions = op.getIntegerSet();
2207     Block *blockToMove;
2208     if (isTriviallyFalse(affineIfConditions)) {
2209       // The absence, or equivalently, the emptiness of the else region need not
2210       // be checked when affine.if is returning results because if an affine.if
2211       // operation is returning results, it always has a non-empty else region.
2212       if (op.getNumResults() == 0 && !op.hasElse()) {
2213         // If the else region is absent, or equivalently, empty, remove the
2214         // affine.if operation (which is not returning any results).
2215         rewriter.eraseOp(op);
2216         return success();
2217       }
2218       blockToMove = op.getElseBlock();
2219     } else if (isTriviallyTrue(affineIfConditions)) {
2220       blockToMove = op.getThenBlock();
2221     } else {
2222       return failure();
2223     }
2224     Operation *blockToMoveTerminator = blockToMove->getTerminator();
2225     // Promote the "blockToMove" block to the parent operation block between the
2226     // prologue and epilogue of "op".
2227     rewriter.mergeBlockBefore(blockToMove, op);
2228     // Replace the "op" operation with the operands of the
2229     // "blockToMoveTerminator" operation. Note that "blockToMoveTerminator" is
2230     // the affine.yield operation present in the "blockToMove" block. It has no
2231     // operands when affine.if is not returning results and therefore, in that
2232     // case, replaceOp just erases "op". When affine.if is not returning
2233     // results, the affine.yield operation can be omitted. It gets inserted
2234     // implicitly.
2235     rewriter.replaceOp(op, blockToMoveTerminator->getOperands());
2236     // Erase the "blockToMoveTerminator" operation since it is now in the parent
2237     // operation block, which already has its own terminator.
2238     rewriter.eraseOp(blockToMoveTerminator);
2239     return success();
2240   }
2241 };
2242 } // namespace
2243 
2244 LogicalResult AffineIfOp::verify() {
2245   // Verify that we have a condition attribute.
2246   // FIXME: This should be specified in the arguments list in ODS.
2247   auto conditionAttr =
2248       (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
2249   if (!conditionAttr)
2250     return emitOpError("requires an integer set attribute named 'condition'");
2251 
2252   // Verify that there are enough operands for the condition.
2253   IntegerSet condition = conditionAttr.getValue();
2254   if (getNumOperands() != condition.getNumInputs())
2255     return emitOpError("operand count and condition integer set dimension and "
2256                        "symbol count must match");
2257 
2258   // Verify that the operands are valid dimension/symbols.
2259   if (failed(verifyDimAndSymbolIdentifiers(*this, getOperands(),
2260                                            condition.getNumDims())))
2261     return failure();
2262 
2263   return success();
2264 }
2265 
2266 ParseResult AffineIfOp::parse(OpAsmParser &parser, OperationState &result) {
2267   // Parse the condition attribute set.
2268   IntegerSetAttr conditionAttr;
2269   unsigned numDims;
2270   if (parser.parseAttribute(conditionAttr,
2271                             AffineIfOp::getConditionAttrStrName(),
2272                             result.attributes) ||
2273       parseDimAndSymbolList(parser, result.operands, numDims))
2274     return failure();
2275 
2276   // Verify the condition operands.
2277   auto set = conditionAttr.getValue();
2278   if (set.getNumDims() != numDims)
2279     return parser.emitError(
2280         parser.getNameLoc(),
2281         "dim operand count and integer set dim count must match");
2282   if (numDims + set.getNumSymbols() != result.operands.size())
2283     return parser.emitError(
2284         parser.getNameLoc(),
2285         "symbol operand count and integer set symbol count must match");
2286 
2287   if (parser.parseOptionalArrowTypeList(result.types))
2288     return failure();
2289 
2290   // Create the regions for 'then' and 'else'.  The latter must be created even
2291   // if it remains empty for the validity of the operation.
2292   result.regions.reserve(2);
2293   Region *thenRegion = result.addRegion();
2294   Region *elseRegion = result.addRegion();
2295 
2296   // Parse the 'then' region.
2297   if (parser.parseRegion(*thenRegion, {}, {}))
2298     return failure();
2299   AffineIfOp::ensureTerminator(*thenRegion, parser.getBuilder(),
2300                                result.location);
2301 
2302   // If we find an 'else' keyword then parse the 'else' region.
2303   if (!parser.parseOptionalKeyword("else")) {
2304     if (parser.parseRegion(*elseRegion, {}, {}))
2305       return failure();
2306     AffineIfOp::ensureTerminator(*elseRegion, parser.getBuilder(),
2307                                  result.location);
2308   }
2309 
2310   // Parse the optional attribute list.
2311   if (parser.parseOptionalAttrDict(result.attributes))
2312     return failure();
2313 
2314   return success();
2315 }
2316 
2317 void AffineIfOp::print(OpAsmPrinter &p) {
2318   auto conditionAttr =
2319       (*this)->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName());
2320   p << " " << conditionAttr;
2321   printDimAndSymbolList(operand_begin(), operand_end(),
2322                         conditionAttr.getValue().getNumDims(), p);
2323   p.printOptionalArrowTypeList(getResultTypes());
2324   p << ' ';
2325   p.printRegion(getThenRegion(), /*printEntryBlockArgs=*/false,
2326                 /*printBlockTerminators=*/getNumResults());
2327 
2328   // Print the 'else' regions if it has any blocks.
2329   auto &elseRegion = this->getElseRegion();
2330   if (!elseRegion.empty()) {
2331     p << " else ";
2332     p.printRegion(elseRegion,
2333                   /*printEntryBlockArgs=*/false,
2334                   /*printBlockTerminators=*/getNumResults());
2335   }
2336 
2337   // Print the attribute list.
2338   p.printOptionalAttrDict((*this)->getAttrs(),
2339                           /*elidedAttrs=*/getConditionAttrStrName());
2340 }
2341 
2342 IntegerSet AffineIfOp::getIntegerSet() {
2343   return (*this)
2344       ->getAttrOfType<IntegerSetAttr>(getConditionAttrStrName())
2345       .getValue();
2346 }
2347 
2348 void AffineIfOp::setIntegerSet(IntegerSet newSet) {
2349   (*this)->setAttr(getConditionAttrStrName(), IntegerSetAttr::get(newSet));
2350 }
2351 
2352 void AffineIfOp::setConditional(IntegerSet set, ValueRange operands) {
2353   setIntegerSet(set);
2354   (*this)->setOperands(operands);
2355 }
2356 
2357 void AffineIfOp::build(OpBuilder &builder, OperationState &result,
2358                        TypeRange resultTypes, IntegerSet set, ValueRange args,
2359                        bool withElseRegion) {
2360   assert(resultTypes.empty() || withElseRegion);
2361   result.addTypes(resultTypes);
2362   result.addOperands(args);
2363   result.addAttribute(getConditionAttrStrName(), IntegerSetAttr::get(set));
2364 
2365   Region *thenRegion = result.addRegion();
2366   thenRegion->push_back(new Block());
2367   if (resultTypes.empty())
2368     AffineIfOp::ensureTerminator(*thenRegion, builder, result.location);
2369 
2370   Region *elseRegion = result.addRegion();
2371   if (withElseRegion) {
2372     elseRegion->push_back(new Block());
2373     if (resultTypes.empty())
2374       AffineIfOp::ensureTerminator(*elseRegion, builder, result.location);
2375   }
2376 }
2377 
2378 void AffineIfOp::build(OpBuilder &builder, OperationState &result,
2379                        IntegerSet set, ValueRange args, bool withElseRegion) {
2380   AffineIfOp::build(builder, result, /*resultTypes=*/{}, set, args,
2381                     withElseRegion);
2382 }
2383 
2384 /// Canonicalize an affine if op's conditional (integer set + operands).
2385 LogicalResult AffineIfOp::fold(ArrayRef<Attribute>,
2386                                SmallVectorImpl<OpFoldResult> &) {
2387   auto set = getIntegerSet();
2388   SmallVector<Value, 4> operands(getOperands());
2389   canonicalizeSetAndOperands(&set, &operands);
2390 
2391   // Any canonicalization change always leads to either a reduction in the
2392   // number of operands or a change in the number of symbolic operands
2393   // (promotion of dims to symbols).
2394   if (operands.size() < getIntegerSet().getNumInputs() ||
2395       set.getNumSymbols() > getIntegerSet().getNumSymbols()) {
2396     setConditional(set, operands);
2397     return success();
2398   }
2399 
2400   return failure();
2401 }
2402 
2403 void AffineIfOp::getCanonicalizationPatterns(RewritePatternSet &results,
2404                                              MLIRContext *context) {
2405   results.add<SimplifyDeadElse, AlwaysTrueOrFalseIf>(context);
2406 }
2407 
2408 //===----------------------------------------------------------------------===//
2409 // AffineLoadOp
2410 //===----------------------------------------------------------------------===//
2411 
2412 void AffineLoadOp::build(OpBuilder &builder, OperationState &result,
2413                          AffineMap map, ValueRange operands) {
2414   assert(operands.size() == 1 + map.getNumInputs() && "inconsistent operands");
2415   result.addOperands(operands);
2416   if (map)
2417     result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
2418   auto memrefType = operands[0].getType().cast<MemRefType>();
2419   result.types.push_back(memrefType.getElementType());
2420 }
2421 
2422 void AffineLoadOp::build(OpBuilder &builder, OperationState &result,
2423                          Value memref, AffineMap map, ValueRange mapOperands) {
2424   assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
2425   result.addOperands(memref);
2426   result.addOperands(mapOperands);
2427   auto memrefType = memref.getType().cast<MemRefType>();
2428   result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
2429   result.types.push_back(memrefType.getElementType());
2430 }
2431 
2432 void AffineLoadOp::build(OpBuilder &builder, OperationState &result,
2433                          Value memref, ValueRange indices) {
2434   auto memrefType = memref.getType().cast<MemRefType>();
2435   int64_t rank = memrefType.getRank();
2436   // Create identity map for memrefs with at least one dimension or () -> ()
2437   // for zero-dimensional memrefs.
2438   auto map =
2439       rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap();
2440   build(builder, result, memref, map, indices);
2441 }
2442 
2443 ParseResult AffineLoadOp::parse(OpAsmParser &parser, OperationState &result) {
2444   auto &builder = parser.getBuilder();
2445   auto indexTy = builder.getIndexType();
2446 
2447   MemRefType type;
2448   OpAsmParser::UnresolvedOperand memrefInfo;
2449   AffineMapAttr mapAttr;
2450   SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
2451   return failure(
2452       parser.parseOperand(memrefInfo) ||
2453       parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
2454                                     AffineLoadOp::getMapAttrStrName(),
2455                                     result.attributes) ||
2456       parser.parseOptionalAttrDict(result.attributes) ||
2457       parser.parseColonType(type) ||
2458       parser.resolveOperand(memrefInfo, type, result.operands) ||
2459       parser.resolveOperands(mapOperands, indexTy, result.operands) ||
2460       parser.addTypeToList(type.getElementType(), result.types));
2461 }
2462 
2463 void AffineLoadOp::print(OpAsmPrinter &p) {
2464   p << " " << getMemRef() << '[';
2465   if (AffineMapAttr mapAttr =
2466           (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
2467     p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
2468   p << ']';
2469   p.printOptionalAttrDict((*this)->getAttrs(),
2470                           /*elidedAttrs=*/{getMapAttrStrName()});
2471   p << " : " << getMemRefType();
2472 }
2473 
2474 /// Verify common indexing invariants of affine.load, affine.store,
2475 /// affine.vector_load and affine.vector_store.
2476 static LogicalResult
2477 verifyMemoryOpIndexing(Operation *op, AffineMapAttr mapAttr,
2478                        Operation::operand_range mapOperands,
2479                        MemRefType memrefType, unsigned numIndexOperands) {
2480   if (mapAttr) {
2481     AffineMap map = mapAttr.getValue();
2482     if (map.getNumResults() != memrefType.getRank())
2483       return op->emitOpError("affine map num results must equal memref rank");
2484     if (map.getNumInputs() != numIndexOperands)
2485       return op->emitOpError("expects as many subscripts as affine map inputs");
2486   } else {
2487     if (memrefType.getRank() != numIndexOperands)
2488       return op->emitOpError(
2489           "expects the number of subscripts to be equal to memref rank");
2490   }
2491 
2492   Region *scope = getAffineScope(op);
2493   for (auto idx : mapOperands) {
2494     if (!idx.getType().isIndex())
2495       return op->emitOpError("index to load must have 'index' type");
2496     if (!isValidAffineIndexOperand(idx, scope))
2497       return op->emitOpError("index must be a dimension or symbol identifier");
2498   }
2499 
2500   return success();
2501 }
2502 
2503 LogicalResult AffineLoadOp::verify() {
2504   auto memrefType = getMemRefType();
2505   if (getType() != memrefType.getElementType())
2506     return emitOpError("result type must match element type of memref");
2507 
2508   if (failed(verifyMemoryOpIndexing(
2509           getOperation(),
2510           (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
2511           getMapOperands(), memrefType,
2512           /*numIndexOperands=*/getNumOperands() - 1)))
2513     return failure();
2514 
2515   return success();
2516 }
2517 
2518 void AffineLoadOp::getCanonicalizationPatterns(RewritePatternSet &results,
2519                                                MLIRContext *context) {
2520   results.add<SimplifyAffineOp<AffineLoadOp>>(context);
2521 }
2522 
2523 OpFoldResult AffineLoadOp::fold(ArrayRef<Attribute> cstOperands) {
2524   /// load(memrefcast) -> load
2525   if (succeeded(foldMemRefCast(*this)))
2526     return getResult();
2527 
2528   // Fold load from a global constant memref.
2529   auto getGlobalOp = getMemref().getDefiningOp<memref::GetGlobalOp>();
2530   if (!getGlobalOp)
2531     return {};
2532   // Get to the memref.global defining the symbol.
2533   auto *symbolTableOp = getGlobalOp->getParentWithTrait<OpTrait::SymbolTable>();
2534   if (!symbolTableOp)
2535     return {};
2536   auto global = dyn_cast_or_null<memref::GlobalOp>(
2537       SymbolTable::lookupSymbolIn(symbolTableOp, getGlobalOp.nameAttr()));
2538   if (!global)
2539     return {};
2540 
2541   // Check if the global memref is a constant.
2542   auto cstAttr =
2543       global.getConstantInitValue().dyn_cast_or_null<DenseElementsAttr>();
2544   if (!cstAttr)
2545     return {};
2546   // If it's a splat constant, we can fold irrespective of indices.
2547   if (auto splatAttr = cstAttr.dyn_cast<SplatElementsAttr>())
2548     return splatAttr.getSplatValue<Attribute>();
2549   // Otherwise, we can fold only if we know the indices.
2550   if (!getAffineMap().isConstant())
2551     return {};
2552   auto indices = llvm::to_vector<4>(
2553       llvm::map_range(getAffineMap().getConstantResults(),
2554                       [](int64_t v) -> uint64_t { return v; }));
2555   return cstAttr.getValues<Attribute>()[indices];
2556 }
2557 
2558 //===----------------------------------------------------------------------===//
2559 // AffineStoreOp
2560 //===----------------------------------------------------------------------===//
2561 
2562 void AffineStoreOp::build(OpBuilder &builder, OperationState &result,
2563                           Value valueToStore, Value memref, AffineMap map,
2564                           ValueRange mapOperands) {
2565   assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
2566   result.addOperands(valueToStore);
2567   result.addOperands(memref);
2568   result.addOperands(mapOperands);
2569   result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
2570 }
2571 
2572 // Use identity map.
2573 void AffineStoreOp::build(OpBuilder &builder, OperationState &result,
2574                           Value valueToStore, Value memref,
2575                           ValueRange indices) {
2576   auto memrefType = memref.getType().cast<MemRefType>();
2577   int64_t rank = memrefType.getRank();
2578   // Create identity map for memrefs with at least one dimension or () -> ()
2579   // for zero-dimensional memrefs.
2580   auto map =
2581       rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap();
2582   build(builder, result, valueToStore, memref, map, indices);
2583 }
2584 
2585 ParseResult AffineStoreOp::parse(OpAsmParser &parser, OperationState &result) {
2586   auto indexTy = parser.getBuilder().getIndexType();
2587 
2588   MemRefType type;
2589   OpAsmParser::UnresolvedOperand storeValueInfo;
2590   OpAsmParser::UnresolvedOperand memrefInfo;
2591   AffineMapAttr mapAttr;
2592   SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
2593   return failure(parser.parseOperand(storeValueInfo) || parser.parseComma() ||
2594                  parser.parseOperand(memrefInfo) ||
2595                  parser.parseAffineMapOfSSAIds(
2596                      mapOperands, mapAttr, AffineStoreOp::getMapAttrStrName(),
2597                      result.attributes) ||
2598                  parser.parseOptionalAttrDict(result.attributes) ||
2599                  parser.parseColonType(type) ||
2600                  parser.resolveOperand(storeValueInfo, type.getElementType(),
2601                                        result.operands) ||
2602                  parser.resolveOperand(memrefInfo, type, result.operands) ||
2603                  parser.resolveOperands(mapOperands, indexTy, result.operands));
2604 }
2605 
2606 void AffineStoreOp::print(OpAsmPrinter &p) {
2607   p << " " << getValueToStore();
2608   p << ", " << getMemRef() << '[';
2609   if (AffineMapAttr mapAttr =
2610           (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
2611     p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
2612   p << ']';
2613   p.printOptionalAttrDict((*this)->getAttrs(),
2614                           /*elidedAttrs=*/{getMapAttrStrName()});
2615   p << " : " << getMemRefType();
2616 }
2617 
2618 LogicalResult AffineStoreOp::verify() {
2619   // The value to store must have the same type as memref element type.
2620   auto memrefType = getMemRefType();
2621   if (getValueToStore().getType() != memrefType.getElementType())
2622     return emitOpError(
2623         "value to store must have the same type as memref element type");
2624 
2625   if (failed(verifyMemoryOpIndexing(
2626           getOperation(),
2627           (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
2628           getMapOperands(), memrefType,
2629           /*numIndexOperands=*/getNumOperands() - 2)))
2630     return failure();
2631 
2632   return success();
2633 }
2634 
2635 void AffineStoreOp::getCanonicalizationPatterns(RewritePatternSet &results,
2636                                                 MLIRContext *context) {
2637   results.add<SimplifyAffineOp<AffineStoreOp>>(context);
2638 }
2639 
2640 LogicalResult AffineStoreOp::fold(ArrayRef<Attribute> cstOperands,
2641                                   SmallVectorImpl<OpFoldResult> &results) {
2642   /// store(memrefcast) -> store
2643   return foldMemRefCast(*this, getValueToStore());
2644 }
2645 
2646 //===----------------------------------------------------------------------===//
2647 // AffineMinMaxOpBase
2648 //===----------------------------------------------------------------------===//
2649 
2650 template <typename T> static LogicalResult verifyAffineMinMaxOp(T op) {
2651   // Verify that operand count matches affine map dimension and symbol count.
2652   if (op.getNumOperands() !=
2653       op.getMap().getNumDims() + op.getMap().getNumSymbols())
2654     return op.emitOpError(
2655         "operand count and affine map dimension and symbol count must match");
2656   return success();
2657 }
2658 
2659 template <typename T> static void printAffineMinMaxOp(OpAsmPrinter &p, T op) {
2660   p << ' ' << op->getAttr(T::getMapAttrStrName());
2661   auto operands = op.getOperands();
2662   unsigned numDims = op.getMap().getNumDims();
2663   p << '(' << operands.take_front(numDims) << ')';
2664 
2665   if (operands.size() != numDims)
2666     p << '[' << operands.drop_front(numDims) << ']';
2667   p.printOptionalAttrDict(op->getAttrs(),
2668                           /*elidedAttrs=*/{T::getMapAttrStrName()});
2669 }
2670 
2671 template <typename T>
2672 static ParseResult parseAffineMinMaxOp(OpAsmParser &parser,
2673                                        OperationState &result) {
2674   auto &builder = parser.getBuilder();
2675   auto indexType = builder.getIndexType();
2676   SmallVector<OpAsmParser::UnresolvedOperand, 8> dimInfos;
2677   SmallVector<OpAsmParser::UnresolvedOperand, 8> symInfos;
2678   AffineMapAttr mapAttr;
2679   return failure(
2680       parser.parseAttribute(mapAttr, T::getMapAttrStrName(),
2681                             result.attributes) ||
2682       parser.parseOperandList(dimInfos, OpAsmParser::Delimiter::Paren) ||
2683       parser.parseOperandList(symInfos,
2684                               OpAsmParser::Delimiter::OptionalSquare) ||
2685       parser.parseOptionalAttrDict(result.attributes) ||
2686       parser.resolveOperands(dimInfos, indexType, result.operands) ||
2687       parser.resolveOperands(symInfos, indexType, result.operands) ||
2688       parser.addTypeToList(indexType, result.types));
2689 }
2690 
2691 /// Fold an affine min or max operation with the given operands. The operand
2692 /// list may contain nulls, which are interpreted as the operand not being a
2693 /// constant.
2694 template <typename T>
2695 static OpFoldResult foldMinMaxOp(T op, ArrayRef<Attribute> operands) {
2696   static_assert(llvm::is_one_of<T, AffineMinOp, AffineMaxOp>::value,
2697                 "expected affine min or max op");
2698 
2699   // Fold the affine map.
2700   // TODO: Fold more cases:
2701   // min(some_affine, some_affine + constant, ...), etc.
2702   SmallVector<int64_t, 2> results;
2703   auto foldedMap = op.getMap().partialConstantFold(operands, &results);
2704 
2705   // If some of the map results are not constant, try changing the map in-place.
2706   if (results.empty()) {
2707     // If the map is the same, report that folding did not happen.
2708     if (foldedMap == op.getMap())
2709       return {};
2710     op->setAttr("map", AffineMapAttr::get(foldedMap));
2711     return op.getResult();
2712   }
2713 
2714   // Otherwise, completely fold the op into a constant.
2715   auto resultIt = std::is_same<T, AffineMinOp>::value
2716                       ? std::min_element(results.begin(), results.end())
2717                       : std::max_element(results.begin(), results.end());
2718   if (resultIt == results.end())
2719     return {};
2720   return IntegerAttr::get(IndexType::get(op.getContext()), *resultIt);
2721 }
2722 
2723 /// Remove duplicated expressions in affine min/max ops.
2724 template <typename T>
2725 struct DeduplicateAffineMinMaxExpressions : public OpRewritePattern<T> {
2726   using OpRewritePattern<T>::OpRewritePattern;
2727 
2728   LogicalResult matchAndRewrite(T affineOp,
2729                                 PatternRewriter &rewriter) const override {
2730     AffineMap oldMap = affineOp.getAffineMap();
2731 
2732     SmallVector<AffineExpr, 4> newExprs;
2733     for (AffineExpr expr : oldMap.getResults()) {
2734       // This is a linear scan over newExprs, but it should be fine given that
2735       // we typically just have a few expressions per op.
2736       if (!llvm::is_contained(newExprs, expr))
2737         newExprs.push_back(expr);
2738     }
2739 
2740     if (newExprs.size() == oldMap.getNumResults())
2741       return failure();
2742 
2743     auto newMap = AffineMap::get(oldMap.getNumDims(), oldMap.getNumSymbols(),
2744                                  newExprs, rewriter.getContext());
2745     rewriter.replaceOpWithNewOp<T>(affineOp, newMap, affineOp.getMapOperands());
2746 
2747     return success();
2748   }
2749 };
2750 
2751 /// Merge an affine min/max op to its consumers if its consumer is also an
2752 /// affine min/max op.
2753 ///
2754 /// This pattern requires the producer affine min/max op is bound to a
2755 /// dimension/symbol that is used as a standalone expression in the consumer
2756 /// affine op's map.
2757 ///
2758 /// For example, a pattern like the following:
2759 ///
2760 ///   %0 = affine.min affine_map<()[s0] -> (s0 + 16, s0 * 8)> ()[%sym1]
2761 ///   %1 = affine.min affine_map<(d0)[s0] -> (s0 + 4, d0)> (%0)[%sym2]
2762 ///
2763 /// Can be turned into:
2764 ///
2765 ///   %1 = affine.min affine_map<
2766 ///          ()[s0, s1] -> (s0 + 4, s1 + 16, s1 * 8)> ()[%sym2, %sym1]
2767 template <typename T> struct MergeAffineMinMaxOp : public OpRewritePattern<T> {
2768   using OpRewritePattern<T>::OpRewritePattern;
2769 
2770   LogicalResult matchAndRewrite(T affineOp,
2771                                 PatternRewriter &rewriter) const override {
2772     AffineMap oldMap = affineOp.getAffineMap();
2773     ValueRange dimOperands =
2774         affineOp.getMapOperands().take_front(oldMap.getNumDims());
2775     ValueRange symOperands =
2776         affineOp.getMapOperands().take_back(oldMap.getNumSymbols());
2777 
2778     auto newDimOperands = llvm::to_vector<8>(dimOperands);
2779     auto newSymOperands = llvm::to_vector<8>(symOperands);
2780     SmallVector<AffineExpr, 4> newExprs;
2781     SmallVector<T, 4> producerOps;
2782 
2783     // Go over each expression to see whether it's a single dimension/symbol
2784     // with the corresponding operand which is the result of another affine
2785     // min/max op. If So it can be merged into this affine op.
2786     for (AffineExpr expr : oldMap.getResults()) {
2787       if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>()) {
2788         Value symValue = symOperands[symExpr.getPosition()];
2789         if (auto producerOp = symValue.getDefiningOp<T>()) {
2790           producerOps.push_back(producerOp);
2791           continue;
2792         }
2793       } else if (auto dimExpr = expr.dyn_cast<AffineDimExpr>()) {
2794         Value dimValue = dimOperands[dimExpr.getPosition()];
2795         if (auto producerOp = dimValue.getDefiningOp<T>()) {
2796           producerOps.push_back(producerOp);
2797           continue;
2798         }
2799       }
2800       // For the above cases we will remove the expression by merging the
2801       // producer affine min/max's affine expressions. Otherwise we need to
2802       // keep the existing expression.
2803       newExprs.push_back(expr);
2804     }
2805 
2806     if (producerOps.empty())
2807       return failure();
2808 
2809     unsigned numUsedDims = oldMap.getNumDims();
2810     unsigned numUsedSyms = oldMap.getNumSymbols();
2811 
2812     // Now go over all producer affine ops and merge their expressions.
2813     for (T producerOp : producerOps) {
2814       AffineMap producerMap = producerOp.getAffineMap();
2815       unsigned numProducerDims = producerMap.getNumDims();
2816       unsigned numProducerSyms = producerMap.getNumSymbols();
2817 
2818       // Collect all dimension/symbol values.
2819       ValueRange dimValues =
2820           producerOp.getMapOperands().take_front(numProducerDims);
2821       ValueRange symValues =
2822           producerOp.getMapOperands().take_back(numProducerSyms);
2823       newDimOperands.append(dimValues.begin(), dimValues.end());
2824       newSymOperands.append(symValues.begin(), symValues.end());
2825 
2826       // For expressions we need to shift to avoid overlap.
2827       for (AffineExpr expr : producerMap.getResults()) {
2828         newExprs.push_back(expr.shiftDims(numProducerDims, numUsedDims)
2829                                .shiftSymbols(numProducerSyms, numUsedSyms));
2830       }
2831 
2832       numUsedDims += numProducerDims;
2833       numUsedSyms += numProducerSyms;
2834     }
2835 
2836     auto newMap = AffineMap::get(numUsedDims, numUsedSyms, newExprs,
2837                                  rewriter.getContext());
2838     auto newOperands =
2839         llvm::to_vector<8>(llvm::concat<Value>(newDimOperands, newSymOperands));
2840     rewriter.replaceOpWithNewOp<T>(affineOp, newMap, newOperands);
2841 
2842     return success();
2843   }
2844 };
2845 
2846 /// Canonicalize the result expression order of an affine map and return success
2847 /// if the order changed.
2848 ///
2849 /// The function flattens the map's affine expressions to coefficient arrays and
2850 /// sorts them in lexicographic order. A coefficient array contains a multiplier
2851 /// for every dimension/symbol and a constant term. The canonicalization fails
2852 /// if a result expression is not pure or if the flattening requires local
2853 /// variables that, unlike dimensions and symbols, have no global order.
2854 static LogicalResult canonicalizeMapExprAndTermOrder(AffineMap &map) {
2855   SmallVector<SmallVector<int64_t>> flattenedExprs;
2856   for (const AffineExpr &resultExpr : map.getResults()) {
2857     // Fail if the expression is not pure.
2858     if (!resultExpr.isPureAffine())
2859       return failure();
2860 
2861     SimpleAffineExprFlattener flattener(map.getNumDims(), map.getNumSymbols());
2862     flattener.walkPostOrder(resultExpr);
2863 
2864     // Fail if the flattened expression has local variables.
2865     if (flattener.operandExprStack.back().size() !=
2866         map.getNumDims() + map.getNumSymbols() + 1)
2867       return failure();
2868 
2869     flattenedExprs.emplace_back(flattener.operandExprStack.back().begin(),
2870                                 flattener.operandExprStack.back().end());
2871   }
2872 
2873   // Fail if sorting is not necessary.
2874   if (llvm::is_sorted(flattenedExprs))
2875     return failure();
2876 
2877   // Reorder the result expressions according to their flattened form.
2878   SmallVector<unsigned> resultPermutation =
2879       llvm::to_vector(llvm::seq<unsigned>(0, map.getNumResults()));
2880   llvm::sort(resultPermutation, [&](unsigned lhs, unsigned rhs) {
2881     return flattenedExprs[lhs] < flattenedExprs[rhs];
2882   });
2883   SmallVector<AffineExpr> newExprs;
2884   for (unsigned idx : resultPermutation)
2885     newExprs.push_back(map.getResult(idx));
2886 
2887   map = AffineMap::get(map.getNumDims(), map.getNumSymbols(), newExprs,
2888                        map.getContext());
2889   return success();
2890 }
2891 
2892 /// Canonicalize the affine map result expression order of an affine min/max
2893 /// operation.
2894 ///
2895 /// The pattern calls `canonicalizeMapExprAndTermOrder` to order the result
2896 /// expressions and replaces the operation if the order changed.
2897 ///
2898 /// For example, the following operation:
2899 ///
2900 ///   %0 = affine.min affine_map<(d0, d1) -> (d0 + d1, d1 + 16, 32)> (%i0, %i1)
2901 ///
2902 /// Turns into:
2903 ///
2904 ///   %0 = affine.min affine_map<(d0, d1) -> (32, d1 + 16, d0 + d1)> (%i0, %i1)
2905 template <typename T>
2906 struct CanonicalizeAffineMinMaxOpExprAndTermOrder : public OpRewritePattern<T> {
2907   using OpRewritePattern<T>::OpRewritePattern;
2908 
2909   LogicalResult matchAndRewrite(T affineOp,
2910                                 PatternRewriter &rewriter) const override {
2911     AffineMap map = affineOp.getAffineMap();
2912     if (failed(canonicalizeMapExprAndTermOrder(map)))
2913       return failure();
2914 
2915     rewriter.replaceOpWithNewOp<T>(affineOp, map, affineOp.getMapOperands());
2916     return success();
2917   }
2918 };
2919 
2920 template <typename T>
2921 struct CanonicalizeSingleResultAffineMinMaxOp : public OpRewritePattern<T> {
2922   using OpRewritePattern<T>::OpRewritePattern;
2923 
2924   LogicalResult matchAndRewrite(T affineOp,
2925                                 PatternRewriter &rewriter) const override {
2926     if (affineOp.getMap().getNumResults() != 1)
2927       return failure();
2928     rewriter.replaceOpWithNewOp<AffineApplyOp>(affineOp, affineOp.getMap(),
2929                                                affineOp.getOperands());
2930     return success();
2931   }
2932 };
2933 
2934 //===----------------------------------------------------------------------===//
2935 // AffineMinOp
2936 //===----------------------------------------------------------------------===//
2937 //
2938 //   %0 = affine.min (d0) -> (1000, d0 + 512) (%i0)
2939 //
2940 
2941 OpFoldResult AffineMinOp::fold(ArrayRef<Attribute> operands) {
2942   return foldMinMaxOp(*this, operands);
2943 }
2944 
2945 void AffineMinOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2946                                               MLIRContext *context) {
2947   patterns.add<CanonicalizeSingleResultAffineMinMaxOp<AffineMinOp>,
2948                DeduplicateAffineMinMaxExpressions<AffineMinOp>,
2949                MergeAffineMinMaxOp<AffineMinOp>, SimplifyAffineOp<AffineMinOp>,
2950                CanonicalizeAffineMinMaxOpExprAndTermOrder<AffineMinOp>>(
2951       context);
2952 }
2953 
2954 LogicalResult AffineMinOp::verify() { return verifyAffineMinMaxOp(*this); }
2955 
2956 ParseResult AffineMinOp::parse(OpAsmParser &parser, OperationState &result) {
2957   return parseAffineMinMaxOp<AffineMinOp>(parser, result);
2958 }
2959 
2960 void AffineMinOp::print(OpAsmPrinter &p) { printAffineMinMaxOp(p, *this); }
2961 
2962 //===----------------------------------------------------------------------===//
2963 // AffineMaxOp
2964 //===----------------------------------------------------------------------===//
2965 //
2966 //   %0 = affine.max (d0) -> (1000, d0 + 512) (%i0)
2967 //
2968 
2969 OpFoldResult AffineMaxOp::fold(ArrayRef<Attribute> operands) {
2970   return foldMinMaxOp(*this, operands);
2971 }
2972 
2973 void AffineMaxOp::getCanonicalizationPatterns(RewritePatternSet &patterns,
2974                                               MLIRContext *context) {
2975   patterns.add<CanonicalizeSingleResultAffineMinMaxOp<AffineMaxOp>,
2976                DeduplicateAffineMinMaxExpressions<AffineMaxOp>,
2977                MergeAffineMinMaxOp<AffineMaxOp>, SimplifyAffineOp<AffineMaxOp>,
2978                CanonicalizeAffineMinMaxOpExprAndTermOrder<AffineMaxOp>>(
2979       context);
2980 }
2981 
2982 LogicalResult AffineMaxOp::verify() { return verifyAffineMinMaxOp(*this); }
2983 
2984 ParseResult AffineMaxOp::parse(OpAsmParser &parser, OperationState &result) {
2985   return parseAffineMinMaxOp<AffineMaxOp>(parser, result);
2986 }
2987 
2988 void AffineMaxOp::print(OpAsmPrinter &p) { printAffineMinMaxOp(p, *this); }
2989 
2990 //===----------------------------------------------------------------------===//
2991 // AffinePrefetchOp
2992 //===----------------------------------------------------------------------===//
2993 
2994 //
2995 // affine.prefetch %0[%i, %j + 5], read, locality<3>, data : memref<400x400xi32>
2996 //
2997 ParseResult AffinePrefetchOp::parse(OpAsmParser &parser,
2998                                     OperationState &result) {
2999   auto &builder = parser.getBuilder();
3000   auto indexTy = builder.getIndexType();
3001 
3002   MemRefType type;
3003   OpAsmParser::UnresolvedOperand memrefInfo;
3004   IntegerAttr hintInfo;
3005   auto i32Type = parser.getBuilder().getIntegerType(32);
3006   StringRef readOrWrite, cacheType;
3007 
3008   AffineMapAttr mapAttr;
3009   SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
3010   if (parser.parseOperand(memrefInfo) ||
3011       parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
3012                                     AffinePrefetchOp::getMapAttrStrName(),
3013                                     result.attributes) ||
3014       parser.parseComma() || parser.parseKeyword(&readOrWrite) ||
3015       parser.parseComma() || parser.parseKeyword("locality") ||
3016       parser.parseLess() ||
3017       parser.parseAttribute(hintInfo, i32Type,
3018                             AffinePrefetchOp::getLocalityHintAttrStrName(),
3019                             result.attributes) ||
3020       parser.parseGreater() || parser.parseComma() ||
3021       parser.parseKeyword(&cacheType) ||
3022       parser.parseOptionalAttrDict(result.attributes) ||
3023       parser.parseColonType(type) ||
3024       parser.resolveOperand(memrefInfo, type, result.operands) ||
3025       parser.resolveOperands(mapOperands, indexTy, result.operands))
3026     return failure();
3027 
3028   if (!readOrWrite.equals("read") && !readOrWrite.equals("write"))
3029     return parser.emitError(parser.getNameLoc(),
3030                             "rw specifier has to be 'read' or 'write'");
3031   result.addAttribute(
3032       AffinePrefetchOp::getIsWriteAttrStrName(),
3033       parser.getBuilder().getBoolAttr(readOrWrite.equals("write")));
3034 
3035   if (!cacheType.equals("data") && !cacheType.equals("instr"))
3036     return parser.emitError(parser.getNameLoc(),
3037                             "cache type has to be 'data' or 'instr'");
3038 
3039   result.addAttribute(
3040       AffinePrefetchOp::getIsDataCacheAttrStrName(),
3041       parser.getBuilder().getBoolAttr(cacheType.equals("data")));
3042 
3043   return success();
3044 }
3045 
3046 void AffinePrefetchOp::print(OpAsmPrinter &p) {
3047   p << " " << getMemref() << '[';
3048   AffineMapAttr mapAttr =
3049       (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
3050   if (mapAttr)
3051     p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
3052   p << ']' << ", " << (getIsWrite() ? "write" : "read") << ", "
3053     << "locality<" << getLocalityHint() << ">, "
3054     << (getIsDataCache() ? "data" : "instr");
3055   p.printOptionalAttrDict(
3056       (*this)->getAttrs(),
3057       /*elidedAttrs=*/{getMapAttrStrName(), getLocalityHintAttrStrName(),
3058                        getIsDataCacheAttrStrName(), getIsWriteAttrStrName()});
3059   p << " : " << getMemRefType();
3060 }
3061 
3062 LogicalResult AffinePrefetchOp::verify() {
3063   auto mapAttr = (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName());
3064   if (mapAttr) {
3065     AffineMap map = mapAttr.getValue();
3066     if (map.getNumResults() != getMemRefType().getRank())
3067       return emitOpError("affine.prefetch affine map num results must equal"
3068                          " memref rank");
3069     if (map.getNumInputs() + 1 != getNumOperands())
3070       return emitOpError("too few operands");
3071   } else {
3072     if (getNumOperands() != 1)
3073       return emitOpError("too few operands");
3074   }
3075 
3076   Region *scope = getAffineScope(*this);
3077   for (auto idx : getMapOperands()) {
3078     if (!isValidAffineIndexOperand(idx, scope))
3079       return emitOpError("index must be a dimension or symbol identifier");
3080   }
3081   return success();
3082 }
3083 
3084 void AffinePrefetchOp::getCanonicalizationPatterns(RewritePatternSet &results,
3085                                                    MLIRContext *context) {
3086   // prefetch(memrefcast) -> prefetch
3087   results.add<SimplifyAffineOp<AffinePrefetchOp>>(context);
3088 }
3089 
3090 LogicalResult AffinePrefetchOp::fold(ArrayRef<Attribute> cstOperands,
3091                                      SmallVectorImpl<OpFoldResult> &results) {
3092   /// prefetch(memrefcast) -> prefetch
3093   return foldMemRefCast(*this);
3094 }
3095 
3096 //===----------------------------------------------------------------------===//
3097 // AffineParallelOp
3098 //===----------------------------------------------------------------------===//
3099 
3100 void AffineParallelOp::build(OpBuilder &builder, OperationState &result,
3101                              TypeRange resultTypes,
3102                              ArrayRef<arith::AtomicRMWKind> reductions,
3103                              ArrayRef<int64_t> ranges) {
3104   SmallVector<AffineMap> lbs(ranges.size(), builder.getConstantAffineMap(0));
3105   auto ubs = llvm::to_vector<4>(llvm::map_range(ranges, [&](int64_t value) {
3106     return builder.getConstantAffineMap(value);
3107   }));
3108   SmallVector<int64_t> steps(ranges.size(), 1);
3109   build(builder, result, resultTypes, reductions, lbs, /*lbArgs=*/{}, ubs,
3110         /*ubArgs=*/{}, steps);
3111 }
3112 
3113 void AffineParallelOp::build(OpBuilder &builder, OperationState &result,
3114                              TypeRange resultTypes,
3115                              ArrayRef<arith::AtomicRMWKind> reductions,
3116                              ArrayRef<AffineMap> lbMaps, ValueRange lbArgs,
3117                              ArrayRef<AffineMap> ubMaps, ValueRange ubArgs,
3118                              ArrayRef<int64_t> steps) {
3119   assert(llvm::all_of(lbMaps,
3120                       [lbMaps](AffineMap m) {
3121                         return m.getNumDims() == lbMaps[0].getNumDims() &&
3122                                m.getNumSymbols() == lbMaps[0].getNumSymbols();
3123                       }) &&
3124          "expected all lower bounds maps to have the same number of dimensions "
3125          "and symbols");
3126   assert(llvm::all_of(ubMaps,
3127                       [ubMaps](AffineMap m) {
3128                         return m.getNumDims() == ubMaps[0].getNumDims() &&
3129                                m.getNumSymbols() == ubMaps[0].getNumSymbols();
3130                       }) &&
3131          "expected all upper bounds maps to have the same number of dimensions "
3132          "and symbols");
3133   assert((lbMaps.empty() || lbMaps[0].getNumInputs() == lbArgs.size()) &&
3134          "expected lower bound maps to have as many inputs as lower bound "
3135          "operands");
3136   assert((ubMaps.empty() || ubMaps[0].getNumInputs() == ubArgs.size()) &&
3137          "expected upper bound maps to have as many inputs as upper bound "
3138          "operands");
3139 
3140   result.addTypes(resultTypes);
3141 
3142   // Convert the reductions to integer attributes.
3143   SmallVector<Attribute, 4> reductionAttrs;
3144   for (arith::AtomicRMWKind reduction : reductions)
3145     reductionAttrs.push_back(
3146         builder.getI64IntegerAttr(static_cast<int64_t>(reduction)));
3147   result.addAttribute(getReductionsAttrStrName(),
3148                       builder.getArrayAttr(reductionAttrs));
3149 
3150   // Concatenates maps defined in the same input space (same dimensions and
3151   // symbols), assumes there is at least one map.
3152   auto concatMapsSameInput = [&builder](ArrayRef<AffineMap> maps,
3153                                         SmallVectorImpl<int32_t> &groups) {
3154     if (maps.empty())
3155       return AffineMap::get(builder.getContext());
3156     SmallVector<AffineExpr> exprs;
3157     groups.reserve(groups.size() + maps.size());
3158     exprs.reserve(maps.size());
3159     for (AffineMap m : maps) {
3160       llvm::append_range(exprs, m.getResults());
3161       groups.push_back(m.getNumResults());
3162     }
3163     return AffineMap::get(maps[0].getNumDims(), maps[0].getNumSymbols(), exprs,
3164                           maps[0].getContext());
3165   };
3166 
3167   // Set up the bounds.
3168   SmallVector<int32_t> lbGroups, ubGroups;
3169   AffineMap lbMap = concatMapsSameInput(lbMaps, lbGroups);
3170   AffineMap ubMap = concatMapsSameInput(ubMaps, ubGroups);
3171   result.addAttribute(getLowerBoundsMapAttrStrName(),
3172                       AffineMapAttr::get(lbMap));
3173   result.addAttribute(getLowerBoundsGroupsAttrStrName(),
3174                       builder.getI32TensorAttr(lbGroups));
3175   result.addAttribute(getUpperBoundsMapAttrStrName(),
3176                       AffineMapAttr::get(ubMap));
3177   result.addAttribute(getUpperBoundsGroupsAttrStrName(),
3178                       builder.getI32TensorAttr(ubGroups));
3179   result.addAttribute(getStepsAttrStrName(), builder.getI64ArrayAttr(steps));
3180   result.addOperands(lbArgs);
3181   result.addOperands(ubArgs);
3182 
3183   // Create a region and a block for the body.
3184   auto *bodyRegion = result.addRegion();
3185   auto *body = new Block();
3186   // Add all the block arguments.
3187   for (unsigned i = 0, e = steps.size(); i < e; ++i)
3188     body->addArgument(IndexType::get(builder.getContext()), result.location);
3189   bodyRegion->push_back(body);
3190   if (resultTypes.empty())
3191     ensureTerminator(*bodyRegion, builder, result.location);
3192 }
3193 
3194 Region &AffineParallelOp::getLoopBody() { return getRegion(); }
3195 
3196 unsigned AffineParallelOp::getNumDims() { return getSteps().size(); }
3197 
3198 AffineParallelOp::operand_range AffineParallelOp::getLowerBoundsOperands() {
3199   return getOperands().take_front(getLowerBoundsMap().getNumInputs());
3200 }
3201 
3202 AffineParallelOp::operand_range AffineParallelOp::getUpperBoundsOperands() {
3203   return getOperands().drop_front(getLowerBoundsMap().getNumInputs());
3204 }
3205 
3206 AffineMap AffineParallelOp::getLowerBoundMap(unsigned pos) {
3207   auto values = getLowerBoundsGroups().getValues<int32_t>();
3208   unsigned start = 0;
3209   for (unsigned i = 0; i < pos; ++i)
3210     start += values[i];
3211   return getLowerBoundsMap().getSliceMap(start, values[pos]);
3212 }
3213 
3214 AffineMap AffineParallelOp::getUpperBoundMap(unsigned pos) {
3215   auto values = getUpperBoundsGroups().getValues<int32_t>();
3216   unsigned start = 0;
3217   for (unsigned i = 0; i < pos; ++i)
3218     start += values[i];
3219   return getUpperBoundsMap().getSliceMap(start, values[pos]);
3220 }
3221 
3222 AffineValueMap AffineParallelOp::getLowerBoundsValueMap() {
3223   return AffineValueMap(getLowerBoundsMap(), getLowerBoundsOperands());
3224 }
3225 
3226 AffineValueMap AffineParallelOp::getUpperBoundsValueMap() {
3227   return AffineValueMap(getUpperBoundsMap(), getUpperBoundsOperands());
3228 }
3229 
3230 Optional<SmallVector<int64_t, 8>> AffineParallelOp::getConstantRanges() {
3231   if (hasMinMaxBounds())
3232     return llvm::None;
3233 
3234   // Try to convert all the ranges to constant expressions.
3235   SmallVector<int64_t, 8> out;
3236   AffineValueMap rangesValueMap;
3237   AffineValueMap::difference(getUpperBoundsValueMap(), getLowerBoundsValueMap(),
3238                              &rangesValueMap);
3239   out.reserve(rangesValueMap.getNumResults());
3240   for (unsigned i = 0, e = rangesValueMap.getNumResults(); i < e; ++i) {
3241     auto expr = rangesValueMap.getResult(i);
3242     auto cst = expr.dyn_cast<AffineConstantExpr>();
3243     if (!cst)
3244       return llvm::None;
3245     out.push_back(cst.getValue());
3246   }
3247   return out;
3248 }
3249 
3250 Block *AffineParallelOp::getBody() { return &getRegion().front(); }
3251 
3252 OpBuilder AffineParallelOp::getBodyBuilder() {
3253   return OpBuilder(getBody(), std::prev(getBody()->end()));
3254 }
3255 
3256 void AffineParallelOp::setLowerBounds(ValueRange lbOperands, AffineMap map) {
3257   assert(lbOperands.size() == map.getNumInputs() &&
3258          "operands to map must match number of inputs");
3259 
3260   auto ubOperands = getUpperBoundsOperands();
3261 
3262   SmallVector<Value, 4> newOperands(lbOperands);
3263   newOperands.append(ubOperands.begin(), ubOperands.end());
3264   (*this)->setOperands(newOperands);
3265 
3266   setLowerBoundsMapAttr(AffineMapAttr::get(map));
3267 }
3268 
3269 void AffineParallelOp::setUpperBounds(ValueRange ubOperands, AffineMap map) {
3270   assert(ubOperands.size() == map.getNumInputs() &&
3271          "operands to map must match number of inputs");
3272 
3273   SmallVector<Value, 4> newOperands(getLowerBoundsOperands());
3274   newOperands.append(ubOperands.begin(), ubOperands.end());
3275   (*this)->setOperands(newOperands);
3276 
3277   setUpperBoundsMapAttr(AffineMapAttr::get(map));
3278 }
3279 
3280 void AffineParallelOp::setLowerBoundsMap(AffineMap map) {
3281   AffineMap lbMap = getLowerBoundsMap();
3282   assert(lbMap.getNumDims() == map.getNumDims() &&
3283          lbMap.getNumSymbols() == map.getNumSymbols());
3284   (void)lbMap;
3285   setLowerBoundsMapAttr(AffineMapAttr::get(map));
3286 }
3287 
3288 void AffineParallelOp::setUpperBoundsMap(AffineMap map) {
3289   AffineMap ubMap = getUpperBoundsMap();
3290   assert(ubMap.getNumDims() == map.getNumDims() &&
3291          ubMap.getNumSymbols() == map.getNumSymbols());
3292   (void)ubMap;
3293   setUpperBoundsMapAttr(AffineMapAttr::get(map));
3294 }
3295 
3296 void AffineParallelOp::setSteps(ArrayRef<int64_t> newSteps) {
3297   setStepsAttr(getBodyBuilder().getI64ArrayAttr(newSteps));
3298 }
3299 
3300 LogicalResult AffineParallelOp::verify() {
3301   auto numDims = getNumDims();
3302   if (getLowerBoundsGroups().getNumElements() != numDims ||
3303       getUpperBoundsGroups().getNumElements() != numDims ||
3304       getSteps().size() != numDims || getBody()->getNumArguments() != numDims) {
3305     return emitOpError() << "the number of region arguments ("
3306                          << getBody()->getNumArguments()
3307                          << ") and the number of map groups for lower ("
3308                          << getLowerBoundsGroups().getNumElements()
3309                          << ") and upper bound ("
3310                          << getUpperBoundsGroups().getNumElements()
3311                          << "), and the number of steps (" << getSteps().size()
3312                          << ") must all match";
3313   }
3314 
3315   unsigned expectedNumLBResults = 0;
3316   for (APInt v : getLowerBoundsGroups())
3317     expectedNumLBResults += v.getZExtValue();
3318   if (expectedNumLBResults != getLowerBoundsMap().getNumResults())
3319     return emitOpError() << "expected lower bounds map to have "
3320                          << expectedNumLBResults << " results";
3321   unsigned expectedNumUBResults = 0;
3322   for (APInt v : getUpperBoundsGroups())
3323     expectedNumUBResults += v.getZExtValue();
3324   if (expectedNumUBResults != getUpperBoundsMap().getNumResults())
3325     return emitOpError() << "expected upper bounds map to have "
3326                          << expectedNumUBResults << " results";
3327 
3328   if (getReductions().size() != getNumResults())
3329     return emitOpError("a reduction must be specified for each output");
3330 
3331   // Verify reduction  ops are all valid
3332   for (Attribute attr : getReductions()) {
3333     auto intAttr = attr.dyn_cast<IntegerAttr>();
3334     if (!intAttr || !arith::symbolizeAtomicRMWKind(intAttr.getInt()))
3335       return emitOpError("invalid reduction attribute");
3336   }
3337 
3338   // Verify that the bound operands are valid dimension/symbols.
3339   /// Lower bounds.
3340   if (failed(verifyDimAndSymbolIdentifiers(*this, getLowerBoundsOperands(),
3341                                            getLowerBoundsMap().getNumDims())))
3342     return failure();
3343   /// Upper bounds.
3344   if (failed(verifyDimAndSymbolIdentifiers(*this, getUpperBoundsOperands(),
3345                                            getUpperBoundsMap().getNumDims())))
3346     return failure();
3347   return success();
3348 }
3349 
3350 LogicalResult AffineValueMap::canonicalize() {
3351   SmallVector<Value, 4> newOperands{operands};
3352   auto newMap = getAffineMap();
3353   composeAffineMapAndOperands(&newMap, &newOperands);
3354   if (newMap == getAffineMap() && newOperands == operands)
3355     return failure();
3356   reset(newMap, newOperands);
3357   return success();
3358 }
3359 
3360 /// Canonicalize the bounds of the given loop.
3361 static LogicalResult canonicalizeLoopBounds(AffineParallelOp op) {
3362   AffineValueMap lb = op.getLowerBoundsValueMap();
3363   bool lbCanonicalized = succeeded(lb.canonicalize());
3364 
3365   AffineValueMap ub = op.getUpperBoundsValueMap();
3366   bool ubCanonicalized = succeeded(ub.canonicalize());
3367 
3368   // Any canonicalization change always leads to updated map(s).
3369   if (!lbCanonicalized && !ubCanonicalized)
3370     return failure();
3371 
3372   if (lbCanonicalized)
3373     op.setLowerBounds(lb.getOperands(), lb.getAffineMap());
3374   if (ubCanonicalized)
3375     op.setUpperBounds(ub.getOperands(), ub.getAffineMap());
3376 
3377   return success();
3378 }
3379 
3380 LogicalResult AffineParallelOp::fold(ArrayRef<Attribute> operands,
3381                                      SmallVectorImpl<OpFoldResult> &results) {
3382   return canonicalizeLoopBounds(*this);
3383 }
3384 
3385 /// Prints a lower(upper) bound of an affine parallel loop with max(min)
3386 /// conditions in it. `mapAttr` is a flat list of affine expressions and `group`
3387 /// identifies which of the those expressions form max/min groups. `operands`
3388 /// are the SSA values of dimensions and symbols and `keyword` is either "min"
3389 /// or "max".
3390 static void printMinMaxBound(OpAsmPrinter &p, AffineMapAttr mapAttr,
3391                              DenseIntElementsAttr group, ValueRange operands,
3392                              StringRef keyword) {
3393   AffineMap map = mapAttr.getValue();
3394   unsigned numDims = map.getNumDims();
3395   ValueRange dimOperands = operands.take_front(numDims);
3396   ValueRange symOperands = operands.drop_front(numDims);
3397   unsigned start = 0;
3398   for (llvm::APInt groupSize : group) {
3399     if (start != 0)
3400       p << ", ";
3401 
3402     unsigned size = groupSize.getZExtValue();
3403     if (size == 1) {
3404       p.printAffineExprOfSSAIds(map.getResult(start), dimOperands, symOperands);
3405       ++start;
3406     } else {
3407       p << keyword << '(';
3408       AffineMap submap = map.getSliceMap(start, size);
3409       p.printAffineMapOfSSAIds(AffineMapAttr::get(submap), operands);
3410       p << ')';
3411       start += size;
3412     }
3413   }
3414 }
3415 
3416 void AffineParallelOp::print(OpAsmPrinter &p) {
3417   p << " (" << getBody()->getArguments() << ") = (";
3418   printMinMaxBound(p, getLowerBoundsMapAttr(), getLowerBoundsGroupsAttr(),
3419                    getLowerBoundsOperands(), "max");
3420   p << ") to (";
3421   printMinMaxBound(p, getUpperBoundsMapAttr(), getUpperBoundsGroupsAttr(),
3422                    getUpperBoundsOperands(), "min");
3423   p << ')';
3424   SmallVector<int64_t, 8> steps = getSteps();
3425   bool elideSteps = llvm::all_of(steps, [](int64_t step) { return step == 1; });
3426   if (!elideSteps) {
3427     p << " step (";
3428     llvm::interleaveComma(steps, p);
3429     p << ')';
3430   }
3431   if (getNumResults()) {
3432     p << " reduce (";
3433     llvm::interleaveComma(getReductions(), p, [&](auto &attr) {
3434       arith::AtomicRMWKind sym = *arith::symbolizeAtomicRMWKind(
3435           attr.template cast<IntegerAttr>().getInt());
3436       p << "\"" << arith::stringifyAtomicRMWKind(sym) << "\"";
3437     });
3438     p << ") -> (" << getResultTypes() << ")";
3439   }
3440 
3441   p << ' ';
3442   p.printRegion(getRegion(), /*printEntryBlockArgs=*/false,
3443                 /*printBlockTerminators=*/getNumResults());
3444   p.printOptionalAttrDict(
3445       (*this)->getAttrs(),
3446       /*elidedAttrs=*/{AffineParallelOp::getReductionsAttrStrName(),
3447                        AffineParallelOp::getLowerBoundsMapAttrStrName(),
3448                        AffineParallelOp::getLowerBoundsGroupsAttrStrName(),
3449                        AffineParallelOp::getUpperBoundsMapAttrStrName(),
3450                        AffineParallelOp::getUpperBoundsGroupsAttrStrName(),
3451                        AffineParallelOp::getStepsAttrStrName()});
3452 }
3453 
3454 /// Given a list of lists of parsed operands, populates `uniqueOperands` with
3455 /// unique operands. Also populates `replacements with affine expressions of
3456 /// `kind` that can be used to update affine maps previously accepting a
3457 /// `operands` to accept `uniqueOperands` instead.
3458 static ParseResult deduplicateAndResolveOperands(
3459     OpAsmParser &parser,
3460     ArrayRef<SmallVector<OpAsmParser::UnresolvedOperand>> operands,
3461     SmallVectorImpl<Value> &uniqueOperands,
3462     SmallVectorImpl<AffineExpr> &replacements, AffineExprKind kind) {
3463   assert((kind == AffineExprKind::DimId || kind == AffineExprKind::SymbolId) &&
3464          "expected operands to be dim or symbol expression");
3465 
3466   Type indexType = parser.getBuilder().getIndexType();
3467   for (const auto &list : operands) {
3468     SmallVector<Value> valueOperands;
3469     if (parser.resolveOperands(list, indexType, valueOperands))
3470       return failure();
3471     for (Value operand : valueOperands) {
3472       unsigned pos = std::distance(uniqueOperands.begin(),
3473                                    llvm::find(uniqueOperands, operand));
3474       if (pos == uniqueOperands.size())
3475         uniqueOperands.push_back(operand);
3476       replacements.push_back(
3477           kind == AffineExprKind::DimId
3478               ? getAffineDimExpr(pos, parser.getContext())
3479               : getAffineSymbolExpr(pos, parser.getContext()));
3480     }
3481   }
3482   return success();
3483 }
3484 
3485 namespace {
3486 enum class MinMaxKind { Min, Max };
3487 } // namespace
3488 
3489 /// Parses an affine map that can contain a min/max for groups of its results,
3490 /// e.g., max(expr-1, expr-2), expr-3, max(expr-4, expr-5, expr-6). Populates
3491 /// `result` attributes with the map (flat list of expressions) and the grouping
3492 /// (list of integers that specify how many expressions to put into each
3493 /// min/max) attributes. Deduplicates repeated operands.
3494 ///
3495 /// parallel-bound       ::= `(` parallel-group-list `)`
3496 /// parallel-group-list  ::= parallel-group (`,` parallel-group-list)?
3497 /// parallel-group       ::= simple-group | min-max-group
3498 /// simple-group         ::= expr-of-ssa-ids
3499 /// min-max-group        ::= ( `min` | `max` ) `(` expr-of-ssa-ids-list `)`
3500 /// expr-of-ssa-ids-list ::= expr-of-ssa-ids (`,` expr-of-ssa-id-list)?
3501 ///
3502 /// Examples:
3503 ///   (%0, min(%1 + %2, %3), %4, min(%5 floordiv 32, %6))
3504 ///   (%0, max(%1 - 2 * %2))
3505 static ParseResult parseAffineMapWithMinMax(OpAsmParser &parser,
3506                                             OperationState &result,
3507                                             MinMaxKind kind) {
3508   constexpr llvm::StringLiteral tmpAttrStrName = "__pseudo_bound_map";
3509 
3510   StringRef mapName = kind == MinMaxKind::Min
3511                           ? AffineParallelOp::getUpperBoundsMapAttrStrName()
3512                           : AffineParallelOp::getLowerBoundsMapAttrStrName();
3513   StringRef groupsName =
3514       kind == MinMaxKind::Min
3515           ? AffineParallelOp::getUpperBoundsGroupsAttrStrName()
3516           : AffineParallelOp::getLowerBoundsGroupsAttrStrName();
3517 
3518   if (failed(parser.parseLParen()))
3519     return failure();
3520 
3521   if (succeeded(parser.parseOptionalRParen())) {
3522     result.addAttribute(
3523         mapName, AffineMapAttr::get(parser.getBuilder().getEmptyAffineMap()));
3524     result.addAttribute(groupsName, parser.getBuilder().getI32TensorAttr({}));
3525     return success();
3526   }
3527 
3528   SmallVector<AffineExpr> flatExprs;
3529   SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatDimOperands;
3530   SmallVector<SmallVector<OpAsmParser::UnresolvedOperand>> flatSymOperands;
3531   SmallVector<int32_t> numMapsPerGroup;
3532   SmallVector<OpAsmParser::UnresolvedOperand> mapOperands;
3533   auto parseOperands = [&]() {
3534     if (succeeded(parser.parseOptionalKeyword(
3535             kind == MinMaxKind::Min ? "min" : "max"))) {
3536       mapOperands.clear();
3537       AffineMapAttr map;
3538       if (failed(parser.parseAffineMapOfSSAIds(mapOperands, map, tmpAttrStrName,
3539                                                result.attributes,
3540                                                OpAsmParser::Delimiter::Paren)))
3541         return failure();
3542       result.attributes.erase(tmpAttrStrName);
3543       llvm::append_range(flatExprs, map.getValue().getResults());
3544       auto operandsRef = llvm::makeArrayRef(mapOperands);
3545       auto dimsRef = operandsRef.take_front(map.getValue().getNumDims());
3546       SmallVector<OpAsmParser::UnresolvedOperand> dims(dimsRef.begin(),
3547                                                        dimsRef.end());
3548       auto symsRef = operandsRef.drop_front(map.getValue().getNumDims());
3549       SmallVector<OpAsmParser::UnresolvedOperand> syms(symsRef.begin(),
3550                                                        symsRef.end());
3551       flatDimOperands.append(map.getValue().getNumResults(), dims);
3552       flatSymOperands.append(map.getValue().getNumResults(), syms);
3553       numMapsPerGroup.push_back(map.getValue().getNumResults());
3554     } else {
3555       if (failed(parser.parseAffineExprOfSSAIds(flatDimOperands.emplace_back(),
3556                                                 flatSymOperands.emplace_back(),
3557                                                 flatExprs.emplace_back())))
3558         return failure();
3559       numMapsPerGroup.push_back(1);
3560     }
3561     return success();
3562   };
3563   if (parser.parseCommaSeparatedList(parseOperands) || parser.parseRParen())
3564     return failure();
3565 
3566   unsigned totalNumDims = 0;
3567   unsigned totalNumSyms = 0;
3568   for (unsigned i = 0, e = flatExprs.size(); i < e; ++i) {
3569     unsigned numDims = flatDimOperands[i].size();
3570     unsigned numSyms = flatSymOperands[i].size();
3571     flatExprs[i] = flatExprs[i]
3572                        .shiftDims(numDims, totalNumDims)
3573                        .shiftSymbols(numSyms, totalNumSyms);
3574     totalNumDims += numDims;
3575     totalNumSyms += numSyms;
3576   }
3577 
3578   // Deduplicate map operands.
3579   SmallVector<Value> dimOperands, symOperands;
3580   SmallVector<AffineExpr> dimRplacements, symRepacements;
3581   if (deduplicateAndResolveOperands(parser, flatDimOperands, dimOperands,
3582                                     dimRplacements, AffineExprKind::DimId) ||
3583       deduplicateAndResolveOperands(parser, flatSymOperands, symOperands,
3584                                     symRepacements, AffineExprKind::SymbolId))
3585     return failure();
3586 
3587   result.operands.append(dimOperands.begin(), dimOperands.end());
3588   result.operands.append(symOperands.begin(), symOperands.end());
3589 
3590   Builder &builder = parser.getBuilder();
3591   auto flatMap = AffineMap::get(totalNumDims, totalNumSyms, flatExprs,
3592                                 parser.getContext());
3593   flatMap = flatMap.replaceDimsAndSymbols(
3594       dimRplacements, symRepacements, dimOperands.size(), symOperands.size());
3595 
3596   result.addAttribute(mapName, AffineMapAttr::get(flatMap));
3597   result.addAttribute(groupsName, builder.getI32TensorAttr(numMapsPerGroup));
3598   return success();
3599 }
3600 
3601 //
3602 // operation ::= `affine.parallel` `(` ssa-ids `)` `=` parallel-bound
3603 //               `to` parallel-bound steps? region attr-dict?
3604 // steps     ::= `steps` `(` integer-literals `)`
3605 //
3606 ParseResult AffineParallelOp::parse(OpAsmParser &parser,
3607                                     OperationState &result) {
3608   auto &builder = parser.getBuilder();
3609   auto indexType = builder.getIndexType();
3610   SmallVector<OpAsmParser::Argument, 4> ivs;
3611   if (parser.parseArgumentList(ivs, OpAsmParser::Delimiter::Paren) ||
3612       parser.parseEqual() ||
3613       parseAffineMapWithMinMax(parser, result, MinMaxKind::Max) ||
3614       parser.parseKeyword("to") ||
3615       parseAffineMapWithMinMax(parser, result, MinMaxKind::Min))
3616     return failure();
3617 
3618   AffineMapAttr stepsMapAttr;
3619   NamedAttrList stepsAttrs;
3620   SmallVector<OpAsmParser::UnresolvedOperand, 4> stepsMapOperands;
3621   if (failed(parser.parseOptionalKeyword("step"))) {
3622     SmallVector<int64_t, 4> steps(ivs.size(), 1);
3623     result.addAttribute(AffineParallelOp::getStepsAttrStrName(),
3624                         builder.getI64ArrayAttr(steps));
3625   } else {
3626     if (parser.parseAffineMapOfSSAIds(stepsMapOperands, stepsMapAttr,
3627                                       AffineParallelOp::getStepsAttrStrName(),
3628                                       stepsAttrs,
3629                                       OpAsmParser::Delimiter::Paren))
3630       return failure();
3631 
3632     // Convert steps from an AffineMap into an I64ArrayAttr.
3633     SmallVector<int64_t, 4> steps;
3634     auto stepsMap = stepsMapAttr.getValue();
3635     for (const auto &result : stepsMap.getResults()) {
3636       auto constExpr = result.dyn_cast<AffineConstantExpr>();
3637       if (!constExpr)
3638         return parser.emitError(parser.getNameLoc(),
3639                                 "steps must be constant integers");
3640       steps.push_back(constExpr.getValue());
3641     }
3642     result.addAttribute(AffineParallelOp::getStepsAttrStrName(),
3643                         builder.getI64ArrayAttr(steps));
3644   }
3645 
3646   // Parse optional clause of the form: `reduce ("addf", "maxf")`, where the
3647   // quoted strings are a member of the enum AtomicRMWKind.
3648   SmallVector<Attribute, 4> reductions;
3649   if (succeeded(parser.parseOptionalKeyword("reduce"))) {
3650     if (parser.parseLParen())
3651       return failure();
3652     auto parseAttributes = [&]() -> ParseResult {
3653       // Parse a single quoted string via the attribute parsing, and then
3654       // verify it is a member of the enum and convert to it's integer
3655       // representation.
3656       StringAttr attrVal;
3657       NamedAttrList attrStorage;
3658       auto loc = parser.getCurrentLocation();
3659       if (parser.parseAttribute(attrVal, builder.getNoneType(), "reduce",
3660                                 attrStorage))
3661         return failure();
3662       llvm::Optional<arith::AtomicRMWKind> reduction =
3663           arith::symbolizeAtomicRMWKind(attrVal.getValue());
3664       if (!reduction)
3665         return parser.emitError(loc, "invalid reduction value: ") << attrVal;
3666       reductions.push_back(builder.getI64IntegerAttr(
3667           static_cast<int64_t>(reduction.getValue())));
3668       // While we keep getting commas, keep parsing.
3669       return success();
3670     };
3671     if (parser.parseCommaSeparatedList(parseAttributes) || parser.parseRParen())
3672       return failure();
3673   }
3674   result.addAttribute(AffineParallelOp::getReductionsAttrStrName(),
3675                       builder.getArrayAttr(reductions));
3676 
3677   // Parse return types of reductions (if any)
3678   if (parser.parseOptionalArrowTypeList(result.types))
3679     return failure();
3680 
3681   // Now parse the body.
3682   Region *body = result.addRegion();
3683   for (auto &iv : ivs)
3684     iv.type = indexType;
3685   if (parser.parseRegion(*body, ivs) ||
3686       parser.parseOptionalAttrDict(result.attributes))
3687     return failure();
3688 
3689   // Add a terminator if none was parsed.
3690   AffineParallelOp::ensureTerminator(*body, builder, result.location);
3691   return success();
3692 }
3693 
3694 //===----------------------------------------------------------------------===//
3695 // AffineYieldOp
3696 //===----------------------------------------------------------------------===//
3697 
3698 LogicalResult AffineYieldOp::verify() {
3699   auto *parentOp = (*this)->getParentOp();
3700   auto results = parentOp->getResults();
3701   auto operands = getOperands();
3702 
3703   if (!isa<AffineParallelOp, AffineIfOp, AffineForOp>(parentOp))
3704     return emitOpError() << "only terminates affine.if/for/parallel regions";
3705   if (parentOp->getNumResults() != getNumOperands())
3706     return emitOpError() << "parent of yield must have same number of "
3707                             "results as the yield operands";
3708   for (auto it : llvm::zip(results, operands)) {
3709     if (std::get<0>(it).getType() != std::get<1>(it).getType())
3710       return emitOpError() << "types mismatch between yield op and its parent";
3711   }
3712 
3713   return success();
3714 }
3715 
3716 //===----------------------------------------------------------------------===//
3717 // AffineVectorLoadOp
3718 //===----------------------------------------------------------------------===//
3719 
3720 void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result,
3721                                VectorType resultType, AffineMap map,
3722                                ValueRange operands) {
3723   assert(operands.size() == 1 + map.getNumInputs() && "inconsistent operands");
3724   result.addOperands(operands);
3725   if (map)
3726     result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3727   result.types.push_back(resultType);
3728 }
3729 
3730 void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result,
3731                                VectorType resultType, Value memref,
3732                                AffineMap map, ValueRange mapOperands) {
3733   assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
3734   result.addOperands(memref);
3735   result.addOperands(mapOperands);
3736   result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3737   result.types.push_back(resultType);
3738 }
3739 
3740 void AffineVectorLoadOp::build(OpBuilder &builder, OperationState &result,
3741                                VectorType resultType, Value memref,
3742                                ValueRange indices) {
3743   auto memrefType = memref.getType().cast<MemRefType>();
3744   int64_t rank = memrefType.getRank();
3745   // Create identity map for memrefs with at least one dimension or () -> ()
3746   // for zero-dimensional memrefs.
3747   auto map =
3748       rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap();
3749   build(builder, result, resultType, memref, map, indices);
3750 }
3751 
3752 void AffineVectorLoadOp::getCanonicalizationPatterns(RewritePatternSet &results,
3753                                                      MLIRContext *context) {
3754   results.add<SimplifyAffineOp<AffineVectorLoadOp>>(context);
3755 }
3756 
3757 ParseResult AffineVectorLoadOp::parse(OpAsmParser &parser,
3758                                       OperationState &result) {
3759   auto &builder = parser.getBuilder();
3760   auto indexTy = builder.getIndexType();
3761 
3762   MemRefType memrefType;
3763   VectorType resultType;
3764   OpAsmParser::UnresolvedOperand memrefInfo;
3765   AffineMapAttr mapAttr;
3766   SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
3767   return failure(
3768       parser.parseOperand(memrefInfo) ||
3769       parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
3770                                     AffineVectorLoadOp::getMapAttrStrName(),
3771                                     result.attributes) ||
3772       parser.parseOptionalAttrDict(result.attributes) ||
3773       parser.parseColonType(memrefType) || parser.parseComma() ||
3774       parser.parseType(resultType) ||
3775       parser.resolveOperand(memrefInfo, memrefType, result.operands) ||
3776       parser.resolveOperands(mapOperands, indexTy, result.operands) ||
3777       parser.addTypeToList(resultType, result.types));
3778 }
3779 
3780 void AffineVectorLoadOp::print(OpAsmPrinter &p) {
3781   p << " " << getMemRef() << '[';
3782   if (AffineMapAttr mapAttr =
3783           (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
3784     p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
3785   p << ']';
3786   p.printOptionalAttrDict((*this)->getAttrs(),
3787                           /*elidedAttrs=*/{getMapAttrStrName()});
3788   p << " : " << getMemRefType() << ", " << getType();
3789 }
3790 
3791 /// Verify common invariants of affine.vector_load and affine.vector_store.
3792 static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType,
3793                                           VectorType vectorType) {
3794   // Check that memref and vector element types match.
3795   if (memrefType.getElementType() != vectorType.getElementType())
3796     return op->emitOpError(
3797         "requires memref and vector types of the same elemental type");
3798   return success();
3799 }
3800 
3801 LogicalResult AffineVectorLoadOp::verify() {
3802   MemRefType memrefType = getMemRefType();
3803   if (failed(verifyMemoryOpIndexing(
3804           getOperation(),
3805           (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
3806           getMapOperands(), memrefType,
3807           /*numIndexOperands=*/getNumOperands() - 1)))
3808     return failure();
3809 
3810   if (failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType())))
3811     return failure();
3812 
3813   return success();
3814 }
3815 
3816 //===----------------------------------------------------------------------===//
3817 // AffineVectorStoreOp
3818 //===----------------------------------------------------------------------===//
3819 
3820 void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &result,
3821                                 Value valueToStore, Value memref, AffineMap map,
3822                                 ValueRange mapOperands) {
3823   assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
3824   result.addOperands(valueToStore);
3825   result.addOperands(memref);
3826   result.addOperands(mapOperands);
3827   result.addAttribute(getMapAttrStrName(), AffineMapAttr::get(map));
3828 }
3829 
3830 // Use identity map.
3831 void AffineVectorStoreOp::build(OpBuilder &builder, OperationState &result,
3832                                 Value valueToStore, Value memref,
3833                                 ValueRange indices) {
3834   auto memrefType = memref.getType().cast<MemRefType>();
3835   int64_t rank = memrefType.getRank();
3836   // Create identity map for memrefs with at least one dimension or () -> ()
3837   // for zero-dimensional memrefs.
3838   auto map =
3839       rank ? builder.getMultiDimIdentityMap(rank) : builder.getEmptyAffineMap();
3840   build(builder, result, valueToStore, memref, map, indices);
3841 }
3842 void AffineVectorStoreOp::getCanonicalizationPatterns(
3843     RewritePatternSet &results, MLIRContext *context) {
3844   results.add<SimplifyAffineOp<AffineVectorStoreOp>>(context);
3845 }
3846 
3847 ParseResult AffineVectorStoreOp::parse(OpAsmParser &parser,
3848                                        OperationState &result) {
3849   auto indexTy = parser.getBuilder().getIndexType();
3850 
3851   MemRefType memrefType;
3852   VectorType resultType;
3853   OpAsmParser::UnresolvedOperand storeValueInfo;
3854   OpAsmParser::UnresolvedOperand memrefInfo;
3855   AffineMapAttr mapAttr;
3856   SmallVector<OpAsmParser::UnresolvedOperand, 1> mapOperands;
3857   return failure(
3858       parser.parseOperand(storeValueInfo) || parser.parseComma() ||
3859       parser.parseOperand(memrefInfo) ||
3860       parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
3861                                     AffineVectorStoreOp::getMapAttrStrName(),
3862                                     result.attributes) ||
3863       parser.parseOptionalAttrDict(result.attributes) ||
3864       parser.parseColonType(memrefType) || parser.parseComma() ||
3865       parser.parseType(resultType) ||
3866       parser.resolveOperand(storeValueInfo, resultType, result.operands) ||
3867       parser.resolveOperand(memrefInfo, memrefType, result.operands) ||
3868       parser.resolveOperands(mapOperands, indexTy, result.operands));
3869 }
3870 
3871 void AffineVectorStoreOp::print(OpAsmPrinter &p) {
3872   p << " " << getValueToStore();
3873   p << ", " << getMemRef() << '[';
3874   if (AffineMapAttr mapAttr =
3875           (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()))
3876     p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
3877   p << ']';
3878   p.printOptionalAttrDict((*this)->getAttrs(),
3879                           /*elidedAttrs=*/{getMapAttrStrName()});
3880   p << " : " << getMemRefType() << ", " << getValueToStore().getType();
3881 }
3882 
3883 LogicalResult AffineVectorStoreOp::verify() {
3884   MemRefType memrefType = getMemRefType();
3885   if (failed(verifyMemoryOpIndexing(
3886           *this, (*this)->getAttrOfType<AffineMapAttr>(getMapAttrStrName()),
3887           getMapOperands(), memrefType,
3888           /*numIndexOperands=*/getNumOperands() - 2)))
3889     return failure();
3890 
3891   if (failed(verifyVectorMemoryOp(*this, memrefType, getVectorType())))
3892     return failure();
3893 
3894   return success();
3895 }
3896 
3897 //===----------------------------------------------------------------------===//
3898 // TableGen'd op method definitions
3899 //===----------------------------------------------------------------------===//
3900 
3901 #define GET_OP_CLASSES
3902 #include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc"
3903