1 //===- Fusion.cpp - Implementation of linalg Fusion -----------------------===// 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 // This file implements the linalg dialect Fusion on tensors operations pass. 10 // 11 //===----------------------------------------------------------------------===// 12 #include "PassDetail.h" 13 #include "mlir/Dialect/Affine/IR/AffineOps.h" 14 #include "mlir/Dialect/Linalg/IR/LinalgOps.h" 15 #include "mlir/Dialect/Linalg/IR/LinalgTypes.h" 16 #include "mlir/Dialect/Linalg/Passes.h" 17 #include "mlir/Dialect/Linalg/Transforms/Transforms.h" 18 #include "mlir/Dialect/Linalg/Utils/Utils.h" 19 #include "mlir/IR/AffineExpr.h" 20 #include "mlir/IR/AffineMap.h" 21 #include "mlir/IR/PatternMatch.h" 22 #include "mlir/Support/LLVM.h" 23 24 using namespace mlir; 25 using namespace mlir::linalg; 26 27 /// Implementation of fusion of generic ops and indexed_generic ops. 28 // struct FuseGenericOpsOnTensors { 29 static bool areTensorOpsFusable(LinalgOp producer, LinalgOp consumer, 30 unsigned consumerIdx) { 31 // Producer and consumer must have tensor semantics. 32 if (!producer.hasTensorSemantics() || !consumer.hasTensorSemantics()) 33 return false; 34 35 // Verify that 36 // - the producer has all "parallel" iterator type. 37 if (producer.getNumParallelLoops() != producer.getNumLoops()) 38 return false; 39 40 // Get the consumer index map. The number of results of the consumer index 41 // map must match the number of loops of the producer. 42 AffineMap consumerIndexMap = consumer.getIndexingMap(consumerIdx); 43 if (consumerIndexMap.getNumResults() != producer.getNumLoops()) 44 return false; 45 46 // Finally the index_map for the result must be invertible. For now just 47 // verify it is a permutation. 48 AffineMap producerResultIndexMap = producer.getOutputIndexingMap(0); 49 return producerResultIndexMap.isPermutation(); 50 } 51 52 /// Append to `fusedOpIndexingMapAttrs` the indexing maps for the operands of 53 /// the `producer` to use in the fused operation given the indexing map of the 54 /// result of the producer in the consumer. 55 static void getIndexingMapOfProducerOperandsInFusedOp( 56 LinalgOp producer, AffineMap fusedConsumerArgIndexMap, 57 SmallVectorImpl<Attribute> &fusedOpIndexingMapAttrs) { 58 // The indexing map in the consumer op (fusedConsumerArgIndexMap) is a map 59 // from consumer loop -> consumer arg tensor index/producer result tensor 60 // index. The fused loop is same as the consumer loop. For each producer arg 61 // the indexing map to be computed is a map from consumer loop -> producer 62 // arg tensor index. 63 64 AffineMap producerResultIndexMap = producer.getOutputIndexingMap(0); 65 // producerResultIndexMap is a map from producer loop -> tensor index. 66 // Compute the inverse to get map from tensor index -> producer loop. 67 // The inverse is a map from producer result tensor index -> producer loop. 68 AffineMap invProducerResultIndexMap = 69 inversePermutation(producerResultIndexMap); 70 assert(invProducerResultIndexMap && 71 "expected producer result indexig map to be invertible"); 72 for (unsigned argNum : llvm::seq<unsigned>(0, producer.getNumInputs())) { 73 // argMap is a map from producer loop -> producer arg tensor index. 74 AffineMap argMap = producer.getInputIndexingMap(argNum); 75 76 // Compose argMap with invProducerResultIndexMap to get a map from 77 // producer result tensor index -> producer arg tensor index. 78 AffineMap t1 = argMap.compose(invProducerResultIndexMap); 79 80 // Compose t1 with fusedConsumerArgIndexMap gives an indexing map from 81 // consumer loop/ fused loop -> producer arg tensor index. 82 AffineMap indexingMap = t1.compose(fusedConsumerArgIndexMap); 83 fusedOpIndexingMapAttrs.push_back(AffineMapAttr::get(indexingMap)); 84 } 85 } 86 87 /// Generate the region of the fused tensor operation. The region of the fused 88 /// op must be empty. 89 static void generateFusedTensorOpRegion(PatternRewriter &rewriter, 90 Operation *fusedOp, LinalgOp producer, 91 LinalgOp consumer, 92 AffineMap consumerToProducerLoopsMap, 93 unsigned consumerIdx, unsigned nloops) { 94 // Build the region of the fused op. 95 Block &producerBlock = producer.getOperation()->getRegion(0).front(); 96 Block &consumerBlock = consumer.getOperation()->getRegion(0).front(); 97 Block *fusedBlock = new Block(); 98 fusedOp->getRegion(0).push_back(fusedBlock); 99 BlockAndValueMapping mapper; 100 OpBuilder::InsertionGuard guard(rewriter); 101 rewriter.setInsertionPointToStart(fusedBlock); 102 103 // The block arguments are 104 // [index_0, index_1, ... , 105 // consumer_operand_0, ... , consumer_operand_(`consumerIdx`-1), 106 // producer_operand_0, ... , producer_operand_(n-1)], 107 // consumer_operand_(`consumerIdx`), .. consumer_operand_(m-1)] 108 // , where n is the number of producer's operand and m is the number 109 // consumer's operand. 110 // If both `numProducerIndices` and `numConsumerIndices` are zero, this is a 111 // generic op. In this case, there are no indices in block arguments. 112 unsigned numProducerIndices = 113 isa<IndexedGenericOp>(producer.getOperation()) ? nloops : 0; 114 unsigned numConsumerIndices = 115 isa<IndexedGenericOp>(consumer.getOperation()) ? nloops : 0; 116 // Firstly, add all the indices to the block arguments. 117 for (unsigned i = 0, e = std::max(numProducerIndices, numConsumerIndices); 118 i < e; ++i) 119 fusedBlock->addArgument(rewriter.getIndexType()); 120 // Map the arguments for the unmodified args from the consumer. 121 for (auto consumerArg : llvm::enumerate(consumerBlock.getArguments())) { 122 if (consumerArg.index() == consumerIdx + numConsumerIndices) { 123 // Map the arguments for the args from the producer. 124 for (auto producerArg : llvm::enumerate(producerBlock.getArguments())) { 125 // If producer is an indexed_generic op, map the indices from consumer 126 // loop to producer loop (because the fusedOp is built based on 127 // consumer's perspective). 128 if (producerArg.index() < numProducerIndices) { 129 auto newIndex = rewriter.create<mlir::AffineApplyOp>( 130 producer.getLoc(), 131 consumerToProducerLoopsMap.getSubMap(producerArg.index()), 132 fusedBlock->getArguments().take_front(nloops)); 133 mapper.map(producerArg.value(), newIndex); 134 } else { 135 mapper.map(producerArg.value(), 136 fusedBlock->addArgument(producerArg.value().getType())); 137 } 138 } 139 continue; 140 } 141 142 // If consumer is an indexed_generic op, map the indices to the block 143 // arguments directly. Otherwise, add the same type of arugment and map to 144 // it. 145 if (consumerArg.index() < numConsumerIndices) { 146 mapper.map(consumerArg.value(), 147 fusedBlock->getArgument(consumerArg.index())); 148 } else { 149 mapper.map(consumerArg.value(), 150 fusedBlock->addArgument(consumerArg.value().getType())); 151 } 152 } 153 154 // Add operations from producer (except the yield operation) to the fused 155 // op. 156 for (auto &op : producerBlock.getOperations()) { 157 if (auto yieldOp = dyn_cast<linalg::YieldOp>(op)) { 158 // Lookup the value the yield operation is mapped to. 159 Value yieldVal = yieldOp.getOperand(0); 160 if (Value clonedVal = mapper.lookupOrNull(yieldVal)) 161 mapper.map(consumerBlock.getArgument(consumerIdx + numConsumerIndices), 162 clonedVal); 163 continue; 164 } 165 rewriter.clone(op, mapper); 166 } 167 for (auto &op : consumerBlock.getOperations()) 168 rewriter.clone(op, mapper); 169 } 170 171 static Optional<SmallVector<Value, 1>> 172 fuseTensorOpsImpl(LinalgOp producer, LinalgOp consumer, unsigned consumerIdx, 173 PatternRewriter &rewriter, 174 OperationFolder *folder = nullptr) { 175 if (!areTensorOpsFusable(producer, consumer, consumerIdx)) 176 return llvm::None; 177 178 unsigned numFusedOperands = 179 producer.getNumInputs() + consumer.getNumInputs() - 1; 180 181 // Compute the fused operands list, 182 SmallVector<Value, 2> fusedOperands; 183 fusedOperands.reserve(numFusedOperands); 184 auto consumerOperands = consumer.getInputs(); 185 auto producerOperands = producer.getInputs(); 186 fusedOperands.assign(consumerOperands.begin(), 187 std::next(consumerOperands.begin(), consumerIdx)); 188 fusedOperands.append(producerOperands.begin(), producerOperands.end()); 189 fusedOperands.append(std::next(consumerOperands.begin(), consumerIdx + 1), 190 consumerOperands.end()); 191 192 // Compute indexing_maps for the fused operation. The indexing_maps for the 193 // operands of the consumers that arent fused are the same. The 194 // indexing_maps for the producers need to be computed based on the 195 // indexing_map of the operand at consumerIdx in the consumer. 196 SmallVector<Attribute, 4> fusedIndexMaps; 197 auto consumerIndexMaps = consumer.indexing_maps(); 198 fusedIndexMaps.reserve(fusedOperands.size() + consumer.getNumOutputs()); 199 fusedIndexMaps.assign(consumerIndexMaps.begin(), 200 std::next(consumerIndexMaps.begin(), consumerIdx)); 201 // Compute indexing maps for the producer args in the fused operation. 202 getIndexingMapOfProducerOperandsInFusedOp( 203 producer, consumer.getInputIndexingMap(consumerIdx), fusedIndexMaps); 204 205 // Append the indexing maps for the remaining consumer operands. 206 fusedIndexMaps.append(std::next(consumerIndexMaps.begin(), consumerIdx + 1), 207 consumerIndexMaps.end()); 208 209 // Generate the fused op. 210 // Tensor-level fusion is only on ops without initTensors and outputBuffers. 211 LinalgOp fusedOp; 212 if (isa<GenericOp>(producer.getOperation()) && 213 isa<GenericOp>(consumer.getOperation())) { 214 fusedOp = rewriter 215 .create<GenericOp>(consumer.getLoc(), 216 consumer.getOperation()->getResultTypes(), 217 /*inputs=*/fusedOperands, 218 /*outputBuffers=*/ValueRange{}, 219 /*initTensors=*/ValueRange{}, 220 rewriter.getArrayAttr(fusedIndexMaps), 221 consumer.iterator_types(), 222 /*doc=*/nullptr, 223 /*library_call=*/nullptr, 224 /*symbol_source=*/nullptr) 225 .getOperation(); 226 } else { 227 fusedOp = 228 rewriter 229 .create<IndexedGenericOp>(consumer.getLoc(), 230 consumer.getOperation()->getResultTypes(), 231 /*inputs=*/fusedOperands, 232 /*outputBuffers=*/ValueRange{}, 233 /*initTensors=*/ValueRange{}, 234 rewriter.getArrayAttr(fusedIndexMaps), 235 consumer.iterator_types(), 236 /*doc=*/nullptr, 237 /*library_call=*/nullptr, 238 /*symbol_source=*/nullptr) 239 .getOperation(); 240 } 241 242 // Construct an AffineMap from consumer loops to producer loops. 243 // consumer loop -> tensor index 244 AffineMap consumerResultIndexMap = consumer.getInputIndexingMap(consumerIdx); 245 // producer loop -> tensor index 246 AffineMap producerResultIndexMap = producer.getOutputIndexingMap(0); 247 // tensor index -> producer loop 248 AffineMap invProducerResultIndexMap = 249 inversePermutation(producerResultIndexMap); 250 assert(invProducerResultIndexMap && 251 "expected producer result indexig map to be invertible"); 252 // consumer loop -> producer loop 253 AffineMap consumerToProducerLoopsMap = 254 invProducerResultIndexMap.compose(consumerResultIndexMap); 255 256 generateFusedTensorOpRegion(rewriter, fusedOp.getOperation(), producer, 257 consumer, consumerToProducerLoopsMap, consumerIdx, 258 consumer.getNumLoops()); 259 return SmallVector<Value, 1>(fusedOp.getOperation()->getResults()); 260 } 261 262 /// Linearize the expressions in `sourceMap` based on the `reassociationMaps` 263 /// provided, given the shape of the source tensor that corresponds to the 264 /// `sourceMap`. Note that this implicitly assumes that the tensors dimensions 265 /// are "row-major" ordered logically. 266 /// 267 /// For example: 268 /// 269 /// %0 = op ... : tensor<?x?x4x5xf32> 270 /// with output index_map `affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>` 271 /// 272 /// and reshape: 273 /// %1 = linalg.tensor_reshape %0 [affine_map<(i, j, k, l) -> (i)>, 274 /// affine_map<(i, j, k, l) -> (j, k, l)>] : 275 /// tensor<?x?x4x5xf32> into tensor<?x?xf32> 276 /// 277 /// would be rewritten into: 278 /// %0 = op ... : tensor<?x?x4x5xf32> 279 /// with output index_map 280 /// `affine_map<(d0, d1, d2, d3) -> (d0, d1 * 20 + d2 * 5 + d3)>` 281 static AffineMap linearizeCollapsedDims(AffineMap sourceMap, 282 ArrayRef<int64_t> sourceShape, 283 ArrayRef<AffineMap> reassociationMaps) { 284 SmallVector<AffineExpr, 4> resultExprs; 285 resultExprs.reserve(reassociationMaps.size()); 286 ArrayRef<AffineExpr> sourceExprs = sourceMap.getResults(); 287 MLIRContext *context = sourceMap.getContext(); 288 289 // Compute the result exprs based on the reassociation maps. 290 for (AffineMap map : reassociationMaps) { 291 ArrayRef<AffineExpr> collapsedDims = map.getResults(); 292 // Assume that they are in-order and contiguous (already checked in 293 // verifier). 294 assert(!collapsedDims.empty()); 295 unsigned startDim = 296 collapsedDims.front().cast<AffineDimExpr>().getPosition(); 297 AffineExpr linearizedExpr = makeCanonicalStridedLayoutExpr( 298 sourceShape.slice(startDim, collapsedDims.size()), 299 sourceExprs.slice(startDim, collapsedDims.size()), context); 300 resultExprs.push_back(linearizedExpr); 301 } 302 return AffineMap::get(sourceMap.getNumDims(), sourceMap.getNumSymbols(), 303 resultExprs, context); 304 } 305 306 /// Checks if the `reshapeOp` can be fused with it consumer (if `asProducer` is 307 /// true) or its producer (if `asProducer` is false) given the indexing map at 308 /// its use. 309 static bool isTensorReshapeOpFoldableByLinearization(TensorReshapeOp reshapeOp, 310 AffineMap useIndexMap, 311 bool asProducer) { 312 RankedTensorType returnType = reshapeOp.getResultType(); 313 RankedTensorType operandType = reshapeOp.getSrcType(); 314 // Reshape is fusable with its consumer (i.e. reshape as a producer) when its 315 // operand is of lesser rank than the result. Fusing when operand has higher 316 // rank will require use of mods and divs in the indexing maps of the fused op 317 // which would make it non-invertible. Similarly reshape is fused with its 318 // producer (i.e. reshape as consumer) only if the return type has lesser 319 // rank. 320 if ((asProducer && reshapeOp.getSrcType().hasStaticShape() && 321 returnType.getRank() < operandType.getRank()) || 322 (!asProducer && reshapeOp.getResultType().hasStaticShape() && 323 operandType.getRank() < returnType.getRank())) 324 return false; 325 return useIndexMap.isPermutation(); 326 } 327 328 /// Based on the type of `op` create a linalg op of the same type, i.e. if `op` 329 /// is a linalg.generic operation, the create a `linalg.generic` operation with 330 /// the given `args`. Expects `op` to be `linalg.generic` or 331 /// `linalg.indexed_generic`. 332 template <typename... Args> 333 static LinalgOp createLinalgOpOfSameType(LinalgOp op, PatternRewriter &rewriter, 334 Args... args) { 335 if (isa<GenericOp>(op.getOperation())) 336 return cast<LinalgOp>(rewriter.create<GenericOp>(args...).getOperation()); 337 if (isa<IndexedGenericOp>(op.getOperation())) 338 return cast<LinalgOp>( 339 rewriter.create<IndexedGenericOp>(args...).getOperation()); 340 llvm_unreachable( 341 "expected only linalg.generic or linalg.indexed_generic ops"); 342 return nullptr; 343 } 344 345 /// Conditions for folding a generic/indexed-generic operation with a reshape op 346 /// by expanding the iteration space dimensionality for tensor operations. These 347 /// are preconditions assumed by `foldReshapeByDimExpansion` which implements 348 /// the following fusion pattern. 349 /// 350 /// Consider 351 /// 352 /// %c = linalg.generic ins(%a, %b : memref<?x?x?xf32>, memref<?x?xf32>) 353 /// indexing_maps = [affine_map<(d0, d1, d2) -> (d1, d0, d2)>, 354 /// affine_map<(d0, d1, d2) -> (d1, d2)>, 355 /// affine_map<(d0, d1, d2) -> (d0, d2, d1)>] 356 /// %d = linalg.tensor_reshape %c 357 /// [affine_map<(d0, d1, d2, d3, d4, d5) -> (d0, d1)>, 358 /// affine_map<(d0, d1, d2, d3, d4, d5) -> (d2)>, 359 /// affine_map<(d0, d1, d2, d3, d4, d5) -> (d3, d4, d5)>] 360 /// : tensor<?x?x?xf32> into tensor<?x?x?x?x?x?xf32> 361 /// 362 /// The reshape can be folded into the `linalgOp` if the 363 /// generic/indexed-generic op loop dimensionality is increased to match the 364 /// result (operand) of the tensor_reshape when the reshape is expanding 365 /// (folding). The indexing_map of the fused tensor in the `linalgOp` and the 366 /// reassociation map helps compute the indexing maps of the modified op. For 367 /// the above example, based on the reassociation map it can be concluded that 368 /// 369 /// - The loop used to access the first dimension of the fused tensor is split 370 /// into two. 371 /// - The loop used to access the second dimension of the fused tensor is kept 372 /// as is. 373 /// - The loop used to access the third dimension of the fused tensor is split 374 /// into three. 375 /// 376 /// i.e. (e0, e1, e2, e3, e4) is the domain of the indexing map of the modified 377 /// op, then 378 /// 379 /// d0 -> e0, e1 380 /// d1 -> e2, e3, e4 381 /// d2 -> e5 382 /// 383 /// substituting this, the generic op can be rewritten as 384 /// 385 /// %d = linalg.generic ins(%0, %1 : ) 386 /// indexing_maps = 387 /// [affine_map<(e0, e1, e2, e3, e4, e5) -> (e2, e3, e4, e0, e1, e5)>, 388 /// affine_map<(e0, e1, e2, e3, e4, e5) -> (e2, e3, e4, e5)>, 389 /// affine_map<(e0, e1, e2, e3, e4, e5) -> (e0, e1, e5, e2, e3, e4)>] 390 /// 391 /// Since operands to the linalg generic are now 5D, reshapes can be introduced 392 /// to make it consistent 393 /// 394 /// %0 = linalg.tensor_reshape %a 395 /// [affine_map<(e0, e1, e2, e3, e4, e5) -> (e0, e1, e2), 396 /// affine_map<(e0, e1, e2, e3, e4, e5) -> (e3, e4), 397 /// affine_map<(e0, e1, e2, e3, e4, e5) -> (e5)] 398 /// : tensor<?x?x?xf32> into tensor<?x?x?x?x?x?xf32> 399 /// %1 = linalg.tensor_reshape %b 400 /// [affine_map<(e0, e1, e2, e3) -> (e0, e1, e2), 401 /// affine_map<(e0, e1, e2, e3) -> (e3)] 402 /// : tensor<?x?x?xf32> into tensor<?x?x?x?xf32> 403 /// 404 /// The added reshapes are again expanding patterns, so they will get fused 405 /// with its producers if possible. 406 static bool isFusableWithReshapeByDimExpansion(LinalgOp linalgOp, 407 unsigned fusedTensorIndex) { 408 // Is fusable only if: 409 // - The linalgOp is a generic op. 410 // - All the indexing maps for operands in linalgOp are projected 411 // permutations. 412 // - The indexing map at the position representing the fused tensor is a 413 // permutation. 414 // - All the loops in linalgOp are parallel loops. 415 return isa<GenericOp>(linalgOp.getOperation()) && 416 linalgOp.hasTensorSemantics() && 417 llvm::all_of(linalgOp.indexing_maps().getValue().take_front( 418 linalgOp.getNumInputs()), 419 [](Attribute attr) { 420 return attr.cast<AffineMapAttr>() 421 .getValue() 422 .isProjectedPermutation(); 423 }) && 424 linalgOp.getIndexingMap(fusedTensorIndex).isPermutation() && 425 llvm::all_of(linalgOp.iterator_types(), [](Attribute attr) { 426 return attr.cast<StringAttr>().getValue() == 427 getParallelIteratorTypeName(); 428 }); 429 } 430 431 /// Implements the fusion of a tensor_reshape op and a generic/indexed_generic 432 /// op as explained in `isFusableWithReshapeByExpansion`. Assumes that those 433 /// conditions have been satisfied. 434 static Optional<SmallVector<Value, 1>> 435 fuseWithReshapeByExpansion(LinalgOp linalgOp, TensorReshapeOp reshapeOp, 436 unsigned fusedTensorIndex, PatternRewriter &rewriter, 437 OperationFolder *folder = nullptr) { 438 assert(isFusableWithReshapeByDimExpansion(linalgOp, fusedTensorIndex) && 439 "preconditions for fuse operation failed"); 440 // Check if reshape is expanding or collapsing. 441 bool isExpanding = 442 reshapeOp.getSrcType().getRank() < reshapeOp.getResultType().getRank(); 443 RankedTensorType expandedType = 444 isExpanding ? reshapeOp.getResultType() : reshapeOp.getSrcType(); 445 RankedTensorType foldedType = 446 isExpanding ? reshapeOp.getSrcType() : reshapeOp.getResultType(); 447 AffineMap fusedIndexMap = linalgOp.getIndexingMap(fusedTensorIndex); 448 449 // The reshape is folding/expanding consecutive dimensions. Given the indexing 450 // map of the fused tensor find the number of dimensions each of the loops of 451 // the original op is expanded into. Also record the shape of the expanded 452 // dimensions. 453 ArrayRef<int64_t> expandedShape = expandedType.getShape(); 454 SmallVector<unsigned, 4> numFoldedDims(foldedType.getRank(), 0); 455 SmallVector<SmallVector<int64_t, 4>, 4> expandedDimsShape( 456 expandedType.getRank()); 457 auto reassociationMaps = reshapeOp.getReassociationMaps(); 458 for (auto resultExpr : llvm::enumerate(fusedIndexMap.getResults())) { 459 unsigned pos = resultExpr.value().cast<AffineDimExpr>().getPosition(); 460 AffineMap foldedDims = reassociationMaps[resultExpr.index()]; 461 numFoldedDims[pos] = foldedDims.getNumResults(); 462 ArrayRef<int64_t> shape = expandedShape.slice( 463 foldedDims.getResult(0).cast<AffineDimExpr>().getPosition(), 464 numFoldedDims[pos]); 465 expandedDimsShape[pos].assign(shape.begin(), shape.end()); 466 } 467 468 // The remapping of the indices is then the prefix sum (inclusive) of the 469 // numFoldedDims. 470 SmallVector<unsigned, 4> remapping(numFoldedDims.size() + 1, 0); 471 unsigned sum = 0; 472 for (auto numFoldedDim : llvm::enumerate(numFoldedDims)) { 473 sum += numFoldedDim.value(); 474 remapping[numFoldedDim.index() + 1] = sum; 475 } 476 477 SmallVector<AffineMap, 4> expandedOpIndexingMaps; 478 // Compute the modified indexing maps by replacing every loop (AffineDimExpr) 479 // in the original indexing map with the sequence of loops that it is expanded 480 // to. 481 for (AffineMap indexingMap : linalgOp.getIndexingMaps()) { 482 SmallVector<AffineExpr, 4> newExprs; 483 for (AffineExpr expr : indexingMap.getResults()) { 484 unsigned pos = expr.cast<AffineDimExpr>().getPosition(); 485 for (unsigned newPos : 486 llvm::seq<unsigned>(remapping[pos], remapping[pos + 1])) { 487 newExprs.push_back(rewriter.getAffineDimExpr(newPos)); 488 } 489 } 490 expandedOpIndexingMaps.push_back( 491 AffineMap::get(remapping.back(), indexingMap.getNumSymbols(), newExprs, 492 rewriter.getContext())); 493 } 494 495 // The operands of the expanded op are computed by reshaping the original 496 // operands. The reshape depends on the ordering of the loop used to access 497 // the tensor in the original operation, and are expanded into as many 498 // dimensions as the loop is expanded into (as computed by `remapping`). 499 auto getReshapeInfo = 500 [&](AffineMap operandIndexingMap, 501 SmallVectorImpl<ReassociationIndices> &reassociation, 502 SmallVectorImpl<int64_t> &expandedOpOperandShape) { 503 unsigned reshapeDims = 0; 504 for (AffineExpr expr : operandIndexingMap.getResults()) { 505 unsigned origDim = expr.cast<AffineDimExpr>().getPosition(); 506 auto foldedDims = llvm::seq<int64_t>( 507 reshapeDims, reshapeDims + numFoldedDims[origDim]); 508 reassociation.emplace_back(foldedDims.begin(), foldedDims.end()); 509 expandedOpOperandShape.append(expandedDimsShape[origDim].begin(), 510 expandedDimsShape[origDim].end()); 511 reshapeDims += numFoldedDims[origDim]; 512 } 513 }; 514 SmallVector<Value, 4> expandedOpOperands; 515 for (auto operand : llvm::enumerate(linalgOp.getInputs())) { 516 if (operand.index() == fusedTensorIndex) { 517 expandedOpOperands.push_back(reshapeOp.src()); 518 continue; 519 } 520 AffineMap indexingMap = linalgOp.getIndexingMap(operand.index()); 521 SmallVector<ReassociationIndices, 4> reassociation; 522 SmallVector<int64_t, 4> expandedOperandShape; 523 getReshapeInfo(indexingMap, reassociation, expandedOperandShape); 524 Type expandedOperandType = RankedTensorType::get( 525 expandedOperandShape, 526 operand.value().getType().cast<ShapedType>().getElementType()); 527 if (expandedOperandType != operand.value().getType()) { 528 expandedOpOperands.push_back(rewriter.create<TensorReshapeOp>( 529 linalgOp.getLoc(), expandedOperandType, operand.value(), 530 reassociation)); 531 } else { 532 expandedOpOperands.push_back(operand.value()); 533 } 534 } 535 SmallVector<Type, 1> resultTypes; 536 SmallVector<SmallVector<ReassociationIndices, 4>, 1> resultReassociation; 537 for (auto result : llvm::enumerate(linalgOp.getOperation()->getResults())) { 538 AffineMap indexingMap = 539 linalgOp.getIndexingMap(linalgOp.getNumInputs() + result.index()); 540 SmallVector<ReassociationIndices, 4> reassociation; 541 SmallVector<int64_t, 4> expandedResultShape; 542 getReshapeInfo(indexingMap, reassociation, expandedResultShape); 543 resultTypes.push_back(RankedTensorType::get( 544 expandedResultShape, 545 result.value().getType().cast<ShapedType>().getElementType())); 546 resultReassociation.emplace_back(std::move(reassociation)); 547 } 548 549 // The iterator types of the expanded op are all parallel. 550 SmallVector<StringRef, 4> iteratorTypes(remapping.back(), 551 getParallelIteratorTypeName()); 552 553 LinalgOp fusedOp = createLinalgOpOfSameType( 554 linalgOp, rewriter, linalgOp.getLoc(), resultTypes, 555 /*inputs=*/expandedOpOperands, 556 /*outputBuffers=*/ValueRange{}, 557 /*initTensors=*/ValueRange{}, expandedOpIndexingMaps, iteratorTypes); 558 Region &fusedRegion = fusedOp.getOperation()->getRegion(0); 559 // TODO: Add support for indexed generic op, which would need mapping the 560 // expanded dimensions to the original dimension arguments. 561 rewriter.cloneRegionBefore(linalgOp.getOperation()->getRegion(0), fusedRegion, 562 fusedRegion.begin()); 563 564 // Reshape the result values to their original shape if this is a collapsing 565 // reshape folded into its consumer. 566 SmallVector<Value, 1> resultVals; 567 for (auto result : llvm::enumerate(linalgOp.getOperation()->getResults())) { 568 if (!isExpanding && 569 resultTypes[result.index()] != result.value().getType()) { 570 resultVals.push_back(rewriter.create<TensorReshapeOp>( 571 linalgOp.getLoc(), result.value().getType(), 572 fusedOp.getOperation()->getResult(result.index()), 573 resultReassociation[result.index()])); 574 } else { 575 resultVals.push_back(fusedOp.getOperation()->getResult(result.index())); 576 } 577 } 578 // Assuming a single result. 579 return resultVals; 580 } 581 582 namespace { 583 584 /// Pattern to fold tensor_reshape op with its consumer by using the source of 585 /// the reshape op as the operand in the consumer (instead of the result of the 586 /// tensor_reshapeop) when the tensor_reshape op is collapsing. The 587 /// corresponding index map in the consumer needs to be modified to linearize 588 /// the folded dimension. 589 /// 590 /// For example, 591 /// 592 /// #map0 = affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)> 593 /// %0 = linalg.tensor_reshape %arg0 594 /// [affine_map<(i, j, k, l) -> (i)>, affine_map<(i, j, k, l) -> (j, k)>, 595 /// affine_map<(i, j, k, l) -> (l)>] 596 /// tensor<?x?x?xf32> into tensor<?x?x4x?xf32> 597 /// %1 = linalg.generic { indexing_maps = [#map0, #map0, #map0], ... } 598 /// ins(%0, %arg1 : tensor<?x?x4x?xf32>, tensor<?x?x4x?xf32>) ... 599 /// -> tensor<?x?x4x?xf32> 600 /// 601 /// can be folded into 602 /// 603 /// #map0 = affine_map<(d0, d1, d2, d3) -> (d0, d1 * 4 + d2, d3)> 604 /// #map1 = affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)> 605 /// %0 = linalg.generic { indexing_maps = [#map0, #map1, #map1] ... } 606 /// ins(%arg0, %arg1 : tensor<?x?x?xf32>, tensor<?x?x4x?xf32>) ... 607 /// -> tensor<?x?x4x?xf32> 608 template <typename LinalgOpTy> 609 struct FoldProducerReshapeOpByLinearization 610 : public OpRewritePattern<LinalgOpTy> { 611 using OpRewritePattern<LinalgOpTy>::OpRewritePattern; 612 613 LogicalResult matchAndRewrite(LinalgOpTy op, 614 PatternRewriter &rewriter) const override { 615 if (!op.hasTensorSemantics()) 616 return failure(); 617 LinalgOp linalgOp = cast<LinalgOp>(op.getOperation()); 618 for (auto operand : llvm::enumerate(linalgOp.getInputs())) { 619 TensorReshapeOp reshapeOp = 620 operand.value().getDefiningOp<TensorReshapeOp>(); 621 if (!reshapeOp || 622 !isTensorReshapeOpFoldableByLinearization( 623 reshapeOp, linalgOp.getInputIndexingMap(operand.index()), 624 /*asProducer =*/true)) 625 continue; 626 627 // Compute the fused operands list, 628 SmallVector<Value, 2> fusedOperands(linalgOp.getInputs()); 629 fusedOperands[operand.index()] = reshapeOp.src(); 630 631 // Compute indexing_maps for the fused operation. The indexing_maps for 632 // the operands of the consumers that arent fused are the same. 633 SmallVector<AffineMap, 4> fusedIndexMaps = llvm::to_vector<4>( 634 op.indexing_maps().template getAsValueRange<AffineMapAttr>()); 635 636 // Accepted consumer maps are either identity or permutation. 637 auto invMap = inversePermutation(fusedIndexMaps[operand.index()]); 638 639 // Compute the indexing map to use for the result of the producer. 640 AffineMap modifiedMap = 641 linearizeCollapsedDims(invMap, reshapeOp.getResultType().getShape(), 642 reshapeOp.getReassociationMaps()); 643 for (AffineExpr expr : modifiedMap.getResults()) { 644 if (!expr.isPureAffine()) 645 return failure(); 646 } 647 fusedIndexMaps[operand.index()] = modifiedMap; 648 649 // Further check that the resulting index maps can be fused and 650 // inverted. Without this the resultant op is not legal. 651 if (!inversePermutation(concatAffineMaps(fusedIndexMaps))) 652 return op.emitRemark("fused op loop bound computation failed"); 653 654 rewriter.startRootUpdate(op); 655 op.getOperation()->setOperands(fusedOperands); 656 op.indexing_mapsAttr(rewriter.getAffineMapArrayAttr(fusedIndexMaps)); 657 rewriter.finalizeRootUpdate(op); 658 if (reshapeOp.use_empty()) 659 rewriter.eraseOp(reshapeOp); 660 return success(); 661 } 662 return op.emitRemark("no fusion candidates found"); 663 } 664 }; 665 666 /// Pattern to fuse a tensor_reshape op with its consumer generic op, when the 667 /// reshape op is collapsing dimensions. The dimensionality of the loop in the 668 /// consumer generic op is expanded. 669 struct FoldWithProducerReshapeOpByExpansion 670 : public OpRewritePattern<GenericOp> { 671 using OpRewritePattern<GenericOp>::OpRewritePattern; 672 673 LogicalResult matchAndRewrite(GenericOp genericOp, 674 PatternRewriter &rewriter) const override { 675 LinalgOp linalgOp = cast<LinalgOp>(genericOp.getOperation()); 676 for (auto operand : llvm::enumerate(linalgOp.getInputs())) { 677 TensorReshapeOp reshapeOp = 678 operand.value().getDefiningOp<TensorReshapeOp>(); 679 if (!reshapeOp) 680 continue; 681 682 // Fold only if 683 // - The tensor reshape op is folding. 684 // - All constraints of fusing with reshape by expansion are met. 685 if (reshapeOp.getSrcType().getRank() < 686 reshapeOp.getResultType().getRank() || 687 !isFusableWithReshapeByDimExpansion(linalgOp, operand.index())) 688 continue; 689 690 Optional<SmallVector<Value, 1>> replacementValues = 691 fuseWithReshapeByExpansion(linalgOp, reshapeOp, operand.index(), 692 rewriter); 693 if (!replacementValues) 694 return failure(); 695 rewriter.replaceOp(genericOp, replacementValues.getValue()); 696 if (reshapeOp.use_empty()) 697 rewriter.eraseOp(reshapeOp); 698 return success(); 699 } 700 return failure(); 701 } 702 }; 703 704 /// Pattern to fold tensor_reshape op with its producer. The corresponding index 705 /// map in the consumer needs to be modified to linearize the folded dimension. 706 struct FoldConsumerReshapeOpByLinearization 707 : public OpRewritePattern<TensorReshapeOp> { 708 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 709 710 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 711 PatternRewriter &rewriter) const override { 712 LinalgOp producer = reshapeOp.src().getDefiningOp<LinalgOp>(); 713 if (!producer || 714 !isa<GenericOp, IndexedGenericOp>(producer.getOperation()) || 715 !producer.hasTensorSemantics() || producer.getNumOutputs() != 1 || 716 !isTensorReshapeOpFoldableByLinearization( 717 reshapeOp, producer.getOutputIndexingMap(0), /*asProducer =*/false)) 718 return failure(); 719 // The indexing_maps for the operands of the fused operation are same as 720 // those for the operands of the producer. 721 SmallVector<AffineMap, 4> fusedIndexMaps = llvm::to_vector<4>( 722 producer.indexing_maps().getAsValueRange<AffineMapAttr>()); 723 724 auto invMap = inversePermutation(producer.getOutputIndexingMap(0)); 725 726 // Compute the indexing map to use for the operand of the producer. 727 AffineMap modifiedMap = 728 linearizeCollapsedDims(invMap, reshapeOp.getSrcType().getShape(), 729 reshapeOp.getReassociationMaps()); 730 for (AffineExpr expr : modifiedMap.getResults()) { 731 if (!expr.isPureAffine()) 732 return reshapeOp.emitRemark("fused op indexing map is not affine"); 733 } 734 fusedIndexMaps.back() = modifiedMap; 735 736 // Further check that the resulting index maps can be fused and 737 // inverted. Without this the resultant op is not legal. 738 if (!inversePermutation(concatAffineMaps(fusedIndexMaps))) 739 return reshapeOp.emitRemark("fused op loop bound computation failed"); 740 741 LinalgOp fusedOp = createLinalgOpOfSameType( 742 producer, rewriter, rewriter.getUnknownLoc(), reshapeOp.getResultType(), 743 /*inputs=*/producer.getInputs(), 744 /*outputBuffers=*/ValueRange{}, 745 /*initTensors=*/ValueRange{}, // no init tensors for now. 746 rewriter.getAffineMapArrayAttr(fusedIndexMaps), 747 producer.iterator_types(), 748 /*doc=*/nullptr, 749 /*library_call=*/nullptr, 750 /*symbol_source=*/nullptr); 751 auto &fusedRegion = fusedOp.getOperation()->getRegion(0); 752 rewriter.cloneRegionBefore(producer.getOperation()->getRegion(0), 753 fusedRegion, fusedRegion.begin()); 754 rewriter.replaceOp(reshapeOp, fusedOp.getOperation()->getResults()); 755 if (producer.use_empty()) 756 rewriter.eraseOp(producer); 757 return success(); 758 } 759 }; 760 761 /// Pattern to fold a tensor_reshape op with its producer generic op if the 762 /// tensor_reshape op is expanding, by expanding the dimensionality of the loop 763 /// in the producer op. 764 struct FoldReshapeWithGenericOpByExpansion 765 : public OpRewritePattern<TensorReshapeOp> { 766 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 767 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 768 PatternRewriter &rewriter) const override { 769 // Fold only if 770 // - The tensor reshape op is a expanding case. 771 // - All constraints of fusing with reshape by expansion are met. 772 if (reshapeOp.getSrcType().getRank() > reshapeOp.getResultType().getRank()) 773 return failure(); 774 LinalgOp producer = reshapeOp.src().getDefiningOp<LinalgOp>(); 775 if (!producer || producer.getNumOutputs() != 1 || 776 !isFusableWithReshapeByDimExpansion(producer, producer.getNumInputs())) 777 return failure(); 778 Optional<SmallVector<Value, 1>> replacementValues = 779 fuseWithReshapeByExpansion(producer, reshapeOp, producer.getNumInputs(), 780 rewriter); 781 if (!replacementValues) 782 return failure(); 783 rewriter.replaceOp(reshapeOp, replacementValues.getValue()); 784 if (producer.use_empty()) 785 rewriter.eraseOp(producer); 786 return success(); 787 } 788 }; 789 790 /// Pattern to fold a GenericOp/IndexedGenericOp with a splat constant. 791 template <typename LinalgOpTy> 792 struct FoldSplatConstants : public OpRewritePattern<LinalgOpTy> { 793 using OpRewritePattern<LinalgOpTy>::OpRewritePattern; 794 795 LogicalResult matchAndRewrite(LinalgOpTy op, 796 PatternRewriter &rewriter) const override { 797 if (!op.hasTensorSemantics()) 798 return failure(); 799 LinalgOp linalgOp = cast<LinalgOp>(op.getOperation()); 800 for (auto operand : llvm::enumerate(linalgOp.getInputs())) { 801 ConstantOp constantOp = operand.value().getDefiningOp<ConstantOp>(); 802 if (!constantOp || 803 !constantOp.value().cast<DenseElementsAttr>().isSplat()) 804 continue; 805 806 // The indexing_maps for the operands of the fused operation are same as 807 // those for the operands of the linalgOp without the indexing map at 808 // operand.index() 809 SmallVector<AffineMap, 4> fusedIndexMaps = llvm::to_vector<4>( 810 linalgOp.indexing_maps().getAsValueRange<AffineMapAttr>()); 811 fusedIndexMaps.erase(std::next(fusedIndexMaps.begin(), operand.index())); 812 813 // The operands list is same as the linalgOp with the argument for 814 // constant index dropped. 815 SmallVector<Value, 4> fusedOperands(linalgOp.getInputs()); 816 fusedOperands.erase(std::next(fusedOperands.begin(), operand.index())); 817 818 // Create a constant scalar value from the splat constant. 819 Value scalarConstant = rewriter.create<ConstantOp>( 820 constantOp.getLoc(), 821 constantOp.value().cast<DenseElementsAttr>().getSplatValue()); 822 823 LinalgOp fusedOp = createLinalgOpOfSameType( 824 linalgOp, rewriter, rewriter.getUnknownLoc(), 825 linalgOp.getOperation()->getResultTypes(), 826 /*inputs=*/fusedOperands, 827 /*outputBuffers=*/ValueRange{}, 828 /*initTensors=*/ValueRange{}, // no init tensors for now. 829 rewriter.getAffineMapArrayAttr(fusedIndexMaps), 830 linalgOp.iterator_types(), 831 /*doc=*/nullptr, 832 /*library_call=*/nullptr, 833 /*symbol_source=*/nullptr); 834 835 // Map the block argument corresponding to the replaced argument with the 836 // scalar constant. 837 Region &linalgOpRegion = linalgOp.getOperation()->getRegion(0); 838 Block &entryBlock = *linalgOpRegion.begin(); 839 unsigned argIndex = entryBlock.getNumArguments() - 840 linalgOp.getNumInputs() + operand.index(); 841 BlockAndValueMapping mapping; 842 mapping.map(entryBlock.getArgument(argIndex), scalarConstant); 843 Region &fusedRegion = fusedOp.getOperation()->getRegion(0); 844 rewriter.cloneRegionBefore(linalgOpRegion, fusedRegion, 845 fusedRegion.begin(), mapping); 846 rewriter.replaceOp(linalgOp, fusedOp.getOperation()->getResults()); 847 if (constantOp.use_empty()) 848 rewriter.eraseOp(constantOp); 849 return success(); 850 } 851 return failure(); 852 } 853 }; 854 } // namespace 855 856 Optional<SmallVector<Value, 1>> 857 mlir::linalg::fuseTensorOps(PatternRewriter &rewriter, Operation *consumer, 858 unsigned consumerIdx, OperationFolder *folder) { 859 if (consumerIdx >= consumer->getNumOperands()) 860 return llvm::None; 861 Operation *producer = consumer->getOperand(consumerIdx).getDefiningOp(); 862 if (!producer || producer->getNumResults() != 1) 863 return llvm::None; 864 865 // Fuse when consumer is GenericOp or IndexedGenericOp. 866 if (!isa<GenericOp, IndexedGenericOp>(consumer) || 867 !isa<GenericOp, IndexedGenericOp>(producer)) 868 return llvm::None; 869 870 return fuseTensorOpsImpl(cast<LinalgOp>(producer), cast<LinalgOp>(consumer), 871 consumerIdx, rewriter, folder); 872 } 873 874 namespace { 875 /// Patterns to fuse a generic op, with the producer of its operands. 876 template <typename LinalgOpTy> 877 struct FuseTensorOps : public OpRewritePattern<LinalgOpTy> { 878 using OpRewritePattern<LinalgOpTy>::OpRewritePattern; 879 880 LogicalResult matchAndRewrite(LinalgOpTy op, 881 PatternRewriter &rewriter) const override { 882 // Find the first operand that is defined by another generic op on tensors. 883 for (auto operandNum : 884 llvm::seq<unsigned>(0, op.getOperation()->getNumOperands())) { 885 Operation *producer = 886 op.getOperation()->getOperand(operandNum).getDefiningOp(); 887 if (!producer) 888 continue; 889 Optional<SmallVector<Value, 1>> fusedOpResults = 890 fuseTensorOps(rewriter, op, operandNum); 891 if (fusedOpResults) { 892 rewriter.replaceOp(op, *fusedOpResults); 893 if (producer->use_empty()) 894 rewriter.eraseOp(producer); 895 return success(); 896 } 897 } 898 return failure(); 899 } 900 }; 901 902 /// Pass that fuses generic ops on tensors. Used only for testing. 903 struct FusionOfTensorOpsPass 904 : public LinalgFusionOfTensorOpsBase<FusionOfTensorOpsPass> { 905 void runOnOperation() override { 906 OwningRewritePatternList patterns; 907 Operation *op = getOperation(); 908 populateLinalgTensorOpsFusionPatterns(op->getContext(), patterns); 909 applyPatternsAndFoldGreedily(op->getRegions(), patterns); 910 } 911 }; 912 913 /// Pass to test folding of reshape op with generic/indexed_generic ops by 914 /// linearization. 915 struct FoldReshapeOpsByLinearizationPass 916 : public LinalgFoldReshapeOpsByLinearizationBase< 917 FoldReshapeOpsByLinearizationPass> { 918 void runOnOperation() override { 919 OwningRewritePatternList patterns; 920 Operation *op = getOperation(); 921 populateFoldReshapeOpsByLinearizationPatterns(op->getContext(), patterns); 922 applyPatternsAndFoldGreedily(op->getRegions(), patterns); 923 } 924 }; 925 926 } // namespace 927 928 void mlir::populateFoldReshapeOpsByLinearizationPatterns( 929 MLIRContext *context, OwningRewritePatternList &patterns) { 930 patterns.insert<FoldProducerReshapeOpByLinearization<GenericOp>, 931 FoldProducerReshapeOpByLinearization<IndexedGenericOp>, 932 FoldConsumerReshapeOpByLinearization>(context); 933 } 934 935 void mlir::populateFoldReshapeOpsByExpansionPatterns( 936 MLIRContext *context, OwningRewritePatternList &patterns) { 937 patterns.insert<FoldReshapeWithGenericOpByExpansion, 938 FoldWithProducerReshapeOpByExpansion>(context); 939 } 940 941 void mlir::populateLinalgTensorOpsFusionPatterns( 942 MLIRContext *context, OwningRewritePatternList &patterns) { 943 patterns.insert<FuseTensorOps<GenericOp>, FuseTensorOps<IndexedGenericOp>, 944 FoldSplatConstants<GenericOp>, 945 FoldSplatConstants<IndexedGenericOp>>(context); 946 populateFoldReshapeOpsByExpansionPatterns(context, patterns); 947 GenericOp::getCanonicalizationPatterns(patterns, context); 948 IndexedGenericOp::getCanonicalizationPatterns(patterns, context); 949 TensorReshapeOp::getCanonicalizationPatterns(patterns, context); 950 } 951 952 std::unique_ptr<Pass> mlir::createLinalgFusionOfTensorOpsPass() { 953 return std::make_unique<FusionOfTensorOpsPass>(); 954 } 955 956 std::unique_ptr<Pass> mlir::createFoldReshapeOpsByLinearizationPass() { 957 return std::make_unique<FoldReshapeOpsByLinearizationPass>(); 958 } 959