1 //===- AffineMap.cpp - MLIR Affine Map Classes ----------------------------===//
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/IR/AffineMap.h"
10 #include "AffineMapDetail.h"
11 #include "mlir/IR/Attributes.h"
12 #include "mlir/IR/StandardTypes.h"
13 #include "mlir/Support/LogicalResult.h"
14 #include "mlir/Support/MathExtras.h"
15 #include "llvm/ADT/SmallSet.h"
16 #include "llvm/ADT/StringRef.h"
17 #include "llvm/Support/raw_ostream.h"
18 
19 using namespace mlir;
20 
21 namespace {
22 
23 // AffineExprConstantFolder evaluates an affine expression using constant
24 // operands passed in 'operandConsts'. Returns an IntegerAttr attribute
25 // representing the constant value of the affine expression evaluated on
26 // constant 'operandConsts', or nullptr if it can't be folded.
27 class AffineExprConstantFolder {
28 public:
29   AffineExprConstantFolder(unsigned numDims, ArrayRef<Attribute> operandConsts)
30       : numDims(numDims), operandConsts(operandConsts) {}
31 
32   /// Attempt to constant fold the specified affine expr, or return null on
33   /// failure.
34   IntegerAttr constantFold(AffineExpr expr) {
35     if (auto result = constantFoldImpl(expr))
36       return IntegerAttr::get(IndexType::get(expr.getContext()), *result);
37     return nullptr;
38   }
39 
40 private:
41   Optional<int64_t> constantFoldImpl(AffineExpr expr) {
42     switch (expr.getKind()) {
43     case AffineExprKind::Add:
44       return constantFoldBinExpr(
45           expr, [](int64_t lhs, int64_t rhs) { return lhs + rhs; });
46     case AffineExprKind::Mul:
47       return constantFoldBinExpr(
48           expr, [](int64_t lhs, int64_t rhs) { return lhs * rhs; });
49     case AffineExprKind::Mod:
50       return constantFoldBinExpr(
51           expr, [](int64_t lhs, int64_t rhs) { return mod(lhs, rhs); });
52     case AffineExprKind::FloorDiv:
53       return constantFoldBinExpr(
54           expr, [](int64_t lhs, int64_t rhs) { return floorDiv(lhs, rhs); });
55     case AffineExprKind::CeilDiv:
56       return constantFoldBinExpr(
57           expr, [](int64_t lhs, int64_t rhs) { return ceilDiv(lhs, rhs); });
58     case AffineExprKind::Constant:
59       return expr.cast<AffineConstantExpr>().getValue();
60     case AffineExprKind::DimId:
61       if (auto attr = operandConsts[expr.cast<AffineDimExpr>().getPosition()]
62                           .dyn_cast_or_null<IntegerAttr>())
63         return attr.getInt();
64       return llvm::None;
65     case AffineExprKind::SymbolId:
66       if (auto attr = operandConsts[numDims +
67                                     expr.cast<AffineSymbolExpr>().getPosition()]
68                           .dyn_cast_or_null<IntegerAttr>())
69         return attr.getInt();
70       return llvm::None;
71     }
72     llvm_unreachable("Unknown AffineExpr");
73   }
74 
75   // TODO: Change these to operate on APInts too.
76   Optional<int64_t> constantFoldBinExpr(AffineExpr expr,
77                                         int64_t (*op)(int64_t, int64_t)) {
78     auto binOpExpr = expr.cast<AffineBinaryOpExpr>();
79     if (auto lhs = constantFoldImpl(binOpExpr.getLHS()))
80       if (auto rhs = constantFoldImpl(binOpExpr.getRHS()))
81         return op(*lhs, *rhs);
82     return llvm::None;
83   }
84 
85   // The number of dimension operands in AffineMap containing this expression.
86   unsigned numDims;
87   // The constant valued operands used to evaluate this AffineExpr.
88   ArrayRef<Attribute> operandConsts;
89 };
90 
91 } // end anonymous namespace
92 
93 /// Returns a single constant result affine map.
94 AffineMap AffineMap::getConstantMap(int64_t val, MLIRContext *context) {
95   return get(/*dimCount=*/0, /*symbolCount=*/0,
96              {getAffineConstantExpr(val, context)});
97 }
98 
99 /// Returns an identity affine map (d0, ..., dn) -> (dp, ..., dn) on the most
100 /// minor dimensions.
101 AffineMap AffineMap::getMinorIdentityMap(unsigned dims, unsigned results,
102                                          MLIRContext *context) {
103   assert(dims >= results && "Dimension mismatch");
104   auto id = AffineMap::getMultiDimIdentityMap(dims, context);
105   return AffineMap::get(dims, 0, id.getResults().take_back(results), context);
106 }
107 
108 bool AffineMap::isMinorIdentity() const {
109   return *this ==
110          getMinorIdentityMap(getNumDims(), getNumResults(), getContext());
111 }
112 
113 /// Returns an AffineMap representing a permutation.
114 AffineMap AffineMap::getPermutationMap(ArrayRef<unsigned> permutation,
115                                        MLIRContext *context) {
116   assert(!permutation.empty() &&
117          "Cannot create permutation map from empty permutation vector");
118   SmallVector<AffineExpr, 4> affExprs;
119   for (auto index : permutation)
120     affExprs.push_back(getAffineDimExpr(index, context));
121   auto m = std::max_element(permutation.begin(), permutation.end());
122   auto permutationMap = AffineMap::get(*m + 1, 0, affExprs, context);
123   assert(permutationMap.isPermutation() && "Invalid permutation vector");
124   return permutationMap;
125 }
126 
127 template <typename AffineExprContainer>
128 static void getMaxDimAndSymbol(ArrayRef<AffineExprContainer> exprsList,
129                                int64_t &maxDim, int64_t &maxSym) {
130   for (const auto &exprs : exprsList) {
131     for (auto expr : exprs) {
132       expr.walk([&maxDim, &maxSym](AffineExpr e) {
133         if (auto d = e.dyn_cast<AffineDimExpr>())
134           maxDim = std::max(maxDim, static_cast<int64_t>(d.getPosition()));
135         if (auto s = e.dyn_cast<AffineSymbolExpr>())
136           maxSym = std::max(maxSym, static_cast<int64_t>(s.getPosition()));
137       });
138     }
139   }
140 }
141 
142 template <typename AffineExprContainer>
143 static SmallVector<AffineMap, 4>
144 inferFromExprList(ArrayRef<AffineExprContainer> exprsList) {
145   assert(!exprsList.empty());
146   assert(!exprsList[0].empty());
147   auto context = exprsList[0][0].getContext();
148   int64_t maxDim = -1, maxSym = -1;
149   getMaxDimAndSymbol(exprsList, maxDim, maxSym);
150   SmallVector<AffineMap, 4> maps;
151   maps.reserve(exprsList.size());
152   for (const auto &exprs : exprsList)
153     maps.push_back(AffineMap::get(/*dimCount=*/maxDim + 1,
154                                   /*symbolCount=*/maxSym + 1, exprs, context));
155   return maps;
156 }
157 
158 SmallVector<AffineMap, 4>
159 AffineMap::inferFromExprList(ArrayRef<ArrayRef<AffineExpr>> exprsList) {
160   return ::inferFromExprList(exprsList);
161 }
162 
163 SmallVector<AffineMap, 4>
164 AffineMap::inferFromExprList(ArrayRef<SmallVector<AffineExpr, 4>> exprsList) {
165   return ::inferFromExprList(exprsList);
166 }
167 
168 AffineMap AffineMap::getMultiDimIdentityMap(unsigned numDims,
169                                             MLIRContext *context) {
170   SmallVector<AffineExpr, 4> dimExprs;
171   dimExprs.reserve(numDims);
172   for (unsigned i = 0; i < numDims; ++i)
173     dimExprs.push_back(mlir::getAffineDimExpr(i, context));
174   return get(/*dimCount=*/numDims, /*symbolCount=*/0, dimExprs, context);
175 }
176 
177 MLIRContext *AffineMap::getContext() const { return map->context; }
178 
179 bool AffineMap::isIdentity() const {
180   if (getNumDims() != getNumResults())
181     return false;
182   ArrayRef<AffineExpr> results = getResults();
183   for (unsigned i = 0, numDims = getNumDims(); i < numDims; ++i) {
184     auto expr = results[i].dyn_cast<AffineDimExpr>();
185     if (!expr || expr.getPosition() != i)
186       return false;
187   }
188   return true;
189 }
190 
191 bool AffineMap::isEmpty() const {
192   return getNumDims() == 0 && getNumSymbols() == 0 && getNumResults() == 0;
193 }
194 
195 bool AffineMap::isSingleConstant() const {
196   return getNumResults() == 1 && getResult(0).isa<AffineConstantExpr>();
197 }
198 
199 int64_t AffineMap::getSingleConstantResult() const {
200   assert(isSingleConstant() && "map must have a single constant result");
201   return getResult(0).cast<AffineConstantExpr>().getValue();
202 }
203 
204 unsigned AffineMap::getNumDims() const {
205   assert(map && "uninitialized map storage");
206   return map->numDims;
207 }
208 unsigned AffineMap::getNumSymbols() const {
209   assert(map && "uninitialized map storage");
210   return map->numSymbols;
211 }
212 unsigned AffineMap::getNumResults() const {
213   assert(map && "uninitialized map storage");
214   return map->results.size();
215 }
216 unsigned AffineMap::getNumInputs() const {
217   assert(map && "uninitialized map storage");
218   return map->numDims + map->numSymbols;
219 }
220 
221 ArrayRef<AffineExpr> AffineMap::getResults() const {
222   assert(map && "uninitialized map storage");
223   return map->results;
224 }
225 AffineExpr AffineMap::getResult(unsigned idx) const {
226   assert(map && "uninitialized map storage");
227   return map->results[idx];
228 }
229 
230 /// Folds the results of the application of an affine map on the provided
231 /// operands to a constant if possible. Returns false if the folding happens,
232 /// true otherwise.
233 LogicalResult
234 AffineMap::constantFold(ArrayRef<Attribute> operandConstants,
235                         SmallVectorImpl<Attribute> &results) const {
236   // Attempt partial folding.
237   SmallVector<int64_t, 2> integers;
238   partialConstantFold(operandConstants, &integers);
239 
240   // If all expressions folded to a constant, populate results with attributes
241   // containing those constants.
242   if (integers.empty())
243     return failure();
244 
245   auto range = llvm::map_range(integers, [this](int64_t i) {
246     return IntegerAttr::get(IndexType::get(getContext()), i);
247   });
248   results.append(range.begin(), range.end());
249   return success();
250 }
251 
252 AffineMap
253 AffineMap::partialConstantFold(ArrayRef<Attribute> operandConstants,
254                                SmallVectorImpl<int64_t> *results) const {
255   assert(getNumInputs() == operandConstants.size());
256 
257   // Fold each of the result expressions.
258   AffineExprConstantFolder exprFolder(getNumDims(), operandConstants);
259   SmallVector<AffineExpr, 4> exprs;
260   exprs.reserve(getNumResults());
261 
262   for (auto expr : getResults()) {
263     auto folded = exprFolder.constantFold(expr);
264     // If did not fold to a constant, keep the original expression, and clear
265     // the integer results vector.
266     if (folded) {
267       exprs.push_back(
268           getAffineConstantExpr(folded.getInt(), folded.getContext()));
269       if (results)
270         results->push_back(folded.getInt());
271     } else {
272       exprs.push_back(expr);
273       if (results) {
274         results->clear();
275         results = nullptr;
276       }
277     }
278   }
279 
280   return get(getNumDims(), getNumSymbols(), exprs, getContext());
281 }
282 
283 /// Walk all of the AffineExpr's in this mapping. Each node in an expression
284 /// tree is visited in postorder.
285 void AffineMap::walkExprs(std::function<void(AffineExpr)> callback) const {
286   for (auto expr : getResults())
287     expr.walk(callback);
288 }
289 
290 /// This method substitutes any uses of dimensions and symbols (e.g.
291 /// dim#0 with dimReplacements[0]) in subexpressions and returns the modified
292 /// expression mapping.  Because this can be used to eliminate dims and
293 /// symbols, the client needs to specify the number of dims and symbols in
294 /// the result.  The returned map always has the same number of results.
295 AffineMap AffineMap::replaceDimsAndSymbols(ArrayRef<AffineExpr> dimReplacements,
296                                            ArrayRef<AffineExpr> symReplacements,
297                                            unsigned numResultDims,
298                                            unsigned numResultSyms) const {
299   SmallVector<AffineExpr, 8> results;
300   results.reserve(getNumResults());
301   for (auto expr : getResults())
302     results.push_back(
303         expr.replaceDimsAndSymbols(dimReplacements, symReplacements));
304 
305   return get(numResultDims, numResultSyms, results, getContext());
306 }
307 
308 AffineMap AffineMap::compose(AffineMap map) {
309   assert(getNumDims() == map.getNumResults() && "Number of results mismatch");
310   // Prepare `map` by concatenating the symbols and rewriting its exprs.
311   unsigned numDims = map.getNumDims();
312   unsigned numSymbolsThisMap = getNumSymbols();
313   unsigned numSymbols = numSymbolsThisMap + map.getNumSymbols();
314   SmallVector<AffineExpr, 8> newDims(numDims);
315   for (unsigned idx = 0; idx < numDims; ++idx) {
316     newDims[idx] = getAffineDimExpr(idx, getContext());
317   }
318   SmallVector<AffineExpr, 8> newSymbols(numSymbols);
319   for (unsigned idx = numSymbolsThisMap; idx < numSymbols; ++idx) {
320     newSymbols[idx - numSymbolsThisMap] =
321         getAffineSymbolExpr(idx, getContext());
322   }
323   auto newMap =
324       map.replaceDimsAndSymbols(newDims, newSymbols, numDims, numSymbols);
325   SmallVector<AffineExpr, 8> exprs;
326   exprs.reserve(getResults().size());
327   for (auto expr : getResults())
328     exprs.push_back(expr.compose(newMap));
329   return AffineMap::get(numDims, numSymbols, exprs, map.getContext());
330 }
331 
332 SmallVector<int64_t, 4> AffineMap::compose(ArrayRef<int64_t> values) {
333   assert(getNumSymbols() == 0 && "Expected symbol-less map");
334   SmallVector<AffineExpr, 4> exprs;
335   exprs.reserve(values.size());
336   MLIRContext *ctx = getContext();
337   for (auto v : values)
338     exprs.push_back(getAffineConstantExpr(v, ctx));
339   auto resMap = compose(AffineMap::get(0, 0, exprs, ctx));
340   SmallVector<int64_t, 4> res;
341   res.reserve(resMap.getNumResults());
342   for (auto e : resMap.getResults())
343     res.push_back(e.cast<AffineConstantExpr>().getValue());
344   return res;
345 }
346 
347 bool AffineMap::isProjectedPermutation() {
348   if (getNumSymbols() > 0)
349     return false;
350   SmallVector<bool, 8> seen(getNumInputs(), false);
351   for (auto expr : getResults()) {
352     if (auto dim = expr.dyn_cast<AffineDimExpr>()) {
353       if (seen[dim.getPosition()])
354         return false;
355       seen[dim.getPosition()] = true;
356       continue;
357     }
358     return false;
359   }
360   return true;
361 }
362 
363 bool AffineMap::isPermutation() {
364   if (getNumDims() != getNumResults())
365     return false;
366   return isProjectedPermutation();
367 }
368 
369 AffineMap AffineMap::getSubMap(ArrayRef<unsigned> resultPos) {
370   SmallVector<AffineExpr, 4> exprs;
371   exprs.reserve(resultPos.size());
372   for (auto idx : resultPos)
373     exprs.push_back(getResult(idx));
374   return AffineMap::get(getNumDims(), getNumSymbols(), exprs, getContext());
375 }
376 
377 AffineMap AffineMap::getMajorSubMap(unsigned numResults) {
378   if (numResults == 0)
379     return AffineMap();
380   if (numResults > getNumResults())
381     return *this;
382   return getSubMap(llvm::to_vector<4>(llvm::seq<unsigned>(0, numResults)));
383 }
384 
385 AffineMap AffineMap::getMinorSubMap(unsigned numResults) {
386   if (numResults == 0)
387     return AffineMap();
388   if (numResults > getNumResults())
389     return *this;
390   return getSubMap(llvm::to_vector<4>(
391       llvm::seq<unsigned>(getNumResults() - numResults, getNumResults())));
392 }
393 
394 AffineMap mlir::simplifyAffineMap(AffineMap map) {
395   SmallVector<AffineExpr, 8> exprs;
396   for (auto e : map.getResults()) {
397     exprs.push_back(
398         simplifyAffineExpr(e, map.getNumDims(), map.getNumSymbols()));
399   }
400   return AffineMap::get(map.getNumDims(), map.getNumSymbols(), exprs,
401                         map.getContext());
402 }
403 
404 AffineMap mlir::removeDuplicateExprs(AffineMap map) {
405   auto results = map.getResults();
406   SmallVector<AffineExpr, 4> uniqueExprs(results.begin(), results.end());
407   uniqueExprs.erase(std::unique(uniqueExprs.begin(), uniqueExprs.end()),
408                     uniqueExprs.end());
409   return AffineMap::get(map.getNumDims(), map.getNumSymbols(), uniqueExprs,
410                         map.getContext());
411 }
412 
413 AffineMap mlir::inversePermutation(AffineMap map) {
414   if (map.isEmpty())
415     return map;
416   assert(map.getNumSymbols() == 0 && "expected map without symbols");
417   SmallVector<AffineExpr, 4> exprs(map.getNumDims());
418   for (auto en : llvm::enumerate(map.getResults())) {
419     auto expr = en.value();
420     // Skip non-permutations.
421     if (auto d = expr.dyn_cast<AffineDimExpr>()) {
422       if (exprs[d.getPosition()])
423         continue;
424       exprs[d.getPosition()] = getAffineDimExpr(en.index(), d.getContext());
425     }
426   }
427   SmallVector<AffineExpr, 4> seenExprs;
428   seenExprs.reserve(map.getNumDims());
429   for (auto expr : exprs)
430     if (expr)
431       seenExprs.push_back(expr);
432   if (seenExprs.size() != map.getNumInputs())
433     return AffineMap();
434   return AffineMap::get(map.getNumResults(), 0, seenExprs, map.getContext());
435 }
436 
437 AffineMap mlir::concatAffineMaps(ArrayRef<AffineMap> maps) {
438   unsigned numResults = 0, numDims = 0, numSymbols = 0;
439   for (auto m : maps)
440     numResults += m.getNumResults();
441   SmallVector<AffineExpr, 8> results;
442   results.reserve(numResults);
443   for (auto m : maps) {
444     for (auto res : m.getResults())
445       results.push_back(res.shiftSymbols(m.getNumSymbols(), numSymbols));
446 
447     numSymbols += m.getNumSymbols();
448     numDims = std::max(m.getNumDims(), numDims);
449   }
450   return AffineMap::get(numDims, numSymbols, results,
451                         maps.front().getContext());
452 }
453 
454 AffineMap mlir::getProjectedMap(AffineMap map,
455                                 ArrayRef<unsigned> projectedDimensions) {
456   DenseSet<unsigned> projectedDims(projectedDimensions.begin(),
457                                    projectedDimensions.end());
458   MLIRContext *context = map.getContext();
459   SmallVector<AffineExpr, 4> resultExprs;
460   for (auto dim : enumerate(llvm::seq<unsigned>(0, map.getNumDims()))) {
461     if (!projectedDims.count(dim.value()))
462       resultExprs.push_back(getAffineDimExpr(dim.index(), context));
463     else
464       resultExprs.push_back(getAffineConstantExpr(0, context));
465   }
466   return map.compose(AffineMap::get(
467       map.getNumDims() - projectedDimensions.size(), 0, resultExprs, context));
468 }
469 
470 //===----------------------------------------------------------------------===//
471 // MutableAffineMap.
472 //===----------------------------------------------------------------------===//
473 
474 MutableAffineMap::MutableAffineMap(AffineMap map)
475     : numDims(map.getNumDims()), numSymbols(map.getNumSymbols()),
476       context(map.getContext()) {
477   for (auto result : map.getResults())
478     results.push_back(result);
479 }
480 
481 void MutableAffineMap::reset(AffineMap map) {
482   results.clear();
483   numDims = map.getNumDims();
484   numSymbols = map.getNumSymbols();
485   context = map.getContext();
486   for (auto result : map.getResults())
487     results.push_back(result);
488 }
489 
490 bool MutableAffineMap::isMultipleOf(unsigned idx, int64_t factor) const {
491   if (results[idx].isMultipleOf(factor))
492     return true;
493 
494   // TODO: use simplifyAffineExpr and FlatAffineConstraints to
495   // complete this (for a more powerful analysis).
496   return false;
497 }
498 
499 // Simplifies the result affine expressions of this map. The expressions have to
500 // be pure for the simplification implemented.
501 void MutableAffineMap::simplify() {
502   // Simplify each of the results if possible.
503   // TODO: functional-style map
504   for (unsigned i = 0, e = getNumResults(); i < e; i++) {
505     results[i] = simplifyAffineExpr(getResult(i), numDims, numSymbols);
506   }
507 }
508 
509 AffineMap MutableAffineMap::getAffineMap() const {
510   return AffineMap::get(numDims, numSymbols, results, context);
511 }
512