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/StandardOps/IR/Ops.h"
12 #include "mlir/IR/Function.h"
13 #include "mlir/IR/IntegerSet.h"
14 #include "mlir/IR/Matchers.h"
15 #include "mlir/IR/OpImplementation.h"
16 #include "mlir/IR/PatternMatch.h"
17 #include "mlir/Transforms/InliningUtils.h"
18 #include "llvm/ADT/SetVector.h"
19 #include "llvm/ADT/SmallBitVector.h"
20 #include "llvm/Support/Debug.h"
21 
22 using namespace mlir;
23 using llvm::dbgs;
24 
25 #define DEBUG_TYPE "affine-analysis"
26 
27 //===----------------------------------------------------------------------===//
28 // AffineDialect Interfaces
29 //===----------------------------------------------------------------------===//
30 
31 namespace {
32 /// This class defines the interface for handling inlining with affine
33 /// operations.
34 struct AffineInlinerInterface : public DialectInlinerInterface {
35   using DialectInlinerInterface::DialectInlinerInterface;
36 
37   //===--------------------------------------------------------------------===//
38   // Analysis Hooks
39   //===--------------------------------------------------------------------===//
40 
41   /// Returns true if the given region 'src' can be inlined into the region
42   /// 'dest' that is attached to an operation registered to the current dialect.
43   bool isLegalToInline(Region *dest, Region *src,
44                        BlockAndValueMapping &valueMapping) const final {
45     // Conservatively don't allow inlining into affine structures.
46     return false;
47   }
48 
49   /// Returns true if the given operation 'op', that is registered to this
50   /// dialect, can be inlined into the given region, false otherwise.
51   bool isLegalToInline(Operation *op, Region *region,
52                        BlockAndValueMapping &valueMapping) const final {
53     // Always allow inlining affine operations into the top-level region of a
54     // function. There are some edge cases when inlining *into* affine
55     // structures, but that is handled in the other 'isLegalToInline' hook
56     // above.
57     // TODO: We should be able to inline into other regions than functions.
58     return isa<FuncOp>(region->getParentOp());
59   }
60 
61   /// Affine regions should be analyzed recursively.
62   bool shouldAnalyzeRecursively(Operation *op) const final { return true; }
63 };
64 } // end anonymous namespace
65 
66 //===----------------------------------------------------------------------===//
67 // AffineDialect
68 //===----------------------------------------------------------------------===//
69 
70 AffineDialect::AffineDialect(MLIRContext *context)
71     : Dialect(getDialectNamespace(), context) {
72   addOperations<AffineDmaStartOp, AffineDmaWaitOp, AffineLoadOp, AffineStoreOp,
73 #define GET_OP_LIST
74 #include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc"
75                 >();
76   addInterfaces<AffineInlinerInterface>();
77 }
78 
79 /// Materialize a single constant operation from a given attribute value with
80 /// the desired resultant type.
81 Operation *AffineDialect::materializeConstant(OpBuilder &builder,
82                                               Attribute value, Type type,
83                                               Location loc) {
84   return builder.create<ConstantOp>(loc, type, value);
85 }
86 
87 /// A utility function to check if a given region is attached to a function.
88 static bool isFunctionRegion(Region *region) {
89   return llvm::isa<FuncOp>(region->getParentOp());
90 }
91 
92 /// A utility function to check if a value is defined at the top level of a
93 /// function. A value of index type defined at the top level is always a valid
94 /// symbol.
95 bool mlir::isTopLevelValue(Value value) {
96   if (auto arg = value.dyn_cast<BlockArgument>())
97     return isFunctionRegion(arg.getOwner()->getParent());
98   return isFunctionRegion(value.getDefiningOp()->getParentRegion());
99 }
100 
101 // Value can be used as a dimension id if it is valid as a symbol, or
102 // it is an induction variable, or it is a result of affine apply operation
103 // with dimension id arguments.
104 bool mlir::isValidDim(Value value) {
105   // The value must be an index type.
106   if (!value.getType().isIndex())
107     return false;
108 
109   if (auto *op = value.getDefiningOp()) {
110     // Top level operation or constant operation is ok.
111     if (isFunctionRegion(op->getParentRegion()) || isa<ConstantOp>(op))
112       return true;
113     // Affine apply operation is ok if all of its operands are ok.
114     if (auto applyOp = dyn_cast<AffineApplyOp>(op))
115       return applyOp.isValidDim();
116     // The dim op is okay if its operand memref/tensor is defined at the top
117     // level.
118     if (auto dimOp = dyn_cast<DimOp>(op))
119       return isTopLevelValue(dimOp.getOperand());
120     return false;
121   }
122   // This value has to be a block argument of a FuncOp, an 'affine.for', or an
123   // 'affine.parallel'.
124   auto *parentOp = value.cast<BlockArgument>().getOwner()->getParentOp();
125   return isa<FuncOp>(parentOp) || isa<AffineForOp>(parentOp) ||
126          isa<AffineParallelOp>(parentOp);
127 }
128 
129 /// Returns true if the 'index' dimension of the `memref` defined by
130 /// `memrefDefOp` is a statically  shaped one or defined using a valid symbol.
131 template <typename AnyMemRefDefOp>
132 static bool isMemRefSizeValidSymbol(AnyMemRefDefOp memrefDefOp,
133                                     unsigned index) {
134   auto memRefType = memrefDefOp.getType();
135   // Statically shaped.
136   if (!memRefType.isDynamicDim(index))
137     return true;
138   // Get the position of the dimension among dynamic dimensions;
139   unsigned dynamicDimPos = memRefType.getDynamicDimIndex(index);
140   return isValidSymbol(
141       *(memrefDefOp.getDynamicSizes().begin() + dynamicDimPos));
142 }
143 
144 /// Returns true if the result of the dim op is a valid symbol.
145 static bool isDimOpValidSymbol(DimOp dimOp) {
146   // The dim op is okay if its operand memref/tensor is defined at the top
147   // level.
148   if (isTopLevelValue(dimOp.getOperand()))
149     return true;
150 
151   // The dim op is also okay if its operand memref/tensor is a view/subview
152   // whose corresponding size is a valid symbol.
153   unsigned index = dimOp.getIndex();
154   if (auto viewOp = dyn_cast<ViewOp>(dimOp.getOperand().getDefiningOp()))
155     return isMemRefSizeValidSymbol<ViewOp>(viewOp, index);
156   if (auto subViewOp = dyn_cast<SubViewOp>(dimOp.getOperand().getDefiningOp()))
157     return isMemRefSizeValidSymbol<SubViewOp>(subViewOp, index);
158   if (auto allocOp = dyn_cast<AllocOp>(dimOp.getOperand().getDefiningOp()))
159     return isMemRefSizeValidSymbol<AllocOp>(allocOp, index);
160   return false;
161 }
162 
163 // Value can be used as a symbol if it is a constant, or it is defined at
164 // the top level, or it is a result of affine apply operation with symbol
165 // arguments, or a result of the dim op on a memref satisfying certain
166 // constraints.
167 bool mlir::isValidSymbol(Value value) {
168   // The value must be an index type.
169   if (!value.getType().isIndex())
170     return false;
171 
172   if (auto *op = value.getDefiningOp()) {
173     // Top level operation or constant operation is ok.
174     if (isFunctionRegion(op->getParentRegion()) || isa<ConstantOp>(op))
175       return true;
176     // Affine apply operation is ok if all of its operands are ok.
177     if (auto applyOp = dyn_cast<AffineApplyOp>(op))
178       return applyOp.isValidSymbol();
179     if (auto dimOp = dyn_cast<DimOp>(op)) {
180       return isDimOpValidSymbol(dimOp);
181     }
182   }
183   // Otherwise, check that the value is a top level value.
184   return isTopLevelValue(value);
185 }
186 
187 // Returns true if 'value' is a valid index to an affine operation (e.g.
188 // affine.load, affine.store, affine.dma_start, affine.dma_wait).
189 // Returns false otherwise.
190 static bool isValidAffineIndexOperand(Value value) {
191   return isValidDim(value) || isValidSymbol(value);
192 }
193 
194 /// Utility function to verify that a set of operands are valid dimension and
195 /// symbol identifiers. The operands should be laid out such that the dimension
196 /// operands are before the symbol operands. This function returns failure if
197 /// there was an invalid operand. An operation is provided to emit any necessary
198 /// errors.
199 template <typename OpTy>
200 static LogicalResult
201 verifyDimAndSymbolIdentifiers(OpTy &op, Operation::operand_range operands,
202                               unsigned numDims) {
203   unsigned opIt = 0;
204   for (auto operand : operands) {
205     if (opIt++ < numDims) {
206       if (!isValidDim(operand))
207         return op.emitOpError("operand cannot be used as a dimension id");
208     } else if (!isValidSymbol(operand)) {
209       return op.emitOpError("operand cannot be used as a symbol");
210     }
211   }
212   return success();
213 }
214 
215 //===----------------------------------------------------------------------===//
216 // AffineApplyOp
217 //===----------------------------------------------------------------------===//
218 
219 AffineValueMap AffineApplyOp::getAffineValueMap() {
220   return AffineValueMap(getAffineMap(), getOperands(), getResult());
221 }
222 
223 static ParseResult parseAffineApplyOp(OpAsmParser &parser,
224                                       OperationState &result) {
225   auto &builder = parser.getBuilder();
226   auto indexTy = builder.getIndexType();
227 
228   AffineMapAttr mapAttr;
229   unsigned numDims;
230   if (parser.parseAttribute(mapAttr, "map", result.attributes) ||
231       parseDimAndSymbolList(parser, result.operands, numDims) ||
232       parser.parseOptionalAttrDict(result.attributes))
233     return failure();
234   auto map = mapAttr.getValue();
235 
236   if (map.getNumDims() != numDims ||
237       numDims + map.getNumSymbols() != result.operands.size()) {
238     return parser.emitError(parser.getNameLoc(),
239                             "dimension or symbol index mismatch");
240   }
241 
242   result.types.append(map.getNumResults(), indexTy);
243   return success();
244 }
245 
246 static void print(OpAsmPrinter &p, AffineApplyOp op) {
247   p << AffineApplyOp::getOperationName() << " " << op.mapAttr();
248   printDimAndSymbolList(op.operand_begin(), op.operand_end(),
249                         op.getAffineMap().getNumDims(), p);
250   p.printOptionalAttrDict(op.getAttrs(), /*elidedAttrs=*/{"map"});
251 }
252 
253 static LogicalResult verify(AffineApplyOp op) {
254   // Check input and output dimensions match.
255   auto map = op.map();
256 
257   // Verify that operand count matches affine map dimension and symbol count.
258   if (op.getNumOperands() != map.getNumDims() + map.getNumSymbols())
259     return op.emitOpError(
260         "operand count and affine map dimension and symbol count must match");
261 
262   // Verify that the map only produces one result.
263   if (map.getNumResults() != 1)
264     return op.emitOpError("mapping must produce one value");
265 
266   return success();
267 }
268 
269 // The result of the affine apply operation can be used as a dimension id if all
270 // its operands are valid dimension ids.
271 bool AffineApplyOp::isValidDim() {
272   return llvm::all_of(getOperands(),
273                       [](Value op) { return mlir::isValidDim(op); });
274 }
275 
276 // The result of the affine apply operation can be used as a symbol if all its
277 // operands are symbols.
278 bool AffineApplyOp::isValidSymbol() {
279   return llvm::all_of(getOperands(),
280                       [](Value op) { return mlir::isValidSymbol(op); });
281 }
282 
283 OpFoldResult AffineApplyOp::fold(ArrayRef<Attribute> operands) {
284   auto map = getAffineMap();
285 
286   // Fold dims and symbols to existing values.
287   auto expr = map.getResult(0);
288   if (auto dim = expr.dyn_cast<AffineDimExpr>())
289     return getOperand(dim.getPosition());
290   if (auto sym = expr.dyn_cast<AffineSymbolExpr>())
291     return getOperand(map.getNumDims() + sym.getPosition());
292 
293   // Otherwise, default to folding the map.
294   SmallVector<Attribute, 1> result;
295   if (failed(map.constantFold(operands, result)))
296     return {};
297   return result[0];
298 }
299 
300 AffineDimExpr AffineApplyNormalizer::renumberOneDim(Value v) {
301   DenseMap<Value, unsigned>::iterator iterPos;
302   bool inserted = false;
303   std::tie(iterPos, inserted) =
304       dimValueToPosition.insert(std::make_pair(v, dimValueToPosition.size()));
305   if (inserted) {
306     reorderedDims.push_back(v);
307   }
308   return getAffineDimExpr(iterPos->second, v.getContext())
309       .cast<AffineDimExpr>();
310 }
311 
312 AffineMap AffineApplyNormalizer::renumber(const AffineApplyNormalizer &other) {
313   SmallVector<AffineExpr, 8> dimRemapping;
314   for (auto v : other.reorderedDims) {
315     auto kvp = other.dimValueToPosition.find(v);
316     if (dimRemapping.size() <= kvp->second)
317       dimRemapping.resize(kvp->second + 1);
318     dimRemapping[kvp->second] = renumberOneDim(kvp->first);
319   }
320   unsigned numSymbols = concatenatedSymbols.size();
321   unsigned numOtherSymbols = other.concatenatedSymbols.size();
322   SmallVector<AffineExpr, 8> symRemapping(numOtherSymbols);
323   for (unsigned idx = 0; idx < numOtherSymbols; ++idx) {
324     symRemapping[idx] =
325         getAffineSymbolExpr(idx + numSymbols, other.affineMap.getContext());
326   }
327   concatenatedSymbols.insert(concatenatedSymbols.end(),
328                              other.concatenatedSymbols.begin(),
329                              other.concatenatedSymbols.end());
330   auto map = other.affineMap;
331   return map.replaceDimsAndSymbols(dimRemapping, symRemapping,
332                                    reorderedDims.size(),
333                                    concatenatedSymbols.size());
334 }
335 
336 // Gather the positions of the operands that are produced by an AffineApplyOp.
337 static llvm::SetVector<unsigned>
338 indicesFromAffineApplyOp(ArrayRef<Value> operands) {
339   llvm::SetVector<unsigned> res;
340   for (auto en : llvm::enumerate(operands))
341     if (isa_and_nonnull<AffineApplyOp>(en.value().getDefiningOp()))
342       res.insert(en.index());
343   return res;
344 }
345 
346 // Support the special case of a symbol coming from an AffineApplyOp that needs
347 // to be composed into the current AffineApplyOp.
348 // This case is handled by rewriting all such symbols into dims for the purpose
349 // of allowing mathematical AffineMap composition.
350 // Returns an AffineMap where symbols that come from an AffineApplyOp have been
351 // rewritten as dims and are ordered after the original dims.
352 // TODO(andydavis,ntv): This promotion makes AffineMap lose track of which
353 // symbols are represented as dims. This loss is static but can still be
354 // recovered dynamically (with `isValidSymbol`). Still this is annoying for the
355 // semi-affine map case. A dynamic canonicalization of all dims that are valid
356 // symbols (a.k.a `canonicalizePromotedSymbols`) into symbols helps and even
357 // results in better simplifications and foldings. But we should evaluate
358 // whether this behavior is what we really want after using more.
359 static AffineMap promoteComposedSymbolsAsDims(AffineMap map,
360                                               ArrayRef<Value> symbols) {
361   if (symbols.empty()) {
362     return map;
363   }
364 
365   // Sanity check on symbols.
366   for (auto sym : symbols) {
367     assert(isValidSymbol(sym) && "Expected only valid symbols");
368     (void)sym;
369   }
370 
371   // Extract the symbol positions that come from an AffineApplyOp and
372   // needs to be rewritten as dims.
373   auto symPositions = indicesFromAffineApplyOp(symbols);
374   if (symPositions.empty()) {
375     return map;
376   }
377 
378   // Create the new map by replacing each symbol at pos by the next new dim.
379   unsigned numDims = map.getNumDims();
380   unsigned numSymbols = map.getNumSymbols();
381   unsigned numNewDims = 0;
382   unsigned numNewSymbols = 0;
383   SmallVector<AffineExpr, 8> symReplacements(numSymbols);
384   for (unsigned i = 0; i < numSymbols; ++i) {
385     symReplacements[i] =
386         symPositions.count(i) > 0
387             ? getAffineDimExpr(numDims + numNewDims++, map.getContext())
388             : getAffineSymbolExpr(numNewSymbols++, map.getContext());
389   }
390   assert(numSymbols >= numNewDims);
391   AffineMap newMap = map.replaceDimsAndSymbols(
392       {}, symReplacements, numDims + numNewDims, numNewSymbols);
393 
394   return newMap;
395 }
396 
397 /// The AffineNormalizer composes AffineApplyOp recursively. Its purpose is to
398 /// keep a correspondence between the mathematical `map` and the `operands` of
399 /// a given AffineApplyOp. This correspondence is maintained by iterating over
400 /// the operands and forming an `auxiliaryMap` that can be composed
401 /// mathematically with `map`. To keep this correspondence in cases where
402 /// symbols are produced by affine.apply operations, we perform a local rewrite
403 /// of symbols as dims.
404 ///
405 /// Rationale for locally rewriting symbols as dims:
406 /// ================================================
407 /// The mathematical composition of AffineMap must always concatenate symbols
408 /// because it does not have enough information to do otherwise. For example,
409 /// composing `(d0)[s0] -> (d0 + s0)` with itself must produce
410 /// `(d0)[s0, s1] -> (d0 + s0 + s1)`.
411 ///
412 /// The result is only equivalent to `(d0)[s0] -> (d0 + 2 * s0)` when
413 /// applied to the same mlir::Value for both s0 and s1.
414 /// As a consequence mathematical composition of AffineMap always concatenates
415 /// symbols.
416 ///
417 /// When AffineMaps are used in AffineApplyOp however, they may specify
418 /// composition via symbols, which is ambiguous mathematically. This corner case
419 /// is handled by locally rewriting such symbols that come from AffineApplyOp
420 /// into dims and composing through dims.
421 /// TODO(andydavis, ntv): Composition via symbols comes at a significant code
422 /// complexity. Alternatively we should investigate whether we want to
423 /// explicitly disallow symbols coming from affine.apply and instead force the
424 /// user to compose symbols beforehand. The annoyances may be small (i.e. 1 or 2
425 /// extra API calls for such uses, which haven't popped up until now) and the
426 /// benefit potentially big: simpler and more maintainable code for a
427 /// non-trivial, recursive, procedure.
428 AffineApplyNormalizer::AffineApplyNormalizer(AffineMap map,
429                                              ArrayRef<Value> operands)
430     : AffineApplyNormalizer() {
431   static_assert(kMaxAffineApplyDepth > 0, "kMaxAffineApplyDepth must be > 0");
432   assert(map.getNumInputs() == operands.size() &&
433          "number of operands does not match the number of map inputs");
434 
435   LLVM_DEBUG(map.print(dbgs() << "\nInput map: "));
436 
437   // Promote symbols that come from an AffineApplyOp to dims by rewriting the
438   // map to always refer to:
439   //   (dims, symbols coming from AffineApplyOp, other symbols).
440   // The order of operands can remain unchanged.
441   // This is a simplification that relies on 2 ordering properties:
442   //   1. rewritten symbols always appear after the original dims in the map;
443   //   2. operands are traversed in order and either dispatched to:
444   //      a. auxiliaryExprs (dims and symbols rewritten as dims);
445   //      b. concatenatedSymbols (all other symbols)
446   // This allows operand order to remain unchanged.
447   unsigned numDimsBeforeRewrite = map.getNumDims();
448   map = promoteComposedSymbolsAsDims(map,
449                                      operands.take_back(map.getNumSymbols()));
450 
451   LLVM_DEBUG(map.print(dbgs() << "\nRewritten map: "));
452 
453   SmallVector<AffineExpr, 8> auxiliaryExprs;
454   bool furtherCompose = (affineApplyDepth() <= kMaxAffineApplyDepth);
455   // We fully spell out the 2 cases below. In this particular instance a little
456   // code duplication greatly improves readability.
457   // Note that the first branch would disappear if we only supported full
458   // composition (i.e. infinite kMaxAffineApplyDepth).
459   if (!furtherCompose) {
460     // 1. Only dispatch dims or symbols.
461     for (auto en : llvm::enumerate(operands)) {
462       auto t = en.value();
463       assert(t.getType().isIndex());
464       bool isDim = (en.index() < map.getNumDims());
465       if (isDim) {
466         // a. The mathematical composition of AffineMap composes dims.
467         auxiliaryExprs.push_back(renumberOneDim(t));
468       } else {
469         // b. The mathematical composition of AffineMap concatenates symbols.
470         //    We do the same for symbol operands.
471         concatenatedSymbols.push_back(t);
472       }
473     }
474   } else {
475     assert(numDimsBeforeRewrite <= operands.size());
476     // 2. Compose AffineApplyOps and dispatch dims or symbols.
477     for (unsigned i = 0, e = operands.size(); i < e; ++i) {
478       auto t = operands[i];
479       auto affineApply = dyn_cast_or_null<AffineApplyOp>(t.getDefiningOp());
480       if (affineApply) {
481         // a. Compose affine.apply operations.
482         LLVM_DEBUG(affineApply.getOperation()->print(
483             dbgs() << "\nCompose AffineApplyOp recursively: "));
484         AffineMap affineApplyMap = affineApply.getAffineMap();
485         SmallVector<Value, 8> affineApplyOperands(
486             affineApply.getOperands().begin(), affineApply.getOperands().end());
487         AffineApplyNormalizer normalizer(affineApplyMap, affineApplyOperands);
488 
489         LLVM_DEBUG(normalizer.affineMap.print(
490             dbgs() << "\nRenumber into current normalizer: "));
491 
492         auto renumberedMap = renumber(normalizer);
493 
494         LLVM_DEBUG(
495             renumberedMap.print(dbgs() << "\nRecursive composition yields: "));
496 
497         auxiliaryExprs.push_back(renumberedMap.getResult(0));
498       } else {
499         if (i < numDimsBeforeRewrite) {
500           // b. The mathematical composition of AffineMap composes dims.
501           auxiliaryExprs.push_back(renumberOneDim(t));
502         } else {
503           // c. The mathematical composition of AffineMap concatenates symbols.
504           //    Note that the map composition will put symbols already present
505           //    in the map before any symbols coming from the auxiliary map, so
506           //    we insert them before any symbols that are due to renumbering,
507           //    and after the proper symbols we have seen already.
508           concatenatedSymbols.insert(
509               std::next(concatenatedSymbols.begin(), numProperSymbols++), t);
510         }
511       }
512     }
513   }
514 
515   // Early exit if `map` is already composed.
516   if (auxiliaryExprs.empty()) {
517     affineMap = map;
518     return;
519   }
520 
521   assert(concatenatedSymbols.size() >= map.getNumSymbols() &&
522          "Unexpected number of concatenated symbols");
523   auto numDims = dimValueToPosition.size();
524   auto numSymbols = concatenatedSymbols.size() - map.getNumSymbols();
525   auto auxiliaryMap =
526       AffineMap::get(numDims, numSymbols, auxiliaryExprs, map.getContext());
527 
528   LLVM_DEBUG(map.print(dbgs() << "\nCompose map: "));
529   LLVM_DEBUG(auxiliaryMap.print(dbgs() << "\nWith map: "));
530   LLVM_DEBUG(map.compose(auxiliaryMap).print(dbgs() << "\nResult: "));
531 
532   // TODO(andydavis,ntv): Disabling simplification results in major speed gains.
533   // Another option is to cache the results as it is expected a lot of redundant
534   // work is performed in practice.
535   affineMap = simplifyAffineMap(map.compose(auxiliaryMap));
536 
537   LLVM_DEBUG(affineMap.print(dbgs() << "\nSimplified result: "));
538   LLVM_DEBUG(dbgs() << "\n");
539 }
540 
541 void AffineApplyNormalizer::normalize(AffineMap *otherMap,
542                                       SmallVectorImpl<Value> *otherOperands) {
543   AffineApplyNormalizer other(*otherMap, *otherOperands);
544   *otherMap = renumber(other);
545 
546   otherOperands->reserve(reorderedDims.size() + concatenatedSymbols.size());
547   otherOperands->assign(reorderedDims.begin(), reorderedDims.end());
548   otherOperands->append(concatenatedSymbols.begin(), concatenatedSymbols.end());
549 }
550 
551 /// Implements `map` and `operands` composition and simplification to support
552 /// `makeComposedAffineApply`. This can be called to achieve the same effects
553 /// on `map` and `operands` without creating an AffineApplyOp that needs to be
554 /// immediately deleted.
555 static void composeAffineMapAndOperands(AffineMap *map,
556                                         SmallVectorImpl<Value> *operands) {
557   AffineApplyNormalizer normalizer(*map, *operands);
558   auto normalizedMap = normalizer.getAffineMap();
559   auto normalizedOperands = normalizer.getOperands();
560   canonicalizeMapAndOperands(&normalizedMap, &normalizedOperands);
561   *map = normalizedMap;
562   *operands = normalizedOperands;
563   assert(*map);
564 }
565 
566 void mlir::fullyComposeAffineMapAndOperands(AffineMap *map,
567                                             SmallVectorImpl<Value> *operands) {
568   while (llvm::any_of(*operands, [](Value v) {
569     return isa_and_nonnull<AffineApplyOp>(v.getDefiningOp());
570   })) {
571     composeAffineMapAndOperands(map, operands);
572   }
573 }
574 
575 AffineApplyOp mlir::makeComposedAffineApply(OpBuilder &b, Location loc,
576                                             AffineMap map,
577                                             ArrayRef<Value> operands) {
578   AffineMap normalizedMap = map;
579   SmallVector<Value, 8> normalizedOperands(operands.begin(), operands.end());
580   composeAffineMapAndOperands(&normalizedMap, &normalizedOperands);
581   assert(normalizedMap);
582   return b.create<AffineApplyOp>(loc, normalizedMap, normalizedOperands);
583 }
584 
585 // A symbol may appear as a dim in affine.apply operations. This function
586 // canonicalizes dims that are valid symbols into actual symbols.
587 template <class MapOrSet>
588 static void canonicalizePromotedSymbols(MapOrSet *mapOrSet,
589                                         SmallVectorImpl<Value> *operands) {
590   if (!mapOrSet || operands->empty())
591     return;
592 
593   assert(mapOrSet->getNumInputs() == operands->size() &&
594          "map/set inputs must match number of operands");
595 
596   auto *context = mapOrSet->getContext();
597   SmallVector<Value, 8> resultOperands;
598   resultOperands.reserve(operands->size());
599   SmallVector<Value, 8> remappedSymbols;
600   remappedSymbols.reserve(operands->size());
601   unsigned nextDim = 0;
602   unsigned nextSym = 0;
603   unsigned oldNumSyms = mapOrSet->getNumSymbols();
604   SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims());
605   for (unsigned i = 0, e = mapOrSet->getNumInputs(); i != e; ++i) {
606     if (i < mapOrSet->getNumDims()) {
607       if (isValidSymbol((*operands)[i])) {
608         // This is a valid symbol that appears as a dim, canonicalize it.
609         dimRemapping[i] = getAffineSymbolExpr(oldNumSyms + nextSym++, context);
610         remappedSymbols.push_back((*operands)[i]);
611       } else {
612         dimRemapping[i] = getAffineDimExpr(nextDim++, context);
613         resultOperands.push_back((*operands)[i]);
614       }
615     } else {
616       resultOperands.push_back((*operands)[i]);
617     }
618   }
619 
620   resultOperands.append(remappedSymbols.begin(), remappedSymbols.end());
621   *operands = resultOperands;
622   *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, {}, nextDim,
623                                               oldNumSyms + nextSym);
624 
625   assert(mapOrSet->getNumInputs() == operands->size() &&
626          "map/set inputs must match number of operands");
627 }
628 
629 // Works for either an affine map or an integer set.
630 template <class MapOrSet>
631 static void canonicalizeMapOrSetAndOperands(MapOrSet *mapOrSet,
632                                             SmallVectorImpl<Value> *operands) {
633   static_assert(llvm::is_one_of<MapOrSet, AffineMap, IntegerSet>::value,
634                 "Argument must be either of AffineMap or IntegerSet type");
635 
636   if (!mapOrSet || operands->empty())
637     return;
638 
639   assert(mapOrSet->getNumInputs() == operands->size() &&
640          "map/set inputs must match number of operands");
641 
642   canonicalizePromotedSymbols<MapOrSet>(mapOrSet, operands);
643 
644   // Check to see what dims are used.
645   llvm::SmallBitVector usedDims(mapOrSet->getNumDims());
646   llvm::SmallBitVector usedSyms(mapOrSet->getNumSymbols());
647   mapOrSet->walkExprs([&](AffineExpr expr) {
648     if (auto dimExpr = expr.dyn_cast<AffineDimExpr>())
649       usedDims[dimExpr.getPosition()] = true;
650     else if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>())
651       usedSyms[symExpr.getPosition()] = true;
652   });
653 
654   auto *context = mapOrSet->getContext();
655 
656   SmallVector<Value, 8> resultOperands;
657   resultOperands.reserve(operands->size());
658 
659   llvm::SmallDenseMap<Value, AffineExpr, 8> seenDims;
660   SmallVector<AffineExpr, 8> dimRemapping(mapOrSet->getNumDims());
661   unsigned nextDim = 0;
662   for (unsigned i = 0, e = mapOrSet->getNumDims(); i != e; ++i) {
663     if (usedDims[i]) {
664       // Remap dim positions for duplicate operands.
665       auto it = seenDims.find((*operands)[i]);
666       if (it == seenDims.end()) {
667         dimRemapping[i] = getAffineDimExpr(nextDim++, context);
668         resultOperands.push_back((*operands)[i]);
669         seenDims.insert(std::make_pair((*operands)[i], dimRemapping[i]));
670       } else {
671         dimRemapping[i] = it->second;
672       }
673     }
674   }
675   llvm::SmallDenseMap<Value, AffineExpr, 8> seenSymbols;
676   SmallVector<AffineExpr, 8> symRemapping(mapOrSet->getNumSymbols());
677   unsigned nextSym = 0;
678   for (unsigned i = 0, e = mapOrSet->getNumSymbols(); i != e; ++i) {
679     if (!usedSyms[i])
680       continue;
681     // Handle constant operands (only needed for symbolic operands since
682     // constant operands in dimensional positions would have already been
683     // promoted to symbolic positions above).
684     IntegerAttr operandCst;
685     if (matchPattern((*operands)[i + mapOrSet->getNumDims()],
686                      m_Constant(&operandCst))) {
687       symRemapping[i] =
688           getAffineConstantExpr(operandCst.getValue().getSExtValue(), context);
689       continue;
690     }
691     // Remap symbol positions for duplicate operands.
692     auto it = seenSymbols.find((*operands)[i + mapOrSet->getNumDims()]);
693     if (it == seenSymbols.end()) {
694       symRemapping[i] = getAffineSymbolExpr(nextSym++, context);
695       resultOperands.push_back((*operands)[i + mapOrSet->getNumDims()]);
696       seenSymbols.insert(std::make_pair((*operands)[i + mapOrSet->getNumDims()],
697                                         symRemapping[i]));
698     } else {
699       symRemapping[i] = it->second;
700     }
701   }
702   *mapOrSet = mapOrSet->replaceDimsAndSymbols(dimRemapping, symRemapping,
703                                               nextDim, nextSym);
704   *operands = resultOperands;
705 }
706 
707 void mlir::canonicalizeMapAndOperands(AffineMap *map,
708                                       SmallVectorImpl<Value> *operands) {
709   canonicalizeMapOrSetAndOperands<AffineMap>(map, operands);
710 }
711 
712 void mlir::canonicalizeSetAndOperands(IntegerSet *set,
713                                       SmallVectorImpl<Value> *operands) {
714   canonicalizeMapOrSetAndOperands<IntegerSet>(set, operands);
715 }
716 
717 namespace {
718 /// Simplify AffineApply, AffineLoad, and AffineStore operations by composing
719 /// maps that supply results into them.
720 ///
721 template <typename AffineOpTy>
722 struct SimplifyAffineOp : public OpRewritePattern<AffineOpTy> {
723   using OpRewritePattern<AffineOpTy>::OpRewritePattern;
724 
725   /// Replace the affine op with another instance of it with the supplied
726   /// map and mapOperands.
727   void replaceAffineOp(PatternRewriter &rewriter, AffineOpTy affineOp,
728                        AffineMap map, ArrayRef<Value> mapOperands) const;
729 
730   LogicalResult matchAndRewrite(AffineOpTy affineOp,
731                                 PatternRewriter &rewriter) const override {
732     static_assert(llvm::is_one_of<AffineOpTy, AffineLoadOp, AffinePrefetchOp,
733                                   AffineStoreOp, AffineApplyOp, AffineMinOp,
734                                   AffineMaxOp>::value,
735                   "affine load/store/apply/prefetch/min/max op expected");
736     auto map = affineOp.getAffineMap();
737     AffineMap oldMap = map;
738     auto oldOperands = affineOp.getMapOperands();
739     SmallVector<Value, 8> resultOperands(oldOperands);
740     composeAffineMapAndOperands(&map, &resultOperands);
741     if (map == oldMap && std::equal(oldOperands.begin(), oldOperands.end(),
742                                     resultOperands.begin()))
743       return failure();
744 
745     replaceAffineOp(rewriter, affineOp, map, resultOperands);
746     return success();
747   }
748 };
749 
750 // Specialize the template to account for the different build signatures for
751 // affine load, store, and apply ops.
752 template <>
753 void SimplifyAffineOp<AffineLoadOp>::replaceAffineOp(
754     PatternRewriter &rewriter, AffineLoadOp load, AffineMap map,
755     ArrayRef<Value> mapOperands) const {
756   rewriter.replaceOpWithNewOp<AffineLoadOp>(load, load.getMemRef(), map,
757                                             mapOperands);
758 }
759 template <>
760 void SimplifyAffineOp<AffinePrefetchOp>::replaceAffineOp(
761     PatternRewriter &rewriter, AffinePrefetchOp prefetch, AffineMap map,
762     ArrayRef<Value> mapOperands) const {
763   rewriter.replaceOpWithNewOp<AffinePrefetchOp>(
764       prefetch, prefetch.memref(), map, mapOperands,
765       prefetch.localityHint().getZExtValue(), prefetch.isWrite(),
766       prefetch.isDataCache());
767 }
768 template <>
769 void SimplifyAffineOp<AffineStoreOp>::replaceAffineOp(
770     PatternRewriter &rewriter, AffineStoreOp store, AffineMap map,
771     ArrayRef<Value> mapOperands) const {
772   rewriter.replaceOpWithNewOp<AffineStoreOp>(
773       store, store.getValueToStore(), store.getMemRef(), map, mapOperands);
774 }
775 
776 // Generic version for ops that don't have extra operands.
777 template <typename AffineOpTy>
778 void SimplifyAffineOp<AffineOpTy>::replaceAffineOp(
779     PatternRewriter &rewriter, AffineOpTy op, AffineMap map,
780     ArrayRef<Value> mapOperands) const {
781   rewriter.replaceOpWithNewOp<AffineOpTy>(op, map, mapOperands);
782 }
783 } // end anonymous namespace.
784 
785 void AffineApplyOp::getCanonicalizationPatterns(
786     OwningRewritePatternList &results, MLIRContext *context) {
787   results.insert<SimplifyAffineOp<AffineApplyOp>>(context);
788 }
789 
790 //===----------------------------------------------------------------------===//
791 // Common canonicalization pattern support logic
792 //===----------------------------------------------------------------------===//
793 
794 /// This is a common class used for patterns of the form
795 /// "someop(memrefcast) -> someop".  It folds the source of any memref_cast
796 /// into the root operation directly.
797 static LogicalResult foldMemRefCast(Operation *op) {
798   bool folded = false;
799   for (OpOperand &operand : op->getOpOperands()) {
800     auto cast = dyn_cast_or_null<MemRefCastOp>(operand.get().getDefiningOp());
801     if (cast && !cast.getOperand().getType().isa<UnrankedMemRefType>()) {
802       operand.set(cast.getOperand());
803       folded = true;
804     }
805   }
806   return success(folded);
807 }
808 
809 //===----------------------------------------------------------------------===//
810 // AffineDmaStartOp
811 //===----------------------------------------------------------------------===//
812 
813 // TODO(b/133776335) Check that map operands are loop IVs or symbols.
814 void AffineDmaStartOp::build(Builder *builder, OperationState &result,
815                              Value srcMemRef, AffineMap srcMap,
816                              ValueRange srcIndices, Value destMemRef,
817                              AffineMap dstMap, ValueRange destIndices,
818                              Value tagMemRef, AffineMap tagMap,
819                              ValueRange tagIndices, Value numElements,
820                              Value stride, Value elementsPerStride) {
821   result.addOperands(srcMemRef);
822   result.addAttribute(getSrcMapAttrName(), AffineMapAttr::get(srcMap));
823   result.addOperands(srcIndices);
824   result.addOperands(destMemRef);
825   result.addAttribute(getDstMapAttrName(), AffineMapAttr::get(dstMap));
826   result.addOperands(destIndices);
827   result.addOperands(tagMemRef);
828   result.addAttribute(getTagMapAttrName(), AffineMapAttr::get(tagMap));
829   result.addOperands(tagIndices);
830   result.addOperands(numElements);
831   if (stride) {
832     result.addOperands({stride, elementsPerStride});
833   }
834 }
835 
836 void AffineDmaStartOp::print(OpAsmPrinter &p) {
837   p << "affine.dma_start " << getSrcMemRef() << '[';
838   p.printAffineMapOfSSAIds(getSrcMapAttr(), getSrcIndices());
839   p << "], " << getDstMemRef() << '[';
840   p.printAffineMapOfSSAIds(getDstMapAttr(), getDstIndices());
841   p << "], " << getTagMemRef() << '[';
842   p.printAffineMapOfSSAIds(getTagMapAttr(), getTagIndices());
843   p << "], " << getNumElements();
844   if (isStrided()) {
845     p << ", " << getStride();
846     p << ", " << getNumElementsPerStride();
847   }
848   p << " : " << getSrcMemRefType() << ", " << getDstMemRefType() << ", "
849     << getTagMemRefType();
850 }
851 
852 // Parse AffineDmaStartOp.
853 // Ex:
854 //   affine.dma_start %src[%i, %j], %dst[%k, %l], %tag[%index], %size,
855 //     %stride, %num_elt_per_stride
856 //       : memref<3076 x f32, 0>, memref<1024 x f32, 2>, memref<1 x i32>
857 //
858 ParseResult AffineDmaStartOp::parse(OpAsmParser &parser,
859                                     OperationState &result) {
860   OpAsmParser::OperandType srcMemRefInfo;
861   AffineMapAttr srcMapAttr;
862   SmallVector<OpAsmParser::OperandType, 4> srcMapOperands;
863   OpAsmParser::OperandType dstMemRefInfo;
864   AffineMapAttr dstMapAttr;
865   SmallVector<OpAsmParser::OperandType, 4> dstMapOperands;
866   OpAsmParser::OperandType tagMemRefInfo;
867   AffineMapAttr tagMapAttr;
868   SmallVector<OpAsmParser::OperandType, 4> tagMapOperands;
869   OpAsmParser::OperandType numElementsInfo;
870   SmallVector<OpAsmParser::OperandType, 2> strideInfo;
871 
872   SmallVector<Type, 3> types;
873   auto indexType = parser.getBuilder().getIndexType();
874 
875   // Parse and resolve the following list of operands:
876   // *) dst memref followed by its affine maps operands (in square brackets).
877   // *) src memref followed by its affine map operands (in square brackets).
878   // *) tag memref followed by its affine map operands (in square brackets).
879   // *) number of elements transferred by DMA operation.
880   if (parser.parseOperand(srcMemRefInfo) ||
881       parser.parseAffineMapOfSSAIds(srcMapOperands, srcMapAttr,
882                                     getSrcMapAttrName(), result.attributes) ||
883       parser.parseComma() || parser.parseOperand(dstMemRefInfo) ||
884       parser.parseAffineMapOfSSAIds(dstMapOperands, dstMapAttr,
885                                     getDstMapAttrName(), result.attributes) ||
886       parser.parseComma() || parser.parseOperand(tagMemRefInfo) ||
887       parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr,
888                                     getTagMapAttrName(), result.attributes) ||
889       parser.parseComma() || parser.parseOperand(numElementsInfo))
890     return failure();
891 
892   // Parse optional stride and elements per stride.
893   if (parser.parseTrailingOperandList(strideInfo)) {
894     return failure();
895   }
896   if (!strideInfo.empty() && strideInfo.size() != 2) {
897     return parser.emitError(parser.getNameLoc(),
898                             "expected two stride related operands");
899   }
900   bool isStrided = strideInfo.size() == 2;
901 
902   if (parser.parseColonTypeList(types))
903     return failure();
904 
905   if (types.size() != 3)
906     return parser.emitError(parser.getNameLoc(), "expected three types");
907 
908   if (parser.resolveOperand(srcMemRefInfo, types[0], result.operands) ||
909       parser.resolveOperands(srcMapOperands, indexType, result.operands) ||
910       parser.resolveOperand(dstMemRefInfo, types[1], result.operands) ||
911       parser.resolveOperands(dstMapOperands, indexType, result.operands) ||
912       parser.resolveOperand(tagMemRefInfo, types[2], result.operands) ||
913       parser.resolveOperands(tagMapOperands, indexType, result.operands) ||
914       parser.resolveOperand(numElementsInfo, indexType, result.operands))
915     return failure();
916 
917   if (isStrided) {
918     if (parser.resolveOperands(strideInfo, indexType, result.operands))
919       return failure();
920   }
921 
922   // Check that src/dst/tag operand counts match their map.numInputs.
923   if (srcMapOperands.size() != srcMapAttr.getValue().getNumInputs() ||
924       dstMapOperands.size() != dstMapAttr.getValue().getNumInputs() ||
925       tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
926     return parser.emitError(parser.getNameLoc(),
927                             "memref operand count not equal to map.numInputs");
928   return success();
929 }
930 
931 LogicalResult AffineDmaStartOp::verify() {
932   if (!getOperand(getSrcMemRefOperandIndex()).getType().isa<MemRefType>())
933     return emitOpError("expected DMA source to be of memref type");
934   if (!getOperand(getDstMemRefOperandIndex()).getType().isa<MemRefType>())
935     return emitOpError("expected DMA destination to be of memref type");
936   if (!getOperand(getTagMemRefOperandIndex()).getType().isa<MemRefType>())
937     return emitOpError("expected DMA tag to be of memref type");
938 
939   // DMAs from different memory spaces supported.
940   if (getSrcMemorySpace() == getDstMemorySpace()) {
941     return emitOpError("DMA should be between different memory spaces");
942   }
943   unsigned numInputsAllMaps = getSrcMap().getNumInputs() +
944                               getDstMap().getNumInputs() +
945                               getTagMap().getNumInputs();
946   if (getNumOperands() != numInputsAllMaps + 3 + 1 &&
947       getNumOperands() != numInputsAllMaps + 3 + 1 + 2) {
948     return emitOpError("incorrect number of operands");
949   }
950 
951   for (auto idx : getSrcIndices()) {
952     if (!idx.getType().isIndex())
953       return emitOpError("src index to dma_start must have 'index' type");
954     if (!isValidAffineIndexOperand(idx))
955       return emitOpError("src index must be a dimension or symbol identifier");
956   }
957   for (auto idx : getDstIndices()) {
958     if (!idx.getType().isIndex())
959       return emitOpError("dst index to dma_start must have 'index' type");
960     if (!isValidAffineIndexOperand(idx))
961       return emitOpError("dst index must be a dimension or symbol identifier");
962   }
963   for (auto idx : getTagIndices()) {
964     if (!idx.getType().isIndex())
965       return emitOpError("tag index to dma_start must have 'index' type");
966     if (!isValidAffineIndexOperand(idx))
967       return emitOpError("tag index must be a dimension or symbol identifier");
968   }
969   return success();
970 }
971 
972 LogicalResult AffineDmaStartOp::fold(ArrayRef<Attribute> cstOperands,
973                                      SmallVectorImpl<OpFoldResult> &results) {
974   /// dma_start(memrefcast) -> dma_start
975   return foldMemRefCast(*this);
976 }
977 
978 //===----------------------------------------------------------------------===//
979 // AffineDmaWaitOp
980 //===----------------------------------------------------------------------===//
981 
982 // TODO(b/133776335) Check that map operands are loop IVs or symbols.
983 void AffineDmaWaitOp::build(Builder *builder, OperationState &result,
984                             Value tagMemRef, AffineMap tagMap,
985                             ValueRange tagIndices, Value numElements) {
986   result.addOperands(tagMemRef);
987   result.addAttribute(getTagMapAttrName(), AffineMapAttr::get(tagMap));
988   result.addOperands(tagIndices);
989   result.addOperands(numElements);
990 }
991 
992 void AffineDmaWaitOp::print(OpAsmPrinter &p) {
993   p << "affine.dma_wait " << getTagMemRef() << '[';
994   SmallVector<Value, 2> operands(getTagIndices());
995   p.printAffineMapOfSSAIds(getTagMapAttr(), operands);
996   p << "], ";
997   p.printOperand(getNumElements());
998   p << " : " << getTagMemRef().getType();
999 }
1000 
1001 // Parse AffineDmaWaitOp.
1002 // Eg:
1003 //   affine.dma_wait %tag[%index], %num_elements
1004 //     : memref<1 x i32, (d0) -> (d0), 4>
1005 //
1006 ParseResult AffineDmaWaitOp::parse(OpAsmParser &parser,
1007                                    OperationState &result) {
1008   OpAsmParser::OperandType tagMemRefInfo;
1009   AffineMapAttr tagMapAttr;
1010   SmallVector<OpAsmParser::OperandType, 2> tagMapOperands;
1011   Type type;
1012   auto indexType = parser.getBuilder().getIndexType();
1013   OpAsmParser::OperandType numElementsInfo;
1014 
1015   // Parse tag memref, its map operands, and dma size.
1016   if (parser.parseOperand(tagMemRefInfo) ||
1017       parser.parseAffineMapOfSSAIds(tagMapOperands, tagMapAttr,
1018                                     getTagMapAttrName(), result.attributes) ||
1019       parser.parseComma() || parser.parseOperand(numElementsInfo) ||
1020       parser.parseColonType(type) ||
1021       parser.resolveOperand(tagMemRefInfo, type, result.operands) ||
1022       parser.resolveOperands(tagMapOperands, indexType, result.operands) ||
1023       parser.resolveOperand(numElementsInfo, indexType, result.operands))
1024     return failure();
1025 
1026   if (!type.isa<MemRefType>())
1027     return parser.emitError(parser.getNameLoc(),
1028                             "expected tag to be of memref type");
1029 
1030   if (tagMapOperands.size() != tagMapAttr.getValue().getNumInputs())
1031     return parser.emitError(parser.getNameLoc(),
1032                             "tag memref operand count != to map.numInputs");
1033   return success();
1034 }
1035 
1036 LogicalResult AffineDmaWaitOp::verify() {
1037   if (!getOperand(0).getType().isa<MemRefType>())
1038     return emitOpError("expected DMA tag to be of memref type");
1039   for (auto idx : getTagIndices()) {
1040     if (!idx.getType().isIndex())
1041       return emitOpError("index to dma_wait must have 'index' type");
1042     if (!isValidAffineIndexOperand(idx))
1043       return emitOpError("index must be a dimension or symbol identifier");
1044   }
1045   return success();
1046 }
1047 
1048 LogicalResult AffineDmaWaitOp::fold(ArrayRef<Attribute> cstOperands,
1049                                     SmallVectorImpl<OpFoldResult> &results) {
1050   /// dma_wait(memrefcast) -> dma_wait
1051   return foldMemRefCast(*this);
1052 }
1053 
1054 //===----------------------------------------------------------------------===//
1055 // AffineForOp
1056 //===----------------------------------------------------------------------===//
1057 
1058 void AffineForOp::build(Builder *builder, OperationState &result,
1059                         ValueRange lbOperands, AffineMap lbMap,
1060                         ValueRange ubOperands, AffineMap ubMap, int64_t step) {
1061   assert(((!lbMap && lbOperands.empty()) ||
1062           lbOperands.size() == lbMap.getNumInputs()) &&
1063          "lower bound operand count does not match the affine map");
1064   assert(((!ubMap && ubOperands.empty()) ||
1065           ubOperands.size() == ubMap.getNumInputs()) &&
1066          "upper bound operand count does not match the affine map");
1067   assert(step > 0 && "step has to be a positive integer constant");
1068 
1069   // Add an attribute for the step.
1070   result.addAttribute(getStepAttrName(),
1071                       builder->getIntegerAttr(builder->getIndexType(), step));
1072 
1073   // Add the lower bound.
1074   result.addAttribute(getLowerBoundAttrName(), AffineMapAttr::get(lbMap));
1075   result.addOperands(lbOperands);
1076 
1077   // Add the upper bound.
1078   result.addAttribute(getUpperBoundAttrName(), AffineMapAttr::get(ubMap));
1079   result.addOperands(ubOperands);
1080 
1081   // Create a region and a block for the body.  The argument of the region is
1082   // the loop induction variable.
1083   Region *bodyRegion = result.addRegion();
1084   Block *body = new Block();
1085   body->addArgument(IndexType::get(builder->getContext()));
1086   bodyRegion->push_back(body);
1087   ensureTerminator(*bodyRegion, *builder, result.location);
1088 
1089   // Set the operands list as resizable so that we can freely modify the bounds.
1090   result.setOperandListToResizable();
1091 }
1092 
1093 void AffineForOp::build(Builder *builder, OperationState &result, int64_t lb,
1094                         int64_t ub, int64_t step) {
1095   auto lbMap = AffineMap::getConstantMap(lb, builder->getContext());
1096   auto ubMap = AffineMap::getConstantMap(ub, builder->getContext());
1097   return build(builder, result, {}, lbMap, {}, ubMap, step);
1098 }
1099 
1100 static LogicalResult verify(AffineForOp op) {
1101   // Check that the body defines as single block argument for the induction
1102   // variable.
1103   auto *body = op.getBody();
1104   if (body->getNumArguments() != 1 || !body->getArgument(0).getType().isIndex())
1105     return op.emitOpError(
1106         "expected body to have a single index argument for the "
1107         "induction variable");
1108 
1109   // Verify that there are enough operands for the bounds.
1110   AffineMap lowerBoundMap = op.getLowerBoundMap(),
1111             upperBoundMap = op.getUpperBoundMap();
1112   if (op.getNumOperands() !=
1113       (lowerBoundMap.getNumInputs() + upperBoundMap.getNumInputs()))
1114     return op.emitOpError(
1115         "operand count must match with affine map dimension and symbol count");
1116 
1117   // Verify that the bound operands are valid dimension/symbols.
1118   /// Lower bound.
1119   if (failed(verifyDimAndSymbolIdentifiers(op, op.getLowerBoundOperands(),
1120                                            op.getLowerBoundMap().getNumDims())))
1121     return failure();
1122   /// Upper bound.
1123   if (failed(verifyDimAndSymbolIdentifiers(op, op.getUpperBoundOperands(),
1124                                            op.getUpperBoundMap().getNumDims())))
1125     return failure();
1126   return success();
1127 }
1128 
1129 /// Parse a for operation loop bounds.
1130 static ParseResult parseBound(bool isLower, OperationState &result,
1131                               OpAsmParser &p) {
1132   // 'min' / 'max' prefixes are generally syntactic sugar, but are required if
1133   // the map has multiple results.
1134   bool failedToParsedMinMax =
1135       failed(p.parseOptionalKeyword(isLower ? "max" : "min"));
1136 
1137   auto &builder = p.getBuilder();
1138   auto boundAttrName = isLower ? AffineForOp::getLowerBoundAttrName()
1139                                : AffineForOp::getUpperBoundAttrName();
1140 
1141   // Parse ssa-id as identity map.
1142   SmallVector<OpAsmParser::OperandType, 1> boundOpInfos;
1143   if (p.parseOperandList(boundOpInfos))
1144     return failure();
1145 
1146   if (!boundOpInfos.empty()) {
1147     // Check that only one operand was parsed.
1148     if (boundOpInfos.size() > 1)
1149       return p.emitError(p.getNameLoc(),
1150                          "expected only one loop bound operand");
1151 
1152     // TODO: improve error message when SSA value is not of index type.
1153     // Currently it is 'use of value ... expects different type than prior uses'
1154     if (p.resolveOperand(boundOpInfos.front(), builder.getIndexType(),
1155                          result.operands))
1156       return failure();
1157 
1158     // Create an identity map using symbol id. This representation is optimized
1159     // for storage. Analysis passes may expand it into a multi-dimensional map
1160     // if desired.
1161     AffineMap map = builder.getSymbolIdentityMap();
1162     result.addAttribute(boundAttrName, AffineMapAttr::get(map));
1163     return success();
1164   }
1165 
1166   // Get the attribute location.
1167   llvm::SMLoc attrLoc = p.getCurrentLocation();
1168 
1169   Attribute boundAttr;
1170   if (p.parseAttribute(boundAttr, builder.getIndexType(), boundAttrName,
1171                        result.attributes))
1172     return failure();
1173 
1174   // Parse full form - affine map followed by dim and symbol list.
1175   if (auto affineMapAttr = boundAttr.dyn_cast<AffineMapAttr>()) {
1176     unsigned currentNumOperands = result.operands.size();
1177     unsigned numDims;
1178     if (parseDimAndSymbolList(p, result.operands, numDims))
1179       return failure();
1180 
1181     auto map = affineMapAttr.getValue();
1182     if (map.getNumDims() != numDims)
1183       return p.emitError(
1184           p.getNameLoc(),
1185           "dim operand count and affine map dim count must match");
1186 
1187     unsigned numDimAndSymbolOperands =
1188         result.operands.size() - currentNumOperands;
1189     if (numDims + map.getNumSymbols() != numDimAndSymbolOperands)
1190       return p.emitError(
1191           p.getNameLoc(),
1192           "symbol operand count and affine map symbol count must match");
1193 
1194     // If the map has multiple results, make sure that we parsed the min/max
1195     // prefix.
1196     if (map.getNumResults() > 1 && failedToParsedMinMax) {
1197       if (isLower) {
1198         return p.emitError(attrLoc, "lower loop bound affine map with "
1199                                     "multiple results requires 'max' prefix");
1200       }
1201       return p.emitError(attrLoc, "upper loop bound affine map with multiple "
1202                                   "results requires 'min' prefix");
1203     }
1204     return success();
1205   }
1206 
1207   // Parse custom assembly form.
1208   if (auto integerAttr = boundAttr.dyn_cast<IntegerAttr>()) {
1209     result.attributes.pop_back();
1210     result.addAttribute(
1211         boundAttrName,
1212         AffineMapAttr::get(builder.getConstantAffineMap(integerAttr.getInt())));
1213     return success();
1214   }
1215 
1216   return p.emitError(
1217       p.getNameLoc(),
1218       "expected valid affine map representation for loop bounds");
1219 }
1220 
1221 static ParseResult parseAffineForOp(OpAsmParser &parser,
1222                                     OperationState &result) {
1223   auto &builder = parser.getBuilder();
1224   OpAsmParser::OperandType inductionVariable;
1225   // Parse the induction variable followed by '='.
1226   if (parser.parseRegionArgument(inductionVariable) || parser.parseEqual())
1227     return failure();
1228 
1229   // Parse loop bounds.
1230   if (parseBound(/*isLower=*/true, result, parser) ||
1231       parser.parseKeyword("to", " between bounds") ||
1232       parseBound(/*isLower=*/false, result, parser))
1233     return failure();
1234 
1235   // Parse the optional loop step, we default to 1 if one is not present.
1236   if (parser.parseOptionalKeyword("step")) {
1237     result.addAttribute(
1238         AffineForOp::getStepAttrName(),
1239         builder.getIntegerAttr(builder.getIndexType(), /*value=*/1));
1240   } else {
1241     llvm::SMLoc stepLoc = parser.getCurrentLocation();
1242     IntegerAttr stepAttr;
1243     if (parser.parseAttribute(stepAttr, builder.getIndexType(),
1244                               AffineForOp::getStepAttrName().data(),
1245                               result.attributes))
1246       return failure();
1247 
1248     if (stepAttr.getValue().getSExtValue() < 0)
1249       return parser.emitError(
1250           stepLoc,
1251           "expected step to be representable as a positive signed integer");
1252   }
1253 
1254   // Parse the body region.
1255   Region *body = result.addRegion();
1256   if (parser.parseRegion(*body, inductionVariable, builder.getIndexType()))
1257     return failure();
1258 
1259   AffineForOp::ensureTerminator(*body, builder, result.location);
1260 
1261   // Parse the optional attribute list.
1262   if (parser.parseOptionalAttrDict(result.attributes))
1263     return failure();
1264 
1265   // Set the operands list as resizable so that we can freely modify the bounds.
1266   result.setOperandListToResizable();
1267   return success();
1268 }
1269 
1270 static void printBound(AffineMapAttr boundMap,
1271                        Operation::operand_range boundOperands,
1272                        const char *prefix, OpAsmPrinter &p) {
1273   AffineMap map = boundMap.getValue();
1274 
1275   // Check if this bound should be printed using custom assembly form.
1276   // The decision to restrict printing custom assembly form to trivial cases
1277   // comes from the will to roundtrip MLIR binary -> text -> binary in a
1278   // lossless way.
1279   // Therefore, custom assembly form parsing and printing is only supported for
1280   // zero-operand constant maps and single symbol operand identity maps.
1281   if (map.getNumResults() == 1) {
1282     AffineExpr expr = map.getResult(0);
1283 
1284     // Print constant bound.
1285     if (map.getNumDims() == 0 && map.getNumSymbols() == 0) {
1286       if (auto constExpr = expr.dyn_cast<AffineConstantExpr>()) {
1287         p << constExpr.getValue();
1288         return;
1289       }
1290     }
1291 
1292     // Print bound that consists of a single SSA symbol if the map is over a
1293     // single symbol.
1294     if (map.getNumDims() == 0 && map.getNumSymbols() == 1) {
1295       if (auto symExpr = expr.dyn_cast<AffineSymbolExpr>()) {
1296         p.printOperand(*boundOperands.begin());
1297         return;
1298       }
1299     }
1300   } else {
1301     // Map has multiple results. Print 'min' or 'max' prefix.
1302     p << prefix << ' ';
1303   }
1304 
1305   // Print the map and its operands.
1306   p << boundMap;
1307   printDimAndSymbolList(boundOperands.begin(), boundOperands.end(),
1308                         map.getNumDims(), p);
1309 }
1310 
1311 static void print(OpAsmPrinter &p, AffineForOp op) {
1312   p << op.getOperationName() << ' ';
1313   p.printOperand(op.getBody()->getArgument(0));
1314   p << " = ";
1315   printBound(op.getLowerBoundMapAttr(), op.getLowerBoundOperands(), "max", p);
1316   p << " to ";
1317   printBound(op.getUpperBoundMapAttr(), op.getUpperBoundOperands(), "min", p);
1318 
1319   if (op.getStep() != 1)
1320     p << " step " << op.getStep();
1321   p.printRegion(op.region(),
1322                 /*printEntryBlockArgs=*/false,
1323                 /*printBlockTerminators=*/false);
1324   p.printOptionalAttrDict(op.getAttrs(),
1325                           /*elidedAttrs=*/{op.getLowerBoundAttrName(),
1326                                            op.getUpperBoundAttrName(),
1327                                            op.getStepAttrName()});
1328 }
1329 
1330 /// Fold the constant bounds of a loop.
1331 static LogicalResult foldLoopBounds(AffineForOp forOp) {
1332   auto foldLowerOrUpperBound = [&forOp](bool lower) {
1333     // Check to see if each of the operands is the result of a constant.  If
1334     // so, get the value.  If not, ignore it.
1335     SmallVector<Attribute, 8> operandConstants;
1336     auto boundOperands =
1337         lower ? forOp.getLowerBoundOperands() : forOp.getUpperBoundOperands();
1338     for (auto operand : boundOperands) {
1339       Attribute operandCst;
1340       matchPattern(operand, m_Constant(&operandCst));
1341       operandConstants.push_back(operandCst);
1342     }
1343 
1344     AffineMap boundMap =
1345         lower ? forOp.getLowerBoundMap() : forOp.getUpperBoundMap();
1346     assert(boundMap.getNumResults() >= 1 &&
1347            "bound maps should have at least one result");
1348     SmallVector<Attribute, 4> foldedResults;
1349     if (failed(boundMap.constantFold(operandConstants, foldedResults)))
1350       return failure();
1351 
1352     // Compute the max or min as applicable over the results.
1353     assert(!foldedResults.empty() && "bounds should have at least one result");
1354     auto maxOrMin = foldedResults[0].cast<IntegerAttr>().getValue();
1355     for (unsigned i = 1, e = foldedResults.size(); i < e; i++) {
1356       auto foldedResult = foldedResults[i].cast<IntegerAttr>().getValue();
1357       maxOrMin = lower ? llvm::APIntOps::smax(maxOrMin, foldedResult)
1358                        : llvm::APIntOps::smin(maxOrMin, foldedResult);
1359     }
1360     lower ? forOp.setConstantLowerBound(maxOrMin.getSExtValue())
1361           : forOp.setConstantUpperBound(maxOrMin.getSExtValue());
1362     return success();
1363   };
1364 
1365   // Try to fold the lower bound.
1366   bool folded = false;
1367   if (!forOp.hasConstantLowerBound())
1368     folded |= succeeded(foldLowerOrUpperBound(/*lower=*/true));
1369 
1370   // Try to fold the upper bound.
1371   if (!forOp.hasConstantUpperBound())
1372     folded |= succeeded(foldLowerOrUpperBound(/*lower=*/false));
1373   return success(folded);
1374 }
1375 
1376 /// Canonicalize the bounds of the given loop.
1377 static LogicalResult canonicalizeLoopBounds(AffineForOp forOp) {
1378   SmallVector<Value, 4> lbOperands(forOp.getLowerBoundOperands());
1379   SmallVector<Value, 4> ubOperands(forOp.getUpperBoundOperands());
1380 
1381   auto lbMap = forOp.getLowerBoundMap();
1382   auto ubMap = forOp.getUpperBoundMap();
1383   auto prevLbMap = lbMap;
1384   auto prevUbMap = ubMap;
1385 
1386   canonicalizeMapAndOperands(&lbMap, &lbOperands);
1387   lbMap = removeDuplicateExprs(lbMap);
1388 
1389   canonicalizeMapAndOperands(&ubMap, &ubOperands);
1390   ubMap = removeDuplicateExprs(ubMap);
1391 
1392   // Any canonicalization change always leads to updated map(s).
1393   if (lbMap == prevLbMap && ubMap == prevUbMap)
1394     return failure();
1395 
1396   if (lbMap != prevLbMap)
1397     forOp.setLowerBound(lbOperands, lbMap);
1398   if (ubMap != prevUbMap)
1399     forOp.setUpperBound(ubOperands, ubMap);
1400   return success();
1401 }
1402 
1403 namespace {
1404 /// This is a pattern to fold trivially empty loops.
1405 struct AffineForEmptyLoopFolder : public OpRewritePattern<AffineForOp> {
1406   using OpRewritePattern<AffineForOp>::OpRewritePattern;
1407 
1408   LogicalResult matchAndRewrite(AffineForOp forOp,
1409                                 PatternRewriter &rewriter) const override {
1410     // Check that the body only contains a terminator.
1411     if (!llvm::hasSingleElement(*forOp.getBody()))
1412       return failure();
1413     rewriter.eraseOp(forOp);
1414     return success();
1415   }
1416 };
1417 } // end anonymous namespace
1418 
1419 void AffineForOp::getCanonicalizationPatterns(OwningRewritePatternList &results,
1420                                               MLIRContext *context) {
1421   results.insert<AffineForEmptyLoopFolder>(context);
1422 }
1423 
1424 LogicalResult AffineForOp::fold(ArrayRef<Attribute> operands,
1425                                 SmallVectorImpl<OpFoldResult> &results) {
1426   bool folded = succeeded(foldLoopBounds(*this));
1427   folded |= succeeded(canonicalizeLoopBounds(*this));
1428   return success(folded);
1429 }
1430 
1431 AffineBound AffineForOp::getLowerBound() {
1432   auto lbMap = getLowerBoundMap();
1433   return AffineBound(AffineForOp(*this), 0, lbMap.getNumInputs(), lbMap);
1434 }
1435 
1436 AffineBound AffineForOp::getUpperBound() {
1437   auto lbMap = getLowerBoundMap();
1438   auto ubMap = getUpperBoundMap();
1439   return AffineBound(AffineForOp(*this), lbMap.getNumInputs(), getNumOperands(),
1440                      ubMap);
1441 }
1442 
1443 void AffineForOp::setLowerBound(ValueRange lbOperands, AffineMap map) {
1444   assert(lbOperands.size() == map.getNumInputs());
1445   assert(map.getNumResults() >= 1 && "bound map has at least one result");
1446 
1447   SmallVector<Value, 4> newOperands(lbOperands.begin(), lbOperands.end());
1448 
1449   auto ubOperands = getUpperBoundOperands();
1450   newOperands.append(ubOperands.begin(), ubOperands.end());
1451   getOperation()->setOperands(newOperands);
1452 
1453   setAttr(getLowerBoundAttrName(), AffineMapAttr::get(map));
1454 }
1455 
1456 void AffineForOp::setUpperBound(ValueRange ubOperands, AffineMap map) {
1457   assert(ubOperands.size() == map.getNumInputs());
1458   assert(map.getNumResults() >= 1 && "bound map has at least one result");
1459 
1460   SmallVector<Value, 4> newOperands(getLowerBoundOperands());
1461   newOperands.append(ubOperands.begin(), ubOperands.end());
1462   getOperation()->setOperands(newOperands);
1463 
1464   setAttr(getUpperBoundAttrName(), AffineMapAttr::get(map));
1465 }
1466 
1467 void AffineForOp::setLowerBoundMap(AffineMap map) {
1468   auto lbMap = getLowerBoundMap();
1469   assert(lbMap.getNumDims() == map.getNumDims() &&
1470          lbMap.getNumSymbols() == map.getNumSymbols());
1471   assert(map.getNumResults() >= 1 && "bound map has at least one result");
1472   (void)lbMap;
1473   setAttr(getLowerBoundAttrName(), AffineMapAttr::get(map));
1474 }
1475 
1476 void AffineForOp::setUpperBoundMap(AffineMap map) {
1477   auto ubMap = getUpperBoundMap();
1478   assert(ubMap.getNumDims() == map.getNumDims() &&
1479          ubMap.getNumSymbols() == map.getNumSymbols());
1480   assert(map.getNumResults() >= 1 && "bound map has at least one result");
1481   (void)ubMap;
1482   setAttr(getUpperBoundAttrName(), AffineMapAttr::get(map));
1483 }
1484 
1485 bool AffineForOp::hasConstantLowerBound() {
1486   return getLowerBoundMap().isSingleConstant();
1487 }
1488 
1489 bool AffineForOp::hasConstantUpperBound() {
1490   return getUpperBoundMap().isSingleConstant();
1491 }
1492 
1493 int64_t AffineForOp::getConstantLowerBound() {
1494   return getLowerBoundMap().getSingleConstantResult();
1495 }
1496 
1497 int64_t AffineForOp::getConstantUpperBound() {
1498   return getUpperBoundMap().getSingleConstantResult();
1499 }
1500 
1501 void AffineForOp::setConstantLowerBound(int64_t value) {
1502   setLowerBound({}, AffineMap::getConstantMap(value, getContext()));
1503 }
1504 
1505 void AffineForOp::setConstantUpperBound(int64_t value) {
1506   setUpperBound({}, AffineMap::getConstantMap(value, getContext()));
1507 }
1508 
1509 AffineForOp::operand_range AffineForOp::getLowerBoundOperands() {
1510   return {operand_begin(), operand_begin() + getLowerBoundMap().getNumInputs()};
1511 }
1512 
1513 AffineForOp::operand_range AffineForOp::getUpperBoundOperands() {
1514   return {operand_begin() + getLowerBoundMap().getNumInputs(), operand_end()};
1515 }
1516 
1517 bool AffineForOp::matchingBoundOperandList() {
1518   auto lbMap = getLowerBoundMap();
1519   auto ubMap = getUpperBoundMap();
1520   if (lbMap.getNumDims() != ubMap.getNumDims() ||
1521       lbMap.getNumSymbols() != ubMap.getNumSymbols())
1522     return false;
1523 
1524   unsigned numOperands = lbMap.getNumInputs();
1525   for (unsigned i = 0, e = lbMap.getNumInputs(); i < e; i++) {
1526     // Compare Value 's.
1527     if (getOperand(i) != getOperand(numOperands + i))
1528       return false;
1529   }
1530   return true;
1531 }
1532 
1533 Region &AffineForOp::getLoopBody() { return region(); }
1534 
1535 bool AffineForOp::isDefinedOutsideOfLoop(Value value) {
1536   return !region().isAncestor(value.getParentRegion());
1537 }
1538 
1539 LogicalResult AffineForOp::moveOutOfLoop(ArrayRef<Operation *> ops) {
1540   for (auto *op : ops)
1541     op->moveBefore(*this);
1542   return success();
1543 }
1544 
1545 /// Returns if the provided value is the induction variable of a AffineForOp.
1546 bool mlir::isForInductionVar(Value val) {
1547   return getForInductionVarOwner(val) != AffineForOp();
1548 }
1549 
1550 /// Returns the loop parent of an induction variable. If the provided value is
1551 /// not an induction variable, then return nullptr.
1552 AffineForOp mlir::getForInductionVarOwner(Value val) {
1553   auto ivArg = val.dyn_cast<BlockArgument>();
1554   if (!ivArg || !ivArg.getOwner())
1555     return AffineForOp();
1556   auto *containingInst = ivArg.getOwner()->getParent()->getParentOp();
1557   return dyn_cast<AffineForOp>(containingInst);
1558 }
1559 
1560 /// Extracts the induction variables from a list of AffineForOps and returns
1561 /// them.
1562 void mlir::extractForInductionVars(ArrayRef<AffineForOp> forInsts,
1563                                    SmallVectorImpl<Value> *ivs) {
1564   ivs->reserve(forInsts.size());
1565   for (auto forInst : forInsts)
1566     ivs->push_back(forInst.getInductionVar());
1567 }
1568 
1569 //===----------------------------------------------------------------------===//
1570 // AffineIfOp
1571 //===----------------------------------------------------------------------===//
1572 
1573 namespace {
1574 /// Remove else blocks that have nothing other than the terminator.
1575 struct SimplifyDeadElse : public OpRewritePattern<AffineIfOp> {
1576   using OpRewritePattern<AffineIfOp>::OpRewritePattern;
1577 
1578   LogicalResult matchAndRewrite(AffineIfOp ifOp,
1579                                 PatternRewriter &rewriter) const override {
1580     if (ifOp.elseRegion().empty() ||
1581         !llvm::hasSingleElement(*ifOp.getElseBlock()))
1582       return failure();
1583 
1584     rewriter.startRootUpdate(ifOp);
1585     rewriter.eraseBlock(ifOp.getElseBlock());
1586     rewriter.finalizeRootUpdate(ifOp);
1587     return success();
1588   }
1589 };
1590 } // end anonymous namespace.
1591 
1592 static LogicalResult verify(AffineIfOp op) {
1593   // Verify that we have a condition attribute.
1594   auto conditionAttr =
1595       op.getAttrOfType<IntegerSetAttr>(op.getConditionAttrName());
1596   if (!conditionAttr)
1597     return op.emitOpError(
1598         "requires an integer set attribute named 'condition'");
1599 
1600   // Verify that there are enough operands for the condition.
1601   IntegerSet condition = conditionAttr.getValue();
1602   if (op.getNumOperands() != condition.getNumInputs())
1603     return op.emitOpError(
1604         "operand count and condition integer set dimension and "
1605         "symbol count must match");
1606 
1607   // Verify that the operands are valid dimension/symbols.
1608   if (failed(verifyDimAndSymbolIdentifiers(op, op.getOperands(),
1609                                            condition.getNumDims())))
1610     return failure();
1611 
1612   // Verify that the entry of each child region does not have arguments.
1613   for (auto &region : op.getOperation()->getRegions()) {
1614     for (auto &b : region)
1615       if (b.getNumArguments() != 0)
1616         return op.emitOpError(
1617             "requires that child entry blocks have no arguments");
1618   }
1619   return success();
1620 }
1621 
1622 static ParseResult parseAffineIfOp(OpAsmParser &parser,
1623                                    OperationState &result) {
1624   // Parse the condition attribute set.
1625   IntegerSetAttr conditionAttr;
1626   unsigned numDims;
1627   if (parser.parseAttribute(conditionAttr, AffineIfOp::getConditionAttrName(),
1628                             result.attributes) ||
1629       parseDimAndSymbolList(parser, result.operands, numDims))
1630     return failure();
1631 
1632   // Verify the condition operands.
1633   auto set = conditionAttr.getValue();
1634   if (set.getNumDims() != numDims)
1635     return parser.emitError(
1636         parser.getNameLoc(),
1637         "dim operand count and integer set dim count must match");
1638   if (numDims + set.getNumSymbols() != result.operands.size())
1639     return parser.emitError(
1640         parser.getNameLoc(),
1641         "symbol operand count and integer set symbol count must match");
1642 
1643   // Create the regions for 'then' and 'else'.  The latter must be created even
1644   // if it remains empty for the validity of the operation.
1645   result.regions.reserve(2);
1646   Region *thenRegion = result.addRegion();
1647   Region *elseRegion = result.addRegion();
1648 
1649   // Parse the 'then' region.
1650   if (parser.parseRegion(*thenRegion, {}, {}))
1651     return failure();
1652   AffineIfOp::ensureTerminator(*thenRegion, parser.getBuilder(),
1653                                result.location);
1654 
1655   // If we find an 'else' keyword then parse the 'else' region.
1656   if (!parser.parseOptionalKeyword("else")) {
1657     if (parser.parseRegion(*elseRegion, {}, {}))
1658       return failure();
1659     AffineIfOp::ensureTerminator(*elseRegion, parser.getBuilder(),
1660                                  result.location);
1661   }
1662 
1663   // Parse the optional attribute list.
1664   if (parser.parseOptionalAttrDict(result.attributes))
1665     return failure();
1666 
1667   return success();
1668 }
1669 
1670 static void print(OpAsmPrinter &p, AffineIfOp op) {
1671   auto conditionAttr =
1672       op.getAttrOfType<IntegerSetAttr>(op.getConditionAttrName());
1673   p << "affine.if " << conditionAttr;
1674   printDimAndSymbolList(op.operand_begin(), op.operand_end(),
1675                         conditionAttr.getValue().getNumDims(), p);
1676   p.printRegion(op.thenRegion(),
1677                 /*printEntryBlockArgs=*/false,
1678                 /*printBlockTerminators=*/false);
1679 
1680   // Print the 'else' regions if it has any blocks.
1681   auto &elseRegion = op.elseRegion();
1682   if (!elseRegion.empty()) {
1683     p << " else";
1684     p.printRegion(elseRegion,
1685                   /*printEntryBlockArgs=*/false,
1686                   /*printBlockTerminators=*/false);
1687   }
1688 
1689   // Print the attribute list.
1690   p.printOptionalAttrDict(op.getAttrs(),
1691                           /*elidedAttrs=*/op.getConditionAttrName());
1692 }
1693 
1694 IntegerSet AffineIfOp::getIntegerSet() {
1695   return getAttrOfType<IntegerSetAttr>(getConditionAttrName()).getValue();
1696 }
1697 void AffineIfOp::setIntegerSet(IntegerSet newSet) {
1698   setAttr(getConditionAttrName(), IntegerSetAttr::get(newSet));
1699 }
1700 
1701 void AffineIfOp::setConditional(IntegerSet set, ValueRange operands) {
1702   setIntegerSet(set);
1703   getOperation()->setOperands(operands);
1704 }
1705 
1706 void AffineIfOp::build(Builder *builder, OperationState &result, IntegerSet set,
1707                        ValueRange args, bool withElseRegion) {
1708   result.addOperands(args);
1709   result.addAttribute(getConditionAttrName(), IntegerSetAttr::get(set));
1710   Region *thenRegion = result.addRegion();
1711   Region *elseRegion = result.addRegion();
1712   AffineIfOp::ensureTerminator(*thenRegion, *builder, result.location);
1713   if (withElseRegion)
1714     AffineIfOp::ensureTerminator(*elseRegion, *builder, result.location);
1715 }
1716 
1717 /// Canonicalize an affine if op's conditional (integer set + operands).
1718 LogicalResult AffineIfOp::fold(ArrayRef<Attribute>,
1719                                SmallVectorImpl<OpFoldResult> &) {
1720   auto set = getIntegerSet();
1721   SmallVector<Value, 4> operands(getOperands());
1722   canonicalizeSetAndOperands(&set, &operands);
1723 
1724   // Any canonicalization change always leads to either a reduction in the
1725   // number of operands or a change in the number of symbolic operands
1726   // (promotion of dims to symbols).
1727   if (operands.size() < getIntegerSet().getNumInputs() ||
1728       set.getNumSymbols() > getIntegerSet().getNumSymbols()) {
1729     setConditional(set, operands);
1730     return success();
1731   }
1732 
1733   return failure();
1734 }
1735 
1736 void AffineIfOp::getCanonicalizationPatterns(OwningRewritePatternList &results,
1737                                              MLIRContext *context) {
1738   results.insert<SimplifyDeadElse>(context);
1739 }
1740 
1741 //===----------------------------------------------------------------------===//
1742 // AffineLoadOp
1743 //===----------------------------------------------------------------------===//
1744 
1745 void AffineLoadOp::build(Builder *builder, OperationState &result,
1746                          AffineMap map, ValueRange operands) {
1747   assert(operands.size() == 1 + map.getNumInputs() && "inconsistent operands");
1748   result.addOperands(operands);
1749   if (map)
1750     result.addAttribute(getMapAttrName(), AffineMapAttr::get(map));
1751   auto memrefType = operands[0].getType().cast<MemRefType>();
1752   result.types.push_back(memrefType.getElementType());
1753 }
1754 
1755 void AffineLoadOp::build(Builder *builder, OperationState &result, Value memref,
1756                          AffineMap map, ValueRange mapOperands) {
1757   assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
1758   result.addOperands(memref);
1759   result.addOperands(mapOperands);
1760   auto memrefType = memref.getType().cast<MemRefType>();
1761   result.addAttribute(getMapAttrName(), AffineMapAttr::get(map));
1762   result.types.push_back(memrefType.getElementType());
1763 }
1764 
1765 void AffineLoadOp::build(Builder *builder, OperationState &result, Value memref,
1766                          ValueRange indices) {
1767   auto memrefType = memref.getType().cast<MemRefType>();
1768   auto rank = memrefType.getRank();
1769   // Create identity map for memrefs with at least one dimension or () -> ()
1770   // for zero-dimensional memrefs.
1771   auto map = rank ? builder->getMultiDimIdentityMap(rank)
1772                   : builder->getEmptyAffineMap();
1773   build(builder, result, memref, map, indices);
1774 }
1775 
1776 ParseResult AffineLoadOp::parse(OpAsmParser &parser, OperationState &result) {
1777   auto &builder = parser.getBuilder();
1778   auto indexTy = builder.getIndexType();
1779 
1780   MemRefType type;
1781   OpAsmParser::OperandType memrefInfo;
1782   AffineMapAttr mapAttr;
1783   SmallVector<OpAsmParser::OperandType, 1> mapOperands;
1784   return failure(
1785       parser.parseOperand(memrefInfo) ||
1786       parser.parseAffineMapOfSSAIds(mapOperands, mapAttr, getMapAttrName(),
1787                                     result.attributes) ||
1788       parser.parseOptionalAttrDict(result.attributes) ||
1789       parser.parseColonType(type) ||
1790       parser.resolveOperand(memrefInfo, type, result.operands) ||
1791       parser.resolveOperands(mapOperands, indexTy, result.operands) ||
1792       parser.addTypeToList(type.getElementType(), result.types));
1793 }
1794 
1795 void AffineLoadOp::print(OpAsmPrinter &p) {
1796   p << "affine.load " << getMemRef() << '[';
1797   if (AffineMapAttr mapAttr = getAttrOfType<AffineMapAttr>(getMapAttrName()))
1798     p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
1799   p << ']';
1800   p.printOptionalAttrDict(getAttrs(), /*elidedAttrs=*/{getMapAttrName()});
1801   p << " : " << getMemRefType();
1802 }
1803 
1804 LogicalResult AffineLoadOp::verify() {
1805   if (getType() != getMemRefType().getElementType())
1806     return emitOpError("result type must match element type of memref");
1807 
1808   auto mapAttr = getAttrOfType<AffineMapAttr>(getMapAttrName());
1809   if (mapAttr) {
1810     AffineMap map = getAttrOfType<AffineMapAttr>(getMapAttrName()).getValue();
1811     if (map.getNumResults() != getMemRefType().getRank())
1812       return emitOpError("affine.load affine map num results must equal"
1813                          " memref rank");
1814     if (map.getNumInputs() != getNumOperands() - 1)
1815       return emitOpError("expects as many subscripts as affine map inputs");
1816   } else {
1817     if (getMemRefType().getRank() != getNumOperands() - 1)
1818       return emitOpError(
1819           "expects the number of subscripts to be equal to memref rank");
1820   }
1821 
1822   for (auto idx : getMapOperands()) {
1823     if (!idx.getType().isIndex())
1824       return emitOpError("index to load must have 'index' type");
1825     if (!isValidAffineIndexOperand(idx))
1826       return emitOpError("index must be a dimension or symbol identifier");
1827   }
1828   return success();
1829 }
1830 
1831 void AffineLoadOp::getCanonicalizationPatterns(
1832     OwningRewritePatternList &results, MLIRContext *context) {
1833   results.insert<SimplifyAffineOp<AffineLoadOp>>(context);
1834 }
1835 
1836 OpFoldResult AffineLoadOp::fold(ArrayRef<Attribute> cstOperands) {
1837   /// load(memrefcast) -> load
1838   if (succeeded(foldMemRefCast(*this)))
1839     return getResult();
1840   return OpFoldResult();
1841 }
1842 
1843 //===----------------------------------------------------------------------===//
1844 // AffineStoreOp
1845 //===----------------------------------------------------------------------===//
1846 
1847 void AffineStoreOp::build(Builder *builder, OperationState &result,
1848                           Value valueToStore, Value memref, AffineMap map,
1849                           ValueRange mapOperands) {
1850   assert(map.getNumInputs() == mapOperands.size() && "inconsistent index info");
1851   result.addOperands(valueToStore);
1852   result.addOperands(memref);
1853   result.addOperands(mapOperands);
1854   result.addAttribute(getMapAttrName(), AffineMapAttr::get(map));
1855 }
1856 
1857 // Use identity map.
1858 void AffineStoreOp::build(Builder *builder, OperationState &result,
1859                           Value valueToStore, Value memref,
1860                           ValueRange indices) {
1861   auto memrefType = memref.getType().cast<MemRefType>();
1862   auto rank = memrefType.getRank();
1863   // Create identity map for memrefs with at least one dimension or () -> ()
1864   // for zero-dimensional memrefs.
1865   auto map = rank ? builder->getMultiDimIdentityMap(rank)
1866                   : builder->getEmptyAffineMap();
1867   build(builder, result, valueToStore, memref, map, indices);
1868 }
1869 
1870 ParseResult AffineStoreOp::parse(OpAsmParser &parser, OperationState &result) {
1871   auto indexTy = parser.getBuilder().getIndexType();
1872 
1873   MemRefType type;
1874   OpAsmParser::OperandType storeValueInfo;
1875   OpAsmParser::OperandType memrefInfo;
1876   AffineMapAttr mapAttr;
1877   SmallVector<OpAsmParser::OperandType, 1> mapOperands;
1878   return failure(parser.parseOperand(storeValueInfo) || parser.parseComma() ||
1879                  parser.parseOperand(memrefInfo) ||
1880                  parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
1881                                                getMapAttrName(),
1882                                                result.attributes) ||
1883                  parser.parseOptionalAttrDict(result.attributes) ||
1884                  parser.parseColonType(type) ||
1885                  parser.resolveOperand(storeValueInfo, type.getElementType(),
1886                                        result.operands) ||
1887                  parser.resolveOperand(memrefInfo, type, result.operands) ||
1888                  parser.resolveOperands(mapOperands, indexTy, result.operands));
1889 }
1890 
1891 void AffineStoreOp::print(OpAsmPrinter &p) {
1892   p << "affine.store " << getValueToStore();
1893   p << ", " << getMemRef() << '[';
1894   if (AffineMapAttr mapAttr = getAttrOfType<AffineMapAttr>(getMapAttrName()))
1895     p.printAffineMapOfSSAIds(mapAttr, getMapOperands());
1896   p << ']';
1897   p.printOptionalAttrDict(getAttrs(), /*elidedAttrs=*/{getMapAttrName()});
1898   p << " : " << getMemRefType();
1899 }
1900 
1901 LogicalResult AffineStoreOp::verify() {
1902   // First operand must have same type as memref element type.
1903   if (getValueToStore().getType() != getMemRefType().getElementType())
1904     return emitOpError("first operand must have same type memref element type");
1905 
1906   auto mapAttr = getAttrOfType<AffineMapAttr>(getMapAttrName());
1907   if (mapAttr) {
1908     AffineMap map = mapAttr.getValue();
1909     if (map.getNumResults() != getMemRefType().getRank())
1910       return emitOpError("affine.store affine map num results must equal"
1911                          " memref rank");
1912     if (map.getNumInputs() != getNumOperands() - 2)
1913       return emitOpError("expects as many subscripts as affine map inputs");
1914   } else {
1915     if (getMemRefType().getRank() != getNumOperands() - 2)
1916       return emitOpError(
1917           "expects the number of subscripts to be equal to memref rank");
1918   }
1919 
1920   for (auto idx : getMapOperands()) {
1921     if (!idx.getType().isIndex())
1922       return emitOpError("index to store must have 'index' type");
1923     if (!isValidAffineIndexOperand(idx))
1924       return emitOpError("index must be a dimension or symbol identifier");
1925   }
1926   return success();
1927 }
1928 
1929 void AffineStoreOp::getCanonicalizationPatterns(
1930     OwningRewritePatternList &results, MLIRContext *context) {
1931   results.insert<SimplifyAffineOp<AffineStoreOp>>(context);
1932 }
1933 
1934 LogicalResult AffineStoreOp::fold(ArrayRef<Attribute> cstOperands,
1935                                   SmallVectorImpl<OpFoldResult> &results) {
1936   /// store(memrefcast) -> store
1937   return foldMemRefCast(*this);
1938 }
1939 
1940 //===----------------------------------------------------------------------===//
1941 // AffineMinMaxOpBase
1942 //===----------------------------------------------------------------------===//
1943 
1944 template <typename T>
1945 static LogicalResult verifyAffineMinMaxOp(T op) {
1946   // Verify that operand count matches affine map dimension and symbol count.
1947   if (op.getNumOperands() != op.map().getNumDims() + op.map().getNumSymbols())
1948     return op.emitOpError(
1949         "operand count and affine map dimension and symbol count must match");
1950   return success();
1951 }
1952 
1953 template <typename T>
1954 static void printAffineMinMaxOp(OpAsmPrinter &p, T op) {
1955   p << op.getOperationName() << ' ' << op.getAttr(T::getMapAttrName());
1956   auto operands = op.getOperands();
1957   unsigned numDims = op.map().getNumDims();
1958   p << '(' << operands.take_front(numDims) << ')';
1959 
1960   if (operands.size() != numDims)
1961     p << '[' << operands.drop_front(numDims) << ']';
1962   p.printOptionalAttrDict(op.getAttrs(),
1963                           /*elidedAttrs=*/{T::getMapAttrName()});
1964 }
1965 
1966 template <typename T>
1967 static ParseResult parseAffineMinMaxOp(OpAsmParser &parser,
1968                                        OperationState &result) {
1969   auto &builder = parser.getBuilder();
1970   auto indexType = builder.getIndexType();
1971   SmallVector<OpAsmParser::OperandType, 8> dim_infos;
1972   SmallVector<OpAsmParser::OperandType, 8> sym_infos;
1973   AffineMapAttr mapAttr;
1974   return failure(
1975       parser.parseAttribute(mapAttr, T::getMapAttrName(), result.attributes) ||
1976       parser.parseOperandList(dim_infos, OpAsmParser::Delimiter::Paren) ||
1977       parser.parseOperandList(sym_infos,
1978                               OpAsmParser::Delimiter::OptionalSquare) ||
1979       parser.parseOptionalAttrDict(result.attributes) ||
1980       parser.resolveOperands(dim_infos, indexType, result.operands) ||
1981       parser.resolveOperands(sym_infos, indexType, result.operands) ||
1982       parser.addTypeToList(indexType, result.types));
1983 }
1984 
1985 //===----------------------------------------------------------------------===//
1986 // AffineMinOp
1987 //===----------------------------------------------------------------------===//
1988 //
1989 //   %0 = affine.min (d0) -> (1000, d0 + 512) (%i0)
1990 //
1991 
1992 OpFoldResult AffineMinOp::fold(ArrayRef<Attribute> operands) {
1993   // Fold the affine map.
1994   // TODO(andydavis, ntv) Fold more cases: partial static information,
1995   // min(some_affine, some_affine + constant, ...).
1996   SmallVector<Attribute, 2> results;
1997   if (failed(map().constantFold(operands, results)))
1998     return {};
1999 
2000   // Compute and return min of folded map results.
2001   int64_t min = std::numeric_limits<int64_t>::max();
2002   int minIndex = -1;
2003   for (unsigned i = 0, e = results.size(); i < e; ++i) {
2004     auto intAttr = results[i].cast<IntegerAttr>();
2005     if (intAttr.getInt() < min) {
2006       min = intAttr.getInt();
2007       minIndex = i;
2008     }
2009   }
2010   if (minIndex < 0)
2011     return {};
2012   return results[minIndex];
2013 }
2014 
2015 void AffineMinOp::getCanonicalizationPatterns(
2016     OwningRewritePatternList &patterns, MLIRContext *context) {
2017   patterns.insert<SimplifyAffineOp<AffineMinOp>>(context);
2018 }
2019 
2020 //===----------------------------------------------------------------------===//
2021 // AffineMaxOp
2022 //===----------------------------------------------------------------------===//
2023 //
2024 //   %0 = affine.max (d0) -> (1000, d0 + 512) (%i0)
2025 //
2026 
2027 OpFoldResult AffineMaxOp::fold(ArrayRef<Attribute> operands) {
2028   // Fold the affine map.
2029   // TODO(andydavis, ntv, ouhang) Fold more cases: partial static information,
2030   // max(some_affine, some_affine + constant, ...).
2031   SmallVector<Attribute, 2> results;
2032   if (failed(map().constantFold(operands, results)))
2033     return {};
2034 
2035   // Compute and return max of folded map results.
2036   int64_t max = std::numeric_limits<int64_t>::min();
2037   int maxIndex = -1;
2038   for (unsigned i = 0, e = results.size(); i < e; ++i) {
2039     auto intAttr = results[i].cast<IntegerAttr>();
2040     if (intAttr.getInt() > max) {
2041       max = intAttr.getInt();
2042       maxIndex = i;
2043     }
2044   }
2045   if (maxIndex < 0)
2046     return {};
2047   return results[maxIndex];
2048 }
2049 
2050 void AffineMaxOp::getCanonicalizationPatterns(
2051     OwningRewritePatternList &patterns, MLIRContext *context) {
2052   patterns.insert<SimplifyAffineOp<AffineMaxOp>>(context);
2053 }
2054 
2055 //===----------------------------------------------------------------------===//
2056 // AffinePrefetchOp
2057 //===----------------------------------------------------------------------===//
2058 
2059 //
2060 // affine.prefetch %0[%i, %j + 5], read, locality<3>, data : memref<400x400xi32>
2061 //
2062 static ParseResult parseAffinePrefetchOp(OpAsmParser &parser,
2063                                          OperationState &result) {
2064   auto &builder = parser.getBuilder();
2065   auto indexTy = builder.getIndexType();
2066 
2067   MemRefType type;
2068   OpAsmParser::OperandType memrefInfo;
2069   IntegerAttr hintInfo;
2070   auto i32Type = parser.getBuilder().getIntegerType(32);
2071   StringRef readOrWrite, cacheType;
2072 
2073   AffineMapAttr mapAttr;
2074   SmallVector<OpAsmParser::OperandType, 1> mapOperands;
2075   if (parser.parseOperand(memrefInfo) ||
2076       parser.parseAffineMapOfSSAIds(mapOperands, mapAttr,
2077                                     AffinePrefetchOp::getMapAttrName(),
2078                                     result.attributes) ||
2079       parser.parseComma() || parser.parseKeyword(&readOrWrite) ||
2080       parser.parseComma() || parser.parseKeyword("locality") ||
2081       parser.parseLess() ||
2082       parser.parseAttribute(hintInfo, i32Type,
2083                             AffinePrefetchOp::getLocalityHintAttrName(),
2084                             result.attributes) ||
2085       parser.parseGreater() || parser.parseComma() ||
2086       parser.parseKeyword(&cacheType) ||
2087       parser.parseOptionalAttrDict(result.attributes) ||
2088       parser.parseColonType(type) ||
2089       parser.resolveOperand(memrefInfo, type, result.operands) ||
2090       parser.resolveOperands(mapOperands, indexTy, result.operands))
2091     return failure();
2092 
2093   if (!readOrWrite.equals("read") && !readOrWrite.equals("write"))
2094     return parser.emitError(parser.getNameLoc(),
2095                             "rw specifier has to be 'read' or 'write'");
2096   result.addAttribute(
2097       AffinePrefetchOp::getIsWriteAttrName(),
2098       parser.getBuilder().getBoolAttr(readOrWrite.equals("write")));
2099 
2100   if (!cacheType.equals("data") && !cacheType.equals("instr"))
2101     return parser.emitError(parser.getNameLoc(),
2102                             "cache type has to be 'data' or 'instr'");
2103 
2104   result.addAttribute(
2105       AffinePrefetchOp::getIsDataCacheAttrName(),
2106       parser.getBuilder().getBoolAttr(cacheType.equals("data")));
2107 
2108   return success();
2109 }
2110 
2111 static void print(OpAsmPrinter &p, AffinePrefetchOp op) {
2112   p << AffinePrefetchOp::getOperationName() << " " << op.memref() << '[';
2113   AffineMapAttr mapAttr = op.getAttrOfType<AffineMapAttr>(op.getMapAttrName());
2114   if (mapAttr) {
2115     SmallVector<Value, 2> operands(op.getMapOperands());
2116     p.printAffineMapOfSSAIds(mapAttr, operands);
2117   }
2118   p << ']' << ", " << (op.isWrite() ? "write" : "read") << ", "
2119     << "locality<" << op.localityHint() << ">, "
2120     << (op.isDataCache() ? "data" : "instr");
2121   p.printOptionalAttrDict(
2122       op.getAttrs(),
2123       /*elidedAttrs=*/{op.getMapAttrName(), op.getLocalityHintAttrName(),
2124                        op.getIsDataCacheAttrName(), op.getIsWriteAttrName()});
2125   p << " : " << op.getMemRefType();
2126 }
2127 
2128 static LogicalResult verify(AffinePrefetchOp op) {
2129   auto mapAttr = op.getAttrOfType<AffineMapAttr>(op.getMapAttrName());
2130   if (mapAttr) {
2131     AffineMap map = mapAttr.getValue();
2132     if (map.getNumResults() != op.getMemRefType().getRank())
2133       return op.emitOpError("affine.prefetch affine map num results must equal"
2134                             " memref rank");
2135     if (map.getNumInputs() + 1 != op.getNumOperands())
2136       return op.emitOpError("too few operands");
2137   } else {
2138     if (op.getNumOperands() != 1)
2139       return op.emitOpError("too few operands");
2140   }
2141 
2142   for (auto idx : op.getMapOperands()) {
2143     if (!isValidAffineIndexOperand(idx))
2144       return op.emitOpError("index must be a dimension or symbol identifier");
2145   }
2146   return success();
2147 }
2148 
2149 void AffinePrefetchOp::getCanonicalizationPatterns(
2150     OwningRewritePatternList &results, MLIRContext *context) {
2151   // prefetch(memrefcast) -> prefetch
2152   results.insert<SimplifyAffineOp<AffinePrefetchOp>>(context);
2153 }
2154 
2155 LogicalResult AffinePrefetchOp::fold(ArrayRef<Attribute> cstOperands,
2156                                      SmallVectorImpl<OpFoldResult> &results) {
2157   /// prefetch(memrefcast) -> prefetch
2158   return foldMemRefCast(*this);
2159 }
2160 
2161 //===----------------------------------------------------------------------===//
2162 // AffineParallelOp
2163 //===----------------------------------------------------------------------===//
2164 
2165 void AffineParallelOp::build(Builder *builder, OperationState &result,
2166                              ArrayRef<int64_t> ranges) {
2167   SmallVector<AffineExpr, 8> lbExprs(ranges.size(),
2168                                      builder->getAffineConstantExpr(0));
2169   auto lbMap = AffineMap::get(0, 0, lbExprs, builder->getContext());
2170   SmallVector<AffineExpr, 8> ubExprs;
2171   for (int64_t range : ranges)
2172     ubExprs.push_back(builder->getAffineConstantExpr(range));
2173   auto ubMap = AffineMap::get(0, 0, ubExprs, builder->getContext());
2174   build(builder, result, lbMap, {}, ubMap, {});
2175 }
2176 
2177 void AffineParallelOp::build(Builder *builder, OperationState &result,
2178                              AffineMap lbMap, ValueRange lbArgs,
2179                              AffineMap ubMap, ValueRange ubArgs) {
2180   auto numDims = lbMap.getNumResults();
2181   // Verify that the dimensionality of both maps are the same.
2182   assert(numDims == ubMap.getNumResults() &&
2183          "num dims and num results mismatch");
2184   // Make default step sizes of 1.
2185   SmallVector<int64_t, 8> steps(numDims, 1);
2186   build(builder, result, lbMap, lbArgs, ubMap, ubArgs, steps);
2187 }
2188 
2189 void AffineParallelOp::build(Builder *builder, OperationState &result,
2190                              AffineMap lbMap, ValueRange lbArgs,
2191                              AffineMap ubMap, ValueRange ubArgs,
2192                              ArrayRef<int64_t> steps) {
2193   auto numDims = lbMap.getNumResults();
2194   // Verify that the dimensionality of the maps matches the number of steps.
2195   assert(numDims == ubMap.getNumResults() &&
2196          "num dims and num results mismatch");
2197   assert(numDims == steps.size() && "num dims and num steps mismatch");
2198   result.addAttribute(getLowerBoundsMapAttrName(), AffineMapAttr::get(lbMap));
2199   result.addAttribute(getUpperBoundsMapAttrName(), AffineMapAttr::get(ubMap));
2200   result.addAttribute(getStepsAttrName(), builder->getI64ArrayAttr(steps));
2201   result.addOperands(lbArgs);
2202   result.addOperands(ubArgs);
2203   // Create a region and a block for the body.
2204   auto bodyRegion = result.addRegion();
2205   auto body = new Block();
2206   // Add all the block arguments.
2207   for (unsigned i = 0; i < numDims; ++i)
2208     body->addArgument(IndexType::get(builder->getContext()));
2209   bodyRegion->push_back(body);
2210   ensureTerminator(*bodyRegion, *builder, result.location);
2211 }
2212 
2213 unsigned AffineParallelOp::getNumDims() { return steps().size(); }
2214 
2215 AffineParallelOp::operand_range AffineParallelOp::getLowerBoundsOperands() {
2216   return getOperands().take_front(lowerBoundsMap().getNumInputs());
2217 }
2218 
2219 AffineParallelOp::operand_range AffineParallelOp::getUpperBoundsOperands() {
2220   return getOperands().drop_front(lowerBoundsMap().getNumInputs());
2221 }
2222 
2223 AffineValueMap AffineParallelOp::getLowerBoundsValueMap() {
2224   return AffineValueMap(lowerBoundsMap(), getLowerBoundsOperands());
2225 }
2226 
2227 AffineValueMap AffineParallelOp::getUpperBoundsValueMap() {
2228   return AffineValueMap(upperBoundsMap(), getUpperBoundsOperands());
2229 }
2230 
2231 AffineValueMap AffineParallelOp::getRangesValueMap() {
2232   AffineValueMap out;
2233   AffineValueMap::difference(getUpperBoundsValueMap(), getLowerBoundsValueMap(),
2234                              &out);
2235   return out;
2236 }
2237 
2238 Optional<SmallVector<int64_t, 8>> AffineParallelOp::getConstantRanges() {
2239   // Try to convert all the ranges to constant expressions.
2240   SmallVector<int64_t, 8> out;
2241   AffineValueMap rangesValueMap = getRangesValueMap();
2242   out.reserve(rangesValueMap.getNumResults());
2243   for (unsigned i = 0, e = rangesValueMap.getNumResults(); i < e; ++i) {
2244     auto expr = rangesValueMap.getResult(i);
2245     auto cst = expr.dyn_cast<AffineConstantExpr>();
2246     if (!cst)
2247       return llvm::None;
2248     out.push_back(cst.getValue());
2249   }
2250   return out;
2251 }
2252 
2253 Block *AffineParallelOp::getBody() { return &region().front(); }
2254 
2255 OpBuilder AffineParallelOp::getBodyBuilder() {
2256   return OpBuilder(getBody(), std::prev(getBody()->end()));
2257 }
2258 
2259 void AffineParallelOp::setSteps(ArrayRef<int64_t> newSteps) {
2260   assert(newSteps.size() == getNumDims() && "steps & num dims mismatch");
2261   setAttr(getStepsAttrName(), getBodyBuilder().getI64ArrayAttr(newSteps));
2262 }
2263 
2264 static LogicalResult verify(AffineParallelOp op) {
2265   auto numDims = op.getNumDims();
2266   if (op.lowerBoundsMap().getNumResults() != numDims ||
2267       op.upperBoundsMap().getNumResults() != numDims ||
2268       op.steps().size() != numDims ||
2269       op.getBody()->getNumArguments() != numDims) {
2270     return op.emitOpError("region argument count and num results of upper "
2271                           "bounds, lower bounds, and steps must all match");
2272   }
2273   // Verify that the bound operands are valid dimension/symbols.
2274   /// Lower bounds.
2275   if (failed(verifyDimAndSymbolIdentifiers(op, op.getLowerBoundsOperands(),
2276                                            op.lowerBoundsMap().getNumDims())))
2277     return failure();
2278   /// Upper bounds.
2279   if (failed(verifyDimAndSymbolIdentifiers(op, op.getUpperBoundsOperands(),
2280                                            op.upperBoundsMap().getNumDims())))
2281     return failure();
2282   return success();
2283 }
2284 
2285 static void print(OpAsmPrinter &p, AffineParallelOp op) {
2286   p << op.getOperationName() << " (" << op.getBody()->getArguments() << ") = (";
2287   p.printAffineMapOfSSAIds(op.lowerBoundsMapAttr(),
2288                            op.getLowerBoundsOperands());
2289   p << ") to (";
2290   p.printAffineMapOfSSAIds(op.upperBoundsMapAttr(),
2291                            op.getUpperBoundsOperands());
2292   p << ')';
2293   SmallVector<int64_t, 4> steps;
2294   bool elideSteps = true;
2295   for (auto attr : op.steps()) {
2296     auto step = attr.cast<IntegerAttr>().getInt();
2297     elideSteps &= (step == 1);
2298     steps.push_back(step);
2299   }
2300   if (!elideSteps) {
2301     p << " step (";
2302     llvm::interleaveComma(steps, p);
2303     p << ')';
2304   }
2305   p.printRegion(op.region(), /*printEntryBlockArgs=*/false,
2306                 /*printBlockTerminators=*/false);
2307   p.printOptionalAttrDict(
2308       op.getAttrs(),
2309       /*elidedAttrs=*/{AffineParallelOp::getLowerBoundsMapAttrName(),
2310                        AffineParallelOp::getUpperBoundsMapAttrName(),
2311                        AffineParallelOp::getStepsAttrName()});
2312 }
2313 
2314 //
2315 // operation ::= `affine.parallel` `(` ssa-ids `)` `=` `(` map-of-ssa-ids `)`
2316 //               `to` `(` map-of-ssa-ids `)` steps? region attr-dict?
2317 // steps     ::= `steps` `(` integer-literals `)`
2318 //
2319 static ParseResult parseAffineParallelOp(OpAsmParser &parser,
2320                                          OperationState &result) {
2321   auto &builder = parser.getBuilder();
2322   auto indexType = builder.getIndexType();
2323   AffineMapAttr lowerBoundsAttr, upperBoundsAttr;
2324   SmallVector<OpAsmParser::OperandType, 4> ivs;
2325   SmallVector<OpAsmParser::OperandType, 4> lowerBoundsMapOperands;
2326   SmallVector<OpAsmParser::OperandType, 4> upperBoundsMapOperands;
2327   if (parser.parseRegionArgumentList(ivs, /*requiredOperandCount=*/-1,
2328                                      OpAsmParser::Delimiter::Paren) ||
2329       parser.parseEqual() ||
2330       parser.parseAffineMapOfSSAIds(
2331           lowerBoundsMapOperands, lowerBoundsAttr,
2332           AffineParallelOp::getLowerBoundsMapAttrName(), result.attributes,
2333           OpAsmParser::Delimiter::Paren) ||
2334       parser.resolveOperands(lowerBoundsMapOperands, indexType,
2335                              result.operands) ||
2336       parser.parseKeyword("to") ||
2337       parser.parseAffineMapOfSSAIds(
2338           upperBoundsMapOperands, upperBoundsAttr,
2339           AffineParallelOp::getUpperBoundsMapAttrName(), result.attributes,
2340           OpAsmParser::Delimiter::Paren) ||
2341       parser.resolveOperands(upperBoundsMapOperands, indexType,
2342                              result.operands))
2343     return failure();
2344 
2345   AffineMapAttr stepsMapAttr;
2346   SmallVector<NamedAttribute, 1> stepsAttrs;
2347   SmallVector<OpAsmParser::OperandType, 4> stepsMapOperands;
2348   if (failed(parser.parseOptionalKeyword("step"))) {
2349     SmallVector<int64_t, 4> steps(ivs.size(), 1);
2350     result.addAttribute(AffineParallelOp::getStepsAttrName(),
2351                         builder.getI64ArrayAttr(steps));
2352   } else {
2353     if (parser.parseAffineMapOfSSAIds(stepsMapOperands, stepsMapAttr,
2354                                       AffineParallelOp::getStepsAttrName(),
2355                                       stepsAttrs,
2356                                       OpAsmParser::Delimiter::Paren))
2357       return failure();
2358 
2359     // Convert steps from an AffineMap into an I64ArrayAttr.
2360     SmallVector<int64_t, 4> steps;
2361     auto stepsMap = stepsMapAttr.getValue();
2362     for (const auto &result : stepsMap.getResults()) {
2363       auto constExpr = result.dyn_cast<AffineConstantExpr>();
2364       if (!constExpr)
2365         return parser.emitError(parser.getNameLoc(),
2366                                 "steps must be constant integers");
2367       steps.push_back(constExpr.getValue());
2368     }
2369     result.addAttribute(AffineParallelOp::getStepsAttrName(),
2370                         builder.getI64ArrayAttr(steps));
2371   }
2372 
2373   // Now parse the body.
2374   Region *body = result.addRegion();
2375   SmallVector<Type, 4> types(ivs.size(), indexType);
2376   if (parser.parseRegion(*body, ivs, types) ||
2377       parser.parseOptionalAttrDict(result.attributes))
2378     return failure();
2379 
2380   // Add a terminator if none was parsed.
2381   AffineParallelOp::ensureTerminator(*body, builder, result.location);
2382   return success();
2383 }
2384 
2385 //===----------------------------------------------------------------------===//
2386 // TableGen'd op method definitions
2387 //===----------------------------------------------------------------------===//
2388 
2389 #define GET_OP_CLASSES
2390 #include "mlir/Dialect/Affine/IR/AffineOps.cpp.inc"
2391