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